62 Commits

Author SHA1 Message Date
benvin afa056b454 Merge pull request 'Add boot unit applying a local config' (#34) from benvin/boot-unit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #34
2026-10-05 21:46:58 +11:00
benvin f593f7625d Merge pull request 'Revert agent generations that cut off the control plane' (#35) from benvin/agent-safe-apply into main
Reviewed-on: #35
2026-10-05 21:46:08 +11:00
benvin 7c8bd87ec0 Merge pull request 'ci: use container-rpmbuilder image' (#36) from benvin/rpmbuilder-image into main
Reviewed-on: #36
2026-10-05 21:34:18 +11:00
unkin-agent 9854b0e7b6 ci: use container-rpmbuilder image
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 14:39:45 +11:00
unkin-agent c7e02c089a Apply the cached config without safe-apply and keep reverted generations in memory
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:56:11 +11:00
unkin-agent 9092b463a0 Apply agent generations as a pending try with a revert timer
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:52:32 +11:00
unkin-agent 502d06bdda Cap boot unit restarts so it fails open
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:50 +11:00
unkin-agent 4ad55fc65e Revert agent generations that cut off the control plane
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:20 +11:00
unkin-agent 7fcbb5fad8 Order boot unit before sysinit and retry on failure
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:44:58 +11:00
unkin-agent dc406c4f56 Add boot unit applying a local config
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:41:57 +11:00
benvin 2832dd0e7d Merge pull request 'Rate-limit log sites with LOGLIMIT' (#33) from benvin/loglimit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #33
2026-10-04 18:19:06 +11:00
unkin-agent 15ab32431d Drop the name from named shorewall LOGLIMIT values
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:58:12 +11:00
unkin-agent be391ed385 Keep rule match extras on the limited log rule
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:56:23 +11:00
unkin-agent 2e8d51759d Rate-limit log sites with shorewall LOGLIMIT
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:52:37 +11:00
benvin b8410488d2 Merge pull request 'Create the nftables table in the configured address family' (#32) from benvin/v4-only-family into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #32
2026-10-04 15:43:03 +11:00
benvin 200b4d3bdf Merge pull request 'Honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION' (#31) from benvin/invalid-disposition into main
Reviewed-on: #31
2026-10-04 15:42:40 +11:00
unkin-agent d6dfeb62b6 Merge remote-tracking branch 'origin/main' into benvin/v4-only-family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/engine.go
2026-10-04 15:40:58 +11:00
benvin 11fafc0e6a Merge pull request 'Fix large-batch netlink errors and revert failed try applies' (#30) from benvin/try-apply-error into main
Reviewed-on: #30
2026-10-04 15:40:08 +11:00
unkin-agent d949fc4772 Create the nftables table in the configured address family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:36:03 +11:00
unkin-agent 2ed5b958b4 default dispositions to continue before the nil-conf return; test bad disposition values
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:35:58 +11:00
unkin-agent c3049e3ed4 honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:34:11 +11:00
unkin-agent fe689e99ed Raise netlink buffers for large batches; restore snapshot when try apply fails
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
2026-10-04 15:34:08 +11:00
benvin 1e32bde555 Merge pull request 'Emit the implied ACCEPT for DNAT and REDIRECT rules' (#29) from benvin/dnat-implied-accept into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #29
2026-10-04 15:28:05 +11:00
unkin-agent 87be6bc6af Merge remote-tracking branch 'origin/main' into benvin/dnat-implied-accept
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
2026-10-04 15:22:41 +11:00
benvin 599c483cd8 Merge pull request 'Accept DHCP on dhcp interfaces as shorewall does' (#27) from benvin/dhcp-option into main
Reviewed-on: #27
2026-10-04 15:21:51 +11:00
benvin 110b109d97 Merge pull request 'Include firewall zone in all/any rule expansion' (#28) from benvin/all-includes-fw into main
Reviewed-on: #28
2026-10-04 15:21:24 +11:00
unkin-agent ac63f65f2f Keep rule extras off the DNAT implied accept, expand all/any DNAT sources and match SPORT
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:12:29 +11:00
unkin-agent e3130b6c3b Expand all/any firewall pairs per zone and dedupe fw
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:10:49 +11:00
unkin-agent df9326330a Count firewall-expanded all/any specs for limiter guard
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:08:48 +11:00
unkin-agent b29ed2446e Emit the implied ACCEPT for DNAT and REDIRECT rules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:08:11 +11:00
unkin-agent 0551120eec Scope DHCP accept rules to IPv4
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:07:41 +11:00
unkin-agent ff4c9b63e8 Include firewall zone in all/any rule expansion
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:06:24 +11:00
unkin-agent e48e9079bd Accept DHCP on dhcp interfaces as shorewall does
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:06:11 +11:00
benvin 2dd408c54e Merge pull request 'Create a Gitea release on tag' (#26) from benvin/gitea-releases into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #26
2026-10-04 15:00:18 +11:00
unkin-agent 6fc8ae9256 Run release step on prebuilt tea image
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 14:51:14 +11:00
benvin b849626a8c Merge pull request 'Attach conntrack helpers via ct helper objects' (#25) from benvin/ct-helpers into main
Reviewed-on: #25
2026-10-04 14:17:23 +11:00
unkin-agent da32729556 Create a Gitea release with binary, RPM and checksums on tag
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 13:59:34 +11:00
unkin-agent b85ba61fe6 Merge remote-tracking branch 'origin/main' into benvin/ct-helpers
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
#	internal/nftables/engine.go
2026-10-04 13:58:25 +11:00
benvin 9abaa85d6a Merge pull request 'Compile conntrack rules into raw-priority chains' (#22) from benvin/notrack-raw into main
Reviewed-on: #22
2026-10-04 13:55:36 +11:00
unkin-agent 9078f410ee Reject conntrack fw DEST without an address in prerouting
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:56:29 +10:00
unkin-agent da65c5d0a8 Treat conntrack SOURCE all/any as global like an omitted source
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:54:37 +10:00
unkin-agent 90d2875301 Reject zone exclusions in conntrack entries
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:53:09 +10:00
unkin-agent c557c4b78a Attach conntrack helpers via ct helper objects
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-03 23:50:31 +10:00
unkin-agent bf235291d0 Fail closed on unknown/none conntrack DEST and pair all+ symmetrically
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:50:18 +10:00
unkin-agent 1c379bf5f1 Merge remote-tracking branch 'origin/main' into benvin/notrack-raw 2026-10-03 23:49:40 +10:00
unkin-agent b6d67897de Expand all!zone exclusions and fail closed on unknown zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-03 23:48:08 +10:00
unkin-agent 9432bb05c9 Reject unknown and excluded zone names in rule, blrule and conntrack specs 2026-10-03 23:48:08 +10:00
benvin 9a7d1ab35b Merge pull request 'Bump google/nftables to v0.3.0' (#23) from benvin/nftables-v0.3.0 into main
Reviewed-on: #23
2026-10-03 23:45:34 +10:00
benvin 3c5d23a811 Merge pull request 'Match source/dest zones on conntrack rules' (#24) from benvin/conntrack-zones into benvin/notrack-raw
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Reviewed-on: #24
2026-10-03 23:44:11 +10:00
unkin-agent 337490d995 Reject fw source on prerouting conntrack and pin global omitted-zone entries
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:01:29 +10:00
unkin-agent 6af17a4c02 Restrict raw_output to fw sources and reject unmatched prerouting dest zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:59:14 +10:00
unkin-agent 3c5f1cacd4 Match source/dest zones and addresses on conntrack rules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:56:26 +10:00
unkin-agent 445d14c61e bump google/nftables to v0.3.0
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
Set NAT.Specified on DNAT with a port so compiled exprs match kernel readback.
2026-10-03 22:56:20 +10:00
unkin-agent c9adff32f6 Compile conntrack rules into raw-priority chains
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:54:34 +10:00
benvin 63c46fec81 Merge pull request 'Match dest zone oif on output-chain rules and policies' (#20) from benvin/output-oif into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #20
2026-10-03 22:51:28 +10:00
unkin-agent 5238886c6e Merge remote-tracking branch 'origin/main' into benvin/output-oif
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler_test.go
2026-10-03 22:46:09 +10:00
benvin 532bd80a8a Merge pull request 'Match ORIGDEST in DNAT and filter rules' (#21) from benvin/origdest into main
Reviewed-on: #21
2026-10-03 22:44:37 +10:00
unkin-agent 036021d726 Treat only non-negated addresses as scoping interfaceless zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:20:41 +10:00
unkin-agent ff5a52b9e2 Reject ORIGDEST on forwarded rules and mixed-family lists
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:18:48 +10:00
unkin-agent 44e1ba852e Skip rules and policies for zones without interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:18:10 +10:00
unkin-agent 34cf6dc9ab Match ORIGDEST in DNAT and filter rules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:16:40 +10:00
unkin-agent b5be665902 Match dest zone oif on output-chain rules and policies
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:15:24 +10:00
36 changed files with 3166 additions and 383 deletions
+43 -1
View File
@@ -59,7 +59,7 @@ steps:
cpu: 2 cpu: 2
- name: package - name: package
image: git.unkin.net/unkin/almalinux9-rpmbuilder:latest image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/rpmbuilder:0.1.0-alma9
commands: commands:
- ./scripts/build-rpm.sh ${CI_COMMIT_TAG} - ./scripts/build-rpm.sh ${CI_COMMIT_TAG}
depends_on: [build] depends_on: [build]
@@ -104,3 +104,45 @@ steps:
limits: limits:
memory: 512Mi memory: 512Mi
cpu: 500m cpu: 500m
- name: release
image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/tea:0.1.0
environment:
RELEASER_TOKEN:
from_secret: RELEASER_TOKEN
commands:
- |
tea logins add --name gitea --url https://git.unkin.net --token "$${RELEASER_TOKEN}" --no-version-check
CUR_SHA=$$(git rev-list -n1 "${CI_COMMIT_TAG}")
PREV_TAG=""
for t in $$(git tag --sort=-v:refname); do
[ "$$t" = "${CI_COMMIT_TAG}" ] && continue
[ "$$(git rev-list -n1 "$$t")" = "$$CUR_SHA" ] && continue
if git merge-base --is-ancestor "$$t" "${CI_COMMIT_TAG}" 2>/dev/null; then
PREV_TAG="$$t"; break
fi
done
if [ -n "$$PREV_TAG" ]; then
NOTES=$$(git log "$${PREV_TAG}..${CI_COMMIT_TAG}" --pretty=format:"- %s")
else
NOTES=$$(git log --pretty=format:"- %s")
fi
tea releases create --tag "${CI_COMMIT_TAG}" --title "${CI_COMMIT_TAG}" --note "$${NOTES}" --login gitea --repo "${CI_REPO}"
cp dist/tomswall tomswall-linux-amd64
ASSETS="tomswall-linux-amd64"
RPM=$$(ls dist/*.rpm 2>/dev/null | head -1)
[ -n "$$RPM" ] && ASSETS="$$ASSETS $$RPM"
sha256sum $$ASSETS > sha256sums.txt
tea releases assets create "${CI_COMMIT_TAG}" $$ASSETS sha256sums.txt \
--login gitea --repo "${CI_REPO}"
depends_on: [upload-rpm]
backend_options:
kubernetes:
serviceAccountName: default
resources:
requests:
memory: 128Mi
cpu: 100m
limits:
memory: 512Mi
cpu: 500m
+7
View File
@@ -430,6 +430,13 @@ report the generation applied, giving a fleet-wide "converged / N behind" view.
source/dest disables its rule loudly, never opens it. source/dest disables its rule loudly, never opens it.
- **Adds fail closed, the control plane fails open.** Partial rollout blocks new - **Adds fail closed, the control plane fails open.** Partial rollout blocks new
flows until every hop converges; a dead API leaves the last-good posture running. flows until every hop converges; a dead API leaves the last-good posture running.
- **A generation that severs the API is reverted.** The agent applies as a
`tomswall try` does (on-disk snapshot, systemd revert timer, shared lock), then
reports `applied` over a fresh connection. If that fails at the transport level,
or the apply errors, it records the generation in
`/var/lib/tomswall/reverted.json` (skipped until a newer one arrives), restores
the snapshot and reports `reverted`/`failed`. A failed restore leaves the timer
to retry it and is reported `failed`.
--- ---
+2 -1
View File
@@ -29,7 +29,8 @@ func agentCmd() *cobra.Command {
Long: `Agent runs the control-plane pull loop: it fetches this device's compiled Long: `Agent runs the control-plane pull loop: it fetches this device's compiled
config from tomswallapi, differentially applies it, and reports the applied config from tomswallapi, differentially applies it, and reports the applied
generation back. It caches the last known-good config and, if the control plane generation back. It caches the last known-good config and, if the control plane
is unreachable, keeps applying that cache — it never fails closed. is unreachable, keeps applying that cache — it never fails closed. A new
generation that cuts the agent off from the API is reverted and reported as such.
The agent token defaults to the TOMSWALL_AGENT_TOKEN environment variable, and The agent token defaults to the TOMSWALL_AGENT_TOKEN environment variable, and
the device name defaults to the system hostname.`, the device name defaults to the system hostname.`,
+3 -3
View File
@@ -89,10 +89,10 @@ func tryApply(cfg *config.Config, fallback time.Duration) (string, error) {
return "", err return "", err
} }
if err := engine.Apply(changes); err != nil { if err := engine.Apply(changes); err != nil {
if derr := tryapply.Discard(); derr != nil { if aerr := tryapply.Abort(); aerr != nil {
err = fmt.Errorf("%w (discarding snapshot: %v)", err, derr) return "", fmt.Errorf("applying changes: %w; %v; the revert timer restores the previous ruleset within %s", err, aerr, fallback)
} }
return "", fmt.Errorf("applying changes: %w", err) return "", fmt.Errorf("applying changes: %w: previous ruleset restored", err)
} }
return id, nil return id, nil
} }
+4 -5
View File
@@ -3,19 +3,18 @@ module git.unkin.net/unkin/tomswall
go 1.23 go 1.23
require ( require (
github.com/google/nftables v0.2.0 github.com/google/nftables v0.3.0
github.com/mdlayher/netlink v1.7.2 github.com/mdlayher/netlink v1.7.3-0.20250113171957-fbb4dce95f42
github.com/spf13/cobra v1.8.1 github.com/spf13/cobra v1.8.1
golang.org/x/sys v0.18.0 golang.org/x/sys v0.28.0
gopkg.in/yaml.v3 v3.0.1 gopkg.in/yaml.v3 v3.0.1
) )
require ( require (
github.com/google/go-cmp v0.6.0 // indirect github.com/google/go-cmp v0.6.0 // indirect
github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/josharian/native v1.1.0 // indirect
github.com/mdlayher/socket v0.5.1 // indirect github.com/mdlayher/socket v0.5.1 // indirect
github.com/spf13/pflag v1.0.5 // indirect github.com/spf13/pflag v1.0.5 // indirect
golang.org/x/net v0.23.0 // indirect golang.org/x/net v0.33.0 // indirect
golang.org/x/sync v0.6.0 // indirect golang.org/x/sync v0.6.0 // indirect
) )
+10 -12
View File
@@ -1,14 +1,12 @@
github.com/cpuguy83/go-md2man/v2 v2.0.4/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o= github.com/cpuguy83/go-md2man/v2 v2.0.4/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o=
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/nftables v0.2.0 h1:PbJwaBmbVLzpeldoeUKGkE2RjstrjPKMl6oLrfEJ6/8= github.com/google/nftables v0.3.0 h1:bkyZ0cbpVeMHXOrtlFc8ISmfVqq5gPJukoYieyVmITg=
github.com/google/nftables v0.2.0/go.mod h1:Beg6V6zZ3oEn0JuiUQ4wqwuyqqzasOltcoXPtgLbFp4= github.com/google/nftables v0.3.0/go.mod h1:BCp9FsrbF1Fn/Yu6CLUc9GGZFw/+hsxfluNXXmxBfRM=
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
github.com/josharian/native v1.1.0 h1:uuaP0hAbW7Y4l0ZRQ6C9zfb7Mg1mbFKry/xzDAfmtLA= github.com/mdlayher/netlink v1.7.3-0.20250113171957-fbb4dce95f42 h1:A1Cq6Ysb0GM0tpKMbdCXCIfBclan4oHk1Jb+Hrejirg=
github.com/josharian/native v1.1.0/go.mod h1:7X/raswPFr05uY3HiLlYeyQntB6OO7E/d2Cu7qoaN2w= github.com/mdlayher/netlink v1.7.3-0.20250113171957-fbb4dce95f42/go.mod h1:BB4YCPDOzfy7FniQ/lxuYQ3dgmM2cZumHbK8RpTjN2o=
github.com/mdlayher/netlink v1.7.2 h1:/UtM3ofJap7Vl4QWCPDGXY8d3GIY2UGSDbK+QWmY8/g=
github.com/mdlayher/netlink v1.7.2/go.mod h1:xraEF7uJbxLhc5fpHL4cPe221LI2bdttWlU+ZGLfQSw=
github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos= github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos=
github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ= github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ=
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
@@ -16,14 +14,14 @@ github.com/spf13/cobra v1.8.1 h1:e5/vxKd/rZsfSJMUX1agtjeTDf+qv1/JdBF8gg5k9ZM=
github.com/spf13/cobra v1.8.1/go.mod h1:wHxEcudfqmLYa8iTfL+OuZPbBZkmvliBWKIezN3kD9Y= github.com/spf13/cobra v1.8.1/go.mod h1:wHxEcudfqmLYa8iTfL+OuZPbBZkmvliBWKIezN3kD9Y=
github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA=
github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
github.com/vishvananda/netns v0.0.0-20180720170159-13995c7128cc h1:R83G5ikgLMxrBvLh22JhdfI8K6YXEPHx5P03Uu3DRs4= github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=
github.com/vishvananda/netns v0.0.0-20180720170159-13995c7128cc/go.mod h1:ZjcWmFBXmLKZu9Nxj3WKYEafiSqer2rnvPr0en9UNpI= github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
golang.org/x/net v0.23.0 h1:7EYJ93RZ9vYSZAIb2x3lnuvqO5zneoD6IvWjuhfxjTs= golang.org/x/net v0.33.0 h1:74SYHlV8BIgHIFC/LrYkOGIwL19eTYXQ5wc6TBuO36I=
golang.org/x/net v0.23.0/go.mod h1:JKghWKKOSdJwpW2GEx0Ja7fmaKnMsbu+MWVZTokSYmg= golang.org/x/net v0.33.0/go.mod h1:HXLR5J+9DxmrqMwG9qjGCxZ+zKXxBru04zlTvWlWuN4=
golang.org/x/sync v0.6.0 h1:5BMeUDZ7vkXGfEr1x9B4bRcTH4lpkTkpdh0T/J+qjbQ= golang.org/x/sync v0.6.0 h1:5BMeUDZ7vkXGfEr1x9B4bRcTH4lpkTkpdh0T/J+qjbQ=
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sys v0.18.0 h1:DBdB3niSjOA/O0blCZBqDefyWNYveAYMNF1Wum0DYQ4= golang.org/x/sys v0.28.0 h1:Fksou7UEQUWlKvIdsqzJmUmCX3cZuD2+P3XyyzwMhlA=
golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
+226 -24
View File
@@ -2,8 +2,13 @@ package agent
import ( import (
"context" "context"
"encoding/json"
"errors"
"fmt" "fmt"
"log/slog" "log/slog"
"net/url"
"os"
"path/filepath"
"time" "time"
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
@@ -12,11 +17,17 @@ import (
) )
// Applier applies a translated config to the firewall. Abstracted so the run // Applier applies a translated config to the firewall. Abstracted so the run
// loop is testable without touching the kernel. // loop is testable without touching the kernel. With safe, a change is applied
// as a pending try: revert restores the previous ruleset (a failed revert leaves
// the revert timer armed) and keep drops the snapshot. Both are nil when nothing
// changed or safe is false.
type Applier interface { type Applier interface {
Apply(ctx context.Context, cfg *config.Config) error Apply(ctx context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error)
} }
// revertDelay is when the revert timer fires if the agent dies mid-apply.
const revertDelay = time.Minute
// Agent runs the pull-apply-report loop for one device. // Agent runs the pull-apply-report loop for one device.
type Agent struct { type Agent struct {
Client *Client Client *Client
@@ -25,6 +36,9 @@ type Agent struct {
Applier Applier Applier Applier
// Resolver overrides the DNS resolver (tests); nil derives it per-config. // Resolver overrides the DNS resolver (tests); nil derives it per-config.
Resolver *Resolver Resolver *Resolver
// lastReverted covers a reverted generation whose persistence failed.
lastReverted *reverted
} }
// Run loops until ctx is cancelled, applying one cycle per Interval (and once // Run loops until ctx is cancelled, applying one cycle per Interval (and once
@@ -62,16 +76,31 @@ func (a *Agent) RunOnce(ctx context.Context) error {
return fmt.Errorf("control plane unreachable and no cached config: %w", err) return fmt.Errorf("control plane unreachable and no cached config: %w", err)
} }
// Re-apply last known-good; do not report a generation we didn't fetch. // Re-apply last known-good; do not report a generation we didn't fetch.
return a.applyConfig(ctx, cached, false) return a.applyConfig(ctx, cached, nil)
} }
if err := a.Cache.Write(raw); err != nil { rv, err := a.readReverted()
slog.Warn("agent: caching config failed", "err", err) if err != nil {
return err
} }
return a.applyConfig(ctx, rc, true) if a.lastReverted != nil && (rv == nil || a.lastReverted.Generation > rv.Generation) {
rv = a.lastReverted
}
if rv != nil {
a.reportReverted(ctx, rv)
if rc.Generation <= rv.Generation {
slog.Info("agent: generation was reverted, waiting for a newer one", "generation", rc.Generation)
return nil
}
}
return a.applyConfig(ctx, rc, raw)
} }
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error { // applyConfig applies rc. A fetched config (raw != nil) is applied as a pending
// try, verified by reaching the API through the new ruleset and reverted if that
// fails; only then is it cached. The cached config is the last verified-good one,
// so it is applied plainly: there is nothing to verify it against.
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error {
resolver := a.Resolver resolver := a.Resolver
if resolver == nil { if resolver == nil {
resolver = NewResolver(rc.Resolver) resolver = NewResolver(rc.Resolver)
@@ -82,14 +111,65 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
if err != nil { if err != nil {
return fmt.Errorf("translate: %w", err) return fmt.Errorf("translate: %w", err)
} }
if err := a.Applier.Apply(ctx, cfg); err != nil {
return fmt.Errorf("apply: %w", err) unlock, err := tryapply.Acquire()
if errors.Is(err, tryapply.ErrPending) {
slog.Warn("agent: a 'tomswall try' is pending, skipping cycle")
return nil
}
if err != nil {
return err
}
defer unlock()
revert, keep, err := a.Applier.Apply(ctx, cfg, raw != nil)
if err != nil {
err = fmt.Errorf("apply: %w", err)
if raw != nil && revert != nil {
return a.revertGeneration(ctx, rc.Generation, StatusFailed, err, revert)
}
if revert != nil {
if rerr := revert(); rerr != nil {
err = fmt.Errorf("%w; restore: %v; revert timer pending", err, rerr)
}
}
if raw != nil {
if rerr := a.Client.ReportStatus(ctx, Status{Status: StatusFailed, Generation: rc.Generation, Error: err.Error()}); rerr != nil {
slog.Warn("agent: reporting status failed", "err", rerr)
}
}
return err
}
if raw == nil {
slog.Info("agent: applied cached config", "generation", rc.Generation, "rules", len(cfg.Rules))
return nil
}
if err := a.confirm(ctx, rc.Generation); err != nil {
if keep != nil && ctx.Err() == nil {
return a.revertGeneration(ctx, rc.Generation, StatusReverted, err, revert)
}
// Shutdown is not a verdict on the generation: keep it.
if keep != nil {
if kerr := keep(); kerr != nil {
slog.Warn("agent: dropping snapshot failed", "err", kerr)
}
}
return err
}
if keep != nil {
if err := keep(); err != nil {
return fmt.Errorf("dropping snapshot: %w", err)
}
} }
slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules))
if report { if err := a.Cache.Write(raw); err != nil {
if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { slog.Warn("agent: caching config failed", "err", err)
slog.Warn("agent: reporting status failed", "err", err) }
a.lastReverted = nil
if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) {
slog.Warn("agent: clearing reverted generation failed", "err", err)
} }
// Report the FIB so the control plane can scope router enforcement. // Report the FIB so the control plane can scope router enforcement.
if fib := CollectFIB(ctx); len(fib) > 0 { if fib := CollectFIB(ctx); len(fib) > 0 {
@@ -97,31 +177,153 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
slog.Warn("agent: reporting routes failed", "err", err) slog.Warn("agent: reporting routes failed", "err", err)
} }
} }
return nil
}
// revertGeneration marks generation as reverted before restoring, so a failed
// restore can never lead to re-applying it, then reports status. It is also kept
// in memory in case persisting fails. A failed
// restore leaves the snapshot and timer armed and is reported as failed.
func (a *Agent) revertGeneration(ctx context.Context, generation int64, status string, cause error, revert func() error) error {
rv := &reverted{Generation: generation, Status: status, Error: cause.Error()}
a.lastReverted = rv
if err := a.writeReverted(rv); err != nil {
slog.Error("agent: persisting reverted generation failed", "err", err)
}
suffix := ""
if rerr := revert(); rerr != nil {
suffix = fmt.Sprintf("; restore: %v; revert timer pending", rerr)
rv.Status = StatusFailed
rv.Error += suffix
if werr := a.writeReverted(rv); werr != nil {
slog.Error("agent: persisting reverted generation failed", "err", werr)
}
}
a.reportReverted(ctx, rv)
return fmt.Errorf("generation %d %s: %w%s", generation, rv.Status, cause, suffix)
}
var (
verifyAttempts = 3
verifyDelay = 2 * time.Second
verifyTimeout = 5 * time.Second
)
// errUnreachable means the API could not be reached through the new ruleset.
var errUnreachable = errors.New("control plane unreachable after apply")
// confirm reports generation as applied over a fresh connection, which proves
// the API is reachable through the new ruleset. Any HTTP response counts as
// reachable; only repeated transport failures return errUnreachable.
func (a *Agent) confirm(ctx context.Context, generation int64) error {
var err error
for i := 0; i < verifyAttempts; i++ {
if i > 0 {
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(verifyDelay):
}
}
actx, cancel := context.WithTimeout(ctx, verifyTimeout)
err = a.Client.ReportStatus(actx, Status{Status: StatusApplied, Generation: generation})
cancel()
if ctx.Err() != nil {
return ctx.Err()
}
var uerr *url.Error
if !errors.As(err, &uerr) {
if err != nil {
slog.Warn("agent: reporting status failed", "err", err)
} }
return nil return nil
}
}
return fmt.Errorf("%w: %v", errUnreachable, err)
}
// reverted is a generation rolled back after a failed apply or for severing the
// API; persisted so it is not re-applied until a newer generation arrives.
type reverted struct {
Generation int64 `json:"generation"`
Status string `json:"status,omitempty"`
Error string `json:"error,omitempty"`
Reported bool `json:"reported"`
}
func (a *Agent) revertedPath() string {
return filepath.Join(filepath.Dir(a.Cache.Path), "reverted.json")
}
func (a *Agent) readReverted() (*reverted, error) {
b, err := os.ReadFile(a.revertedPath())
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
var rv reverted
if err := json.Unmarshal(b, &rv); err != nil {
return nil, fmt.Errorf("parsing %s: %w", a.revertedPath(), err)
}
return &rv, nil
}
func (a *Agent) writeReverted(rv *reverted) error {
b, err := json.Marshal(rv)
if err != nil {
return err
}
return tryapply.WriteFile(a.revertedPath(), b)
}
// reportReverted reports rv until the control plane accepts it.
func (a *Agent) reportReverted(ctx context.Context, rv *reverted) {
if rv.Reported {
return
}
status := rv.Status
if status == "" {
status = StatusReverted
}
if err := a.Client.ReportStatus(ctx, Status{Status: status, Generation: rv.Generation, Error: rv.Error}); err != nil {
slog.Warn("agent: reporting reverted generation failed, retrying next cycle", "generation", rv.Generation, "err", err)
return
}
rv.Reported = true
if err := a.writeReverted(rv); err != nil {
slog.Warn("agent: persisting reverted generation failed", "err", err)
}
} }
// EngineApplier applies via the real nftables differential engine. // EngineApplier applies via the real nftables differential engine.
type EngineApplier struct{} type EngineApplier struct{}
// Apply computes and applies the differential change set for cfg. It refuses // Apply computes and applies the differential change set for cfg, with safe
// while a 'tomswall try' awaits confirmation. // under a pending try as 'tomswall try' does. The caller holds the try lock.
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { func (EngineApplier) Apply(_ context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error) {
unlock, err := tryapply.Acquire()
if err != nil {
return err
}
defer unlock()
engine, err := nftables.NewEngine(cfg) engine, err := nftables.NewEngine(cfg)
if err != nil { if err != nil {
return fmt.Errorf("initializing nftables: %w", err) return nil, nil, fmt.Errorf("initializing nftables: %w", err)
} }
changes, err := engine.Plan() changes, err := engine.Plan()
if err != nil { if err != nil {
return fmt.Errorf("computing changes: %w", err) return nil, nil, fmt.Errorf("computing changes: %w", err)
} }
if changes.Empty() { if changes.Empty() {
return nil return nil, nil, nil
} }
return engine.Apply(changes) if !safe {
return nil, nil, engine.Apply(changes)
}
snap, err := engine.Snapshot()
if err != nil {
return nil, nil, fmt.Errorf("snapshotting ruleset: %w", err)
}
// PID 0: 'tomswall confirm' must not signal the agent.
if _, err := tryapply.Arm(snap, 0, revertDelay); err != nil {
return nil, nil, err
}
return tryapply.Abort, tryapply.Discard, engine.Apply(changes)
} }
+2 -2
View File
@@ -132,10 +132,10 @@ type fakeApplier struct {
lastGen int lastGen int
} }
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error { func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config, _ bool) (func() error, func() error, error) {
atomic.AddInt32(&f.count, 1) atomic.AddInt32(&f.count, 1)
f.lastGen = len(cfg.Rules) f.lastGen = len(cfg.Rules)
return nil return nil, nil, nil
} }
const renderedYAML = `generation: 7 const renderedYAML = `generation: 7
+4 -10
View File
@@ -2,7 +2,8 @@ package agent
import ( import (
"os" "os"
"path/filepath"
"git.unkin.net/unkin/tomswall/internal/tryapply"
) )
// Cache persists the last known-good rendered config to disk so the agent can // Cache persists the last known-good rendered config to disk so the agent can
@@ -11,16 +12,9 @@ type Cache struct {
Path string Path string
} }
// Write atomically stores the raw config bytes. // Write durably stores the raw config bytes.
func (c Cache) Write(raw []byte) error { func (c Cache) Write(raw []byte) error {
if err := os.MkdirAll(filepath.Dir(c.Path), 0o755); err != nil { return tryapply.WriteFile(c.Path, raw)
return err
}
tmp := c.Path + ".tmp"
if err := os.WriteFile(tmp, raw, 0o600); err != nil {
return err
}
return os.Rename(tmp, c.Path)
} }
// Read returns the cached config, or (nil, nil) when no cache exists yet. // Read returns the cached config, or (nil, nil) when no cache exists yet.
+26 -4
View File
@@ -26,7 +26,9 @@ func NewClient(baseURL, device, token string) *Client {
BaseURL: baseURL, BaseURL: baseURL,
Device: device, Device: device,
Token: token, Token: token,
HTTP: &http.Client{Timeout: 30 * time.Second}, // No keep-alives: every request, the post-apply check included, opens a
// fresh connection that must pass the current ruleset.
HTTP: &http.Client{Timeout: 30 * time.Second, Transport: noKeepAlive()},
} }
} }
@@ -94,10 +96,24 @@ func (c *Client) ReportRoutes(ctx context.Context, prefixes []string) error {
return nil return nil
} }
// ReportStatus tells the control plane which generation this device has applied. // Status values reported to POST /api/v1/devices/{name}/status.
func (c *Client) ReportStatus(ctx context.Context, generation int64) error { const (
StatusApplied = "applied"
StatusReverted = "reverted"
StatusFailed = "failed"
)
// Status is the outcome of applying one generation.
type Status struct {
Status string `json:"status"`
Generation int64 `json:"generation"`
Error string `json:"error,omitempty"`
}
// ReportStatus tells the control plane the outcome of applying a generation.
func (c *Client) ReportStatus(ctx context.Context, st Status) error {
url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device) url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device)
payload, _ := json.Marshal(map[string]int64{"generation": generation}) payload, _ := json.Marshal(st)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload)) req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
if err != nil { if err != nil {
return err return err
@@ -116,3 +132,9 @@ func (c *Client) ReportStatus(ctx context.Context, generation int64) error {
} }
return nil return nil
} }
func noKeepAlive() http.RoundTripper {
t := http.DefaultTransport.(*http.Transport).Clone()
t.DisableKeepAlives = true
return t
}
+396
View File
@@ -0,0 +1,396 @@
package agent
import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/tryapply"
)
func TestMain(m *testing.M) {
dir, err := os.MkdirTemp("", "tomswall-agent-test")
if err != nil {
panic(err)
}
tryapply.Dir = dir
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
return nil
}
verifyDelay = time.Millisecond
verifyTimeout = time.Second
code := m.Run()
os.RemoveAll(dir)
os.Exit(code)
}
// fakeAPI serves a config generation and records status reports; while cut it
// drops connections to the status endpoint, as a severing ruleset would.
type fakeAPI struct {
*httptest.Server
gen atomic.Int64
cut atomic.Bool
code atomic.Int32
mu sync.Mutex
reports []Status
}
func newFakeAPI(t *testing.T, gen int64) *fakeAPI {
f := &fakeAPI{}
f.gen.Store(gen)
f.code.Store(http.StatusNoContent)
f.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v1/devices/fw-a/config":
_, _ = w.Write([]byte(strings.Replace(renderedYAML, "generation: 7", "generation: "+itoa(f.gen.Load()), 1)))
case "/api/v1/devices/fw-a/status":
if f.cut.Load() {
conn, _, _ := w.(http.Hijacker).Hijack()
conn.Close()
return
}
var st Status
_ = json.NewDecoder(r.Body).Decode(&st)
f.mu.Lock()
f.reports = append(f.reports, st)
f.mu.Unlock()
w.WriteHeader(int(f.code.Load()))
default:
w.WriteHeader(http.StatusNotFound)
}
}))
t.Cleanup(f.Close)
return f
}
// timerCmds records the systemd commands tryapply runs.
var timerCmds []string
func itoa(n int64) string { b, _ := json.Marshal(n); return string(b) }
func (f *fakeAPI) last() Status {
f.mu.Lock()
defer f.mu.Unlock()
if len(f.reports) == 0 {
return Status{}
}
return f.reports[len(f.reports)-1]
}
// fakeEngine always changes the ruleset, when safe under a real tryapply pending
// try; onApply simulates its effect and restoreErr fails the restore.
type fakeEngine struct {
applies, plain, restores int
err, restoreErr error
onApply func()
onRestore func()
}
func (f *fakeEngine) Apply(_ context.Context, _ *config.Config, safe bool) (func() error, func() error, error) {
if !safe {
f.plain++
return nil, nil, f.err
}
if _, err := tryapply.Arm(&nftables.Snapshot{Table: "tomswall"}, 0, time.Minute); err != nil {
return nil, nil, err
}
tryapply.Restore = func(*nftables.Snapshot) error {
f.restores++
if f.onRestore != nil {
f.onRestore()
}
return f.restoreErr
}
f.applies++
if f.onApply != nil {
f.onApply()
}
return tryapply.Abort, tryapply.Discard, f.err
}
// pending reports whether a snapshot is still armed and its timer not stopped since.
func pending(t *testing.T) bool {
t.Helper()
_, err := os.Stat(filepath.Join(tryapply.Dir, "try-snapshot.json"))
armed := len(timerCmds) > 0 && timerCmds[len(timerCmds)-1] == "systemd-run"
if (err == nil) != armed {
t.Fatalf("snapshot present=%v but timer armed=%v", err == nil, armed)
}
return armed
}
func newAgent(t *testing.T, api *fakeAPI, eng *fakeEngine) *Agent {
return &Agent{
Client: NewClient(api.URL, "fw-a", "tok"),
Cache: Cache{Path: filepath.Join(t.TempDir(), "rendered.yaml")},
Applier: eng,
}
}
func cachedGen(t *testing.T, a *Agent) int64 {
rc, err := a.Cache.Read()
if err != nil {
t.Fatal(err)
}
if rc == nil {
return 0
}
return rc.Generation
}
func TestSafeApplyReachableApplies(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.restores != 0 || api.last() != (Status{Status: StatusApplied, Generation: 7}) || cachedGen(t, a) != 7 || pending(t) {
t.Fatalf("restores=%d last=%+v cache=%d", eng.restores, api.last(), cachedGen(t, a))
}
}
func TestSafeApplyUnreachableRevertsAndReports(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, onRestore: func() { api.cut.Store(false) }}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
t.Fatalf("want errUnreachable, got %v", err)
}
if eng.restores != 1 || cachedGen(t, a) != 0 || pending(t) {
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
}
if st := api.last(); st.Status != StatusReverted || st.Generation != 7 || st.Error == "" {
t.Fatalf("last report %+v", st)
}
rv, _ := a.readReverted()
if rv == nil || rv.Generation != 7 || !rv.Reported {
t.Fatalf("persisted %+v", rv)
}
}
func TestSafeApplyRevertReportedOnceReachable(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }}
a := newAgent(t, api, eng)
_ = a.RunOnce(context.Background())
if eng.restores != 1 || api.last().Status != "" {
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
}
api.cut.Store(false)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 1 || api.last() != (Status{Status: StatusReverted, Generation: 7, Error: api.last().Error}) {
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
}
}
func TestSafeApplyShutdownDoesNotRevert(t *testing.T) {
api := newFakeAPI(t, 7)
ctx, cancel := context.WithCancel(context.Background())
eng := &fakeEngine{onApply: func() { api.cut.Store(true); cancel() }}
a := newAgent(t, api, eng)
if err := a.RunOnce(ctx); !errors.Is(err, context.Canceled) {
t.Fatalf("want context.Canceled, got %v", err)
}
if rv, _ := a.readReverted(); eng.restores != 0 || rv != nil || cachedGen(t, a) != 0 {
t.Fatalf("restores=%d reverted=%+v cache=%d", eng.restores, rv, cachedGen(t, a))
}
}
func TestSafeApplyHTTPErrorDoesNotRevert(t *testing.T) {
api := newFakeAPI(t, 7)
api.code.Store(http.StatusInternalServerError)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.restores != 0 || cachedGen(t, a) != 7 {
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
}
}
func TestSafeApplyApplyErrorRestoresAndReportsFailed(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{err: errors.New("netlink: boom")}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil {
t.Fatal("want error")
}
if st := api.last(); eng.restores != 1 || st.Status != StatusFailed || !strings.Contains(st.Error, "boom") || pending(t) {
t.Fatalf("restores=%d last=%+v", eng.restores, st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported {
t.Fatalf("persisted %+v", rv)
}
}
func TestSafeApplyApplyErrorRestoreFailsKeepsTimer(t *testing.T) {
t.Cleanup(func() { _ = tryapply.Discard() })
api := newFakeAPI(t, 7)
eng := &fakeEngine{err: errors.New("netlink: boom"), restoreErr: errors.New("netlink: stuck")}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "restore: restoring snapshot: netlink: stuck") {
t.Fatalf("got %v", err)
}
want := "apply: netlink: boom; restore: restoring snapshot: netlink: stuck; revert timer pending"
if st := api.last(); st != (Status{Status: StatusFailed, Generation: 7, Error: want}) || !pending(t) {
t.Fatalf("last=%+v", st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || rv.Status != StatusFailed {
t.Fatalf("persisted %+v", rv)
}
// The next cycle waits for the timer instead of re-applying.
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
t.Fatalf("err=%v applies=%d", err, eng.applies)
}
}
func TestSafeApplyUnreachableRestoreFailsKeepsTimer(t *testing.T) {
t.Cleanup(func() { _ = tryapply.Discard() })
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, restoreErr: errors.New("netlink: stuck")}
eng.onRestore = func() { api.cut.Store(false) }
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) || !strings.Contains(err.Error(), "revert timer pending") {
t.Fatalf("got %v", err)
}
st := api.last()
if st.Status != StatusFailed || st.Generation != 7 || !strings.HasPrefix(st.Error, errUnreachable.Error()) ||
!strings.HasSuffix(st.Error, "; restore: restoring snapshot: netlink: stuck; revert timer pending") || !pending(t) {
t.Fatalf("last=%+v", st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported || cachedGen(t, a) != 0 {
t.Fatalf("persisted %+v cache=%d", rv, cachedGen(t, a))
}
}
func TestSafeApplyRevertedGenerationSkippedAfterRestart(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.writeReverted(&reverted{Generation: 7, Reported: true}); err != nil {
t.Fatal(err)
}
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 0 {
t.Fatalf("reverted generation re-applied")
}
api.gen.Store(8)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if rv, _ := a.readReverted(); eng.applies != 1 || api.last().Generation != 8 || rv != nil {
t.Fatalf("applies=%d last=%+v reverted=%+v", eng.applies, api.last(), rv)
}
}
func TestSafeApplySkipsWhileTryPending(t *testing.T) {
marker := filepath.Join(tryapply.Dir, "try-snapshot.json")
if err := os.WriteFile(marker, []byte("{}"), 0o600); err != nil {
t.Fatal(err)
}
defer os.Remove(marker)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 0 || api.last().Status != "" {
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
}
}
// failArm makes arming the revert timer fail, as without systemd.
func failArm(t *testing.T) {
orig := tryapply.Run
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
if name == "systemd-run" {
return errors.New("no systemd")
}
return nil
}
t.Cleanup(func() { tryapply.Run = orig })
}
func TestSafeApplyArmFailureReportsFailedAndRetries(t *testing.T) {
failArm(t)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "no systemd") {
t.Fatalf("got %v", err)
}
if st := api.last(); eng.applies != 0 || st.Status != StatusFailed || st.Generation != 7 || pending(t) || cachedGen(t, a) != 0 {
t.Fatalf("applies=%d last=%+v", eng.applies, st)
}
if rv, _ := a.readReverted(); rv != nil || a.lastReverted != nil {
t.Fatalf("arm failure marked generation reverted: %+v", rv)
}
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
return nil
}
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 || cachedGen(t, a) != 7 {
t.Fatalf("retry err=%v applies=%d", err, eng.applies)
}
}
func TestCachedConfigAppliesWithoutArm(t *testing.T) {
failArm(t)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.Cache.Write([]byte(renderedYAML)); err != nil {
t.Fatal(err)
}
api.Close()
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.plain != 1 || eng.applies != 0 || pending(t) {
t.Fatalf("plain=%d safe=%d", eng.plain, eng.applies)
}
}
func TestSafeApplyRevertedKeptInMemoryWhenPersistFails(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onRestore: func() { api.cut.Store(false) }}
a := newAgent(t, api, eng)
// A non-empty directory in its place makes persisting reverted.json fail.
eng.onApply = func() {
api.cut.Store(true)
_ = os.MkdirAll(filepath.Join(a.revertedPath(), "x"), 0o755)
}
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
t.Fatalf("want errUnreachable, got %v", err)
}
if err := os.RemoveAll(a.revertedPath()); err != nil {
t.Fatal(err)
}
if eng.restores != 1 || api.last().Status != StatusReverted || pending(t) {
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
}
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
t.Fatalf("reverted generation re-applied: err=%v applies=%d", err, eng.applies)
}
}
+4 -21
View File
@@ -55,28 +55,11 @@ func (c *Config) validateBlrules() error {
return fmt.Errorf("blrules[%d]: dest required", i) return fmt.Errorf("blrules[%d]: dest required", i)
} }
if r.Source != "all" && r.Source != "any" && r.Source != "none" && if err := c.validateZoneRef(r.Source); err != nil {
!hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { return fmt.Errorf("blrules[%d]: source %w", i, err)
for _, zs := range SplitZoneList(r.Source) {
if _, ok := c.Zones[zs.Zone]; !ok {
return fmt.Errorf("blrules[%d]: source zone %q not defined", i, zs.Zone)
}
if !validAddrList(zs.Addr) {
return fmt.Errorf("blrules[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr)
}
}
}
if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" &&
!hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") {
for _, zs := range SplitZoneList(r.Dest) {
if _, ok := c.Zones[zs.Zone]; !ok {
return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, zs.Zone)
}
if !validAddrList(zs.Addr) {
return fmt.Errorf("blrules[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr)
}
} }
if err := c.validateZoneRef(r.Dest); err != nil {
return fmt.Errorf("blrules[%d]: dest %w", i, err)
} }
} }
return nil return nil
+23
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"regexp"
"strings" "strings"
"gopkg.in/yaml.v3" "gopkg.in/yaml.v3"
@@ -55,10 +56,17 @@ type Settings struct {
AddressFamily AddressFamily `yaml:"address_family,omitempty"` AddressFamily AddressFamily `yaml:"address_family,omitempty"`
IPForwarding bool `yaml:"ip_forwarding"` IPForwarding bool `yaml:"ip_forwarding"`
LogLevel string `yaml:"log_level"` LogLevel string `yaml:"log_level"`
// LogLimit rate-limits every log site, shorewall LOGLIMIT syntax rate/unit[:burst]; unset logs every hit.
LogLimit string `yaml:"log_limit,omitempty"`
TableName string `yaml:"table_name"` TableName string `yaml:"table_name"`
// When true, auto-generate CONTINUE policies for sub-zones to their parent zones. // When true, auto-generate CONTINUE policies for sub-zones to their parent zones.
ImplicitContinue bool `yaml:"implicit_continue,omitempty"` ImplicitContinue bool `yaml:"implicit_continue,omitempty"`
// Verdict for ct state invalid/untracked packets; continue passes them to the rules.
// Unset: invalid drops, untracked continues.
InvalidDisposition PolicyAction `yaml:"invalid_disposition,omitempty"`
UntrackedDisposition PolicyAction `yaml:"untracked_disposition,omitempty"`
} }
// Load reads a config file in YAML or JSON format (detected by extension). // Load reads a config file in YAML or JSON format (detected by extension).
@@ -107,6 +115,8 @@ func (c *Config) applyDefaults() {
} }
} }
var logLimitRe = regexp.MustCompile(`^[1-9][0-9]*/(sec|second|min|minute|hour|day)(:[1-9][0-9]*)?$`)
var validAddressFamilies = map[AddressFamily]bool{ var validAddressFamilies = map[AddressFamily]bool{
FamilyINET: true, FamilyIP: true, FamilyIP6: true, FamilyINET: true, FamilyIP: true, FamilyIP6: true,
} }
@@ -115,6 +125,19 @@ func (c *Config) validateSettings() error {
if !validAddressFamilies[c.Settings.AddressFamily] { if !validAddressFamilies[c.Settings.AddressFamily] {
return fmt.Errorf("unknown address_family %q (use inet, ip, or ip6)", c.Settings.AddressFamily) return fmt.Errorf("unknown address_family %q (use inet, ip, or ip6)", c.Settings.AddressFamily)
} }
if l := c.Settings.LogLimit; l != "" && !logLimitRe.MatchString(l) {
return fmt.Errorf("invalid log_limit %q (use rate/{sec|min|hour|day}[:burst]; per-source s:/d: is not supported)", l)
}
for name, d := range map[string]PolicyAction{
"invalid_disposition": c.Settings.InvalidDisposition,
"untracked_disposition": c.Settings.UntrackedDisposition,
} {
switch d {
case "", PolicyAccept, PolicyDrop, PolicyReject, PolicyContinue:
default:
return fmt.Errorf("unknown %s %q (use accept, drop, reject, or continue)", name, d)
}
}
return nil return nil
} }
+29
View File
@@ -1052,3 +1052,32 @@ func TestSplitZoneList(t *testing.T) {
} }
} }
} }
func TestValidateDispositions(t *testing.T) {
for _, tc := range []struct {
invalid, untracked PolicyAction
wantErr string
}{
{"", "", ""},
{PolicyContinue, PolicyDrop, ""},
{"bogus", "", `unknown invalid_disposition "bogus"`},
{PolicyAccept, "log", `unknown untracked_disposition "log"`},
} {
c := baseConfig()
c.Settings.InvalidDisposition = tc.invalid
c.Settings.UntrackedDisposition = tc.untracked
checkErr(t, c.Validate(), tc.wantErr)
}
}
func TestValidateLogLimit(t *testing.T) {
for v, ok := range map[string]bool{
"": true, "1/sec": true, "1/sec:10": true, "30/minute:5": true, "2/hour": true, "1/day:1": true,
"s:1/sec:10": false, "d:1/sec": false, "1": false, "1/week": false, "0/sec": false, "1/sec:": false,
} {
c := &Config{Settings: Settings{AddressFamily: FamilyINET, LogLimit: v}}
if err := c.validateSettings(); (err == nil) != ok {
t.Errorf("log_limit %q: err = %v, want ok=%v", v, err, ok)
}
}
}
+18 -3
View File
@@ -1,6 +1,9 @@
package config package config
import "fmt" import (
"fmt"
"strings"
)
type ConntrackAction string type ConntrackAction string
@@ -68,8 +71,14 @@ func (c *Config) validateConntrack() error {
return fmt.Errorf("conntrack[%d]: helper name required for helper action", i) return fmt.Errorf("conntrack[%d]: helper name required for helper action", i)
} }
if ct.Source == "" && ct.Dest == "" && ct.Action != ConntrackHelper { if HasZoneExclusion(ct.Source) || HasZoneExclusion(ct.Dest) {
return fmt.Errorf("conntrack[%d]: source or dest required", i) return fmt.Errorf("conntrack[%d]: zone exclusions are not supported in conntrack entries", i)
}
if err := c.validateZoneRef(ct.Source); err != nil {
return fmt.Errorf("conntrack[%d]: source %w", i, err)
}
if err := c.validateZoneRef(ct.Dest); err != nil {
return fmt.Errorf("conntrack[%d]: dest %w", i, err)
} }
if ct.User != "" { if ct.User != "" {
@@ -84,3 +93,9 @@ func (c *Config) validateConntrack() error {
} }
return nil return nil
} }
// HasZoneExclusion reports an all/any zone ref with a "+" or "!" modifier (all+, all!x, any+!x, ...).
func HasZoneExclusion(spec string) bool {
zones, _, _ := strings.Cut(spec, ":")
return (strings.HasPrefix(zones, "all") || strings.HasPrefix(zones, "any")) && strings.ContainsAny(zones[3:], "+!")
}
+45 -4
View File
@@ -53,11 +53,52 @@ func TestValidateConntrack(t *testing.T) {
}, },
}, },
{ {
name: "source or dest required for non-helper", name: "omitted source and dest is valid",
rules: []ConntrackRule{ rules: []ConntrackRule{{Action: ConntrackNoTrack, Proto: "udp", DPort: PortSpec{"53"}}},
{Action: ConntrackDrop},
}, },
wantErr: "source or dest required", {
name: "unknown source zone",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "nte"}},
wantErr: `source zone "nte" not defined`,
},
{
name: "unknown dest zone",
rules: []ConntrackRule{{Action: ConntrackDrop, Source: "net", Dest: "nte:192.0.2.1"}},
wantErr: `dest zone "nte" not defined`,
},
{
name: "all and plain zone forms are valid",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "net,fw", Dest: "all:192.0.2.1"}},
},
{
name: "Source all!net rejected",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "all!net"}},
wantErr: "zone exclusions are not supported in conntrack entries",
},
{
name: "Dest all!net:192.0.2.1 rejected",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Dest: "all!net:192.0.2.1"}},
wantErr: "zone exclusions are not supported in conntrack entries",
},
{
name: "Source all+ rejected",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "all+"}},
wantErr: "zone exclusions are not supported in conntrack entries",
},
{
name: "Dest all+!net rejected",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Dest: "all+!net"}},
wantErr: "zone exclusions are not supported in conntrack entries",
},
{
name: "Source any!net rejected",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Source: "any!net"}},
wantErr: "zone exclusions are not supported in conntrack entries",
},
{
name: "Dest any+ rejected",
rules: []ConntrackRule{{Action: ConntrackNoTrack, Dest: "any+"}},
wantErr: "zone exclusions are not supported in conntrack entries",
}, },
{ {
name: "helper without source/dest is valid", name: "helper without source/dest is valid",
+32 -22
View File
@@ -173,29 +173,13 @@ func (c *Config) validateRules() error {
return fmt.Errorf("rule[%d]: dest required", i) return fmt.Errorf("rule[%d]: dest required", i)
} }
if r.Source != "all" && r.Source != "any" && r.Source != "none" && if err := c.validateZoneRef(r.Source); err != nil {
!hasPrefix(r.Source, "all+") && !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { return fmt.Errorf("rule[%d]: source %w", i, err)
for _, zs := range SplitZoneList(r.Source) {
if _, ok := c.Zones[zs.Zone]; !ok {
return fmt.Errorf("rule[%d]: source zone %q not defined", i, zs.Zone)
}
if !validAddrList(zs.Addr) {
return fmt.Errorf("rule[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr)
}
}
} }
if r.Action != RuleDNAT && r.Action != RuleRedirect && r.Action != RuleNoNAT { if r.Action != RuleDNAT && r.Action != RuleRedirect && r.Action != RuleNoNAT {
if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && if err := c.validateZoneRef(r.Dest); err != nil {
!hasPrefix(r.Dest, "all+") && !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { return fmt.Errorf("rule[%d]: dest %w", i, err)
for _, zs := range SplitZoneList(r.Dest) {
if _, ok := c.Zones[zs.Zone]; !ok {
return fmt.Errorf("rule[%d]: dest zone %q not defined", i, zs.Zone)
}
if !validAddrList(zs.Addr) {
return fmt.Errorf("rule[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr)
}
}
} }
} }
@@ -264,6 +248,32 @@ func zoneFromSpec(spec string) string {
return spec return spec
} }
func hasPrefix(s, prefix string) bool { // validateZoneRef checks a SOURCE/DEST spec: all/any[+][!excluded,...][:addr], none, or a declared zone list.
return len(s) >= len(prefix) && s[:len(prefix)] == prefix func (c *Config) validateZoneRef(spec string) error {
zones, addr, _ := strings.Cut(spec, ":")
base, excl, isExcl := strings.Cut(zones, "!")
switch base {
case "", "none", "all", "all+", "any", "any+":
if base == "" && isExcl {
return fmt.Errorf("%q: exclusion needs all or any", spec)
}
for _, z := range strings.Split(excl, ",") {
if _, ok := c.Zones[strings.TrimSpace(z)]; isExcl && !ok {
return fmt.Errorf("excluded zone %q not defined", z)
}
}
if !validAddrList(addr) {
return fmt.Errorf("%q: '!' may only prefix the whole address list", addr)
}
return nil
}
for _, zs := range SplitZoneList(spec) {
if _, ok := c.Zones[zs.Zone]; !ok {
return fmt.Errorf("zone %q not defined", zs.Zone)
}
if !validAddrList(zs.Addr) {
return fmt.Errorf("%q: '!' may only prefix the whole address list", zs.Addr)
}
}
return nil
} }
+2 -2
View File
@@ -28,7 +28,7 @@ func (e *Engine) FindForeignRules() ([]ForeignRule, error) {
var ourTable *nftables.Table var ourTable *nftables.Table
for _, t := range tables { for _, t := range tables {
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { if t.Name == e.cfg.Settings.TableName && t.Family == e.family() {
ourTable = t ourTable = t
break break
} }
@@ -52,7 +52,7 @@ func (e *Engine) FindForeignRules() ([]ForeignRule, error) {
var foreign []ForeignRule var foreign []ForeignRule
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) chains, err := e.conn.ListChainsOfTableFamily(e.family())
if err != nil { if err != nil {
return nil, fmt.Errorf("listing chains: %w", err) return nil, fmt.Errorf("listing chains: %w", err)
} }
+449 -128
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"net" "net"
"slices"
"sort" "sort"
"strconv" "strconv"
"strings" "strings"
@@ -64,26 +65,68 @@ func (c *Compiler) Compile() (*FirewallState, error) {
return nil, fmt.Errorf("static-nat: %w", err) return nil, fmt.Errorf("static-nat: %w", err)
} }
c.compileMSSClamp(state) c.compileMSSClamp(state)
limitLogs(state, c.cfg.Settings.LogLimit)
return state, nil return state, nil
} }
// limitLogs puts a limit in front of every log expression. A limit stops the
// whole rule, so like shorewall's separate LOG rule, a log followed by an action
// splits into a limited log-only rule and the same rule without the log. A
// LOG-action rule (nothing but rate limits after the log) keeps only the log rule.
func limitLogs(state *FirewallState, spec string) {
if spec == "" {
return
}
for chain, rules := range state.Rules {
var out []ManagedRule
for _, r := range rules {
i := slices.IndexFunc(r.Exprs, func(e expr.Any) bool { _, ok := e.(*expr.Log); return ok })
if i < 0 {
out = append(out, r)
continue
}
logRule := r
logRule.Exprs = slices.Concat(r.Exprs[:i], parseRateLimit(spec), r.Exprs[i:i+1])
out = append(out, logRule)
if slices.ContainsFunc(r.Exprs[i+1:], func(e expr.Any) bool { _, ok := e.(*expr.Limit); return !ok }) {
r.Exprs = slices.Concat(r.Exprs[:i], r.Exprs[i+1:])
out = append(out, r)
}
}
state.Rules[chain] = out
}
}
func (c *Compiler) compileConntrackFastPath(state *FirewallState) error { func (c *Compiler) compileConntrackFastPath(state *FirewallState) error {
invalid := c.cfg.Settings.InvalidDisposition
if invalid == "" {
invalid = config.PolicyDrop
}
for _, chain := range []string{"input", "forward", "output"} { for _, chain := range []string{"input", "forward", "output"} {
state.Rules[chain] = append(state.Rules[chain], state.Rules[chain] = append(state.Rules[chain], ManagedRule{
ManagedRule{
Chain: chain, Chain: chain,
Exprs: append(matchCtState(ctStateEstablished|ctStateRelated), Exprs: append(matchCtState(ctStateEstablished|ctStateRelated),
&expr.Verdict{Kind: expr.VerdictAccept}), &expr.Verdict{Kind: expr.VerdictAccept}),
Tag: "ct:fastpath:" + chain, Tag: "ct:fastpath:" + chain,
}, })
ManagedRule{ for _, d := range []struct {
name string
state uint32
action config.PolicyAction
}{
{"invalid", ctStateInvalid, invalid},
{"untracked", ctStateUntracked, c.cfg.Settings.UntrackedDisposition},
} {
if d.action == "" || d.action == config.PolicyContinue {
continue
}
state.Rules[chain] = append(state.Rules[chain], ManagedRule{
Chain: chain, Chain: chain,
Exprs: append(matchCtState(ctStateInvalid), Exprs: append(matchCtState(d.state), policyVerdict(d.action, c.cfg.Settings.AddressFamily)...),
&expr.Verdict{Kind: expr.VerdictDrop}), Tag: "ct:" + d.name + ":" + chain,
Tag: "ct:invalid:" + chain, })
}, }
)
} }
return nil return nil
} }
@@ -125,41 +168,24 @@ func (c *Compiler) compileDHCP(state *FirewallState) {
continue continue
} }
name := iface.PhysicalName() name := iface.PhysicalName()
// Allow DHCPv4 client traffic (bootpc:68 → bootps:67) dhcp := func(chain, dir string, ifaceMatch []expr.Any) {
state.Rules["input"] = append(state.Rules["input"], ManagedRule{ state.Rules[chain] = append(state.Rules[chain], ManagedRule{
Chain: "input", Chain: chain,
Exprs: append(append(append( Exprs: append(append(append(append(ifaceMatch,
matchIfaceName(true, name), matchNFProto(unix.NFPROTO_IPV4)...),
matchProtoNum(unix.IPPROTO_UDP)...), matchProtoNum(unix.IPPROTO_UDP)...),
matchSPort(68)...), matchDPortRange(67, 68)...),
matchDPort(67)..., &expr.Verdict{Kind: expr.VerdictAccept}),
), Tag: fmt.Sprintf("dhcp:%s:%s", dir, iface.Interface),
Tag: fmt.Sprintf("dhcp:in:%s", iface.Interface),
})
// Allow DHCPv4 server → client replies
state.Rules["input"] = append(state.Rules["input"], ManagedRule{
Chain: "input",
Exprs: append(append(append(append(
matchIfaceName(true, name),
matchProtoNum(unix.IPPROTO_UDP)...),
matchSPort(67)...),
matchDPort(68)...),
&expr.Verdict{Kind: expr.VerdictAccept},
),
Tag: fmt.Sprintf("dhcp:reply:%s", iface.Interface),
})
state.Rules["output"] = append(state.Rules["output"], ManagedRule{
Chain: "output",
Exprs: append(append(append(append(
matchIfaceName(false, name),
matchProtoNum(unix.IPPROTO_UDP)...),
matchSPort(68)...),
matchDPort(67)...),
&expr.Verdict{Kind: expr.VerdictAccept},
),
Tag: fmt.Sprintf("dhcp:out:%s", iface.Interface),
}) })
} }
// shorewall: udp dport 67:68 both ways between fw and iface, forwarded back out a bridge
dhcp("input", "in", matchIfaceName(true, name))
dhcp("output", "out", matchIfaceName(false, name))
if iface.Options.Bridge {
dhcp("forward", "fwd", append(matchIfaceName(true, name), matchIfaceName(false, name)...))
}
}
} }
func (c *Compiler) compileIntraZone(state *FirewallState) error { func (c *Compiler) compileIntraZone(state *FirewallState) error {
@@ -205,47 +231,151 @@ func (c *Compiler) compileBlrules(state *FirewallState) error {
} }
if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, if err := c.compileOneRule(state, tag, rule.Source, rule.Dest,
rule.Proto, rule.DPort, rule.SPort, rule.Proto, rule.DPort, rule.SPort,
action, rule.Log, "", fwZone, ""); err != nil { action, rule.Log, "", "", fwZone, ""); err != nil {
return fmt.Errorf("blrule[%d]: %w", i, err) return fmt.Errorf("blrule[%d]: %w", i, err)
} }
} }
return nil return nil
} }
// helperProtos is each kernel helper's default transport, used when an entry gives no proto.
var helperProtos = map[string]string{
"amanda": "udp", "ftp": "tcp", "irc": "tcp", "netbios-ns": "udp", "pptp": "tcp",
"Q.931": "tcp", "RAS": "udp", "sane": "tcp", "sip": "udp", "snmp": "udp", "tftp": "udp",
}
func helperObjName(helper, proto string) string {
if proto != helperProtos[helper] {
return helper + "-" + proto
}
return helper
}
// expandHelper declares a ct helper object per (helper, proto) and returns one entry per proto.
func expandHelper(state *FirewallState, ct config.ConntrackRule) ([]config.ConntrackRule, error) {
proto := ct.Proto
if proto == "" {
if proto = helperProtos[ct.Helper]; proto == "" {
return nil, fmt.Errorf("proto required for helper %q", ct.Helper)
}
}
var out []config.ConntrackRule
for _, p := range strings.Split(proto, ",") {
p = strings.TrimSpace(p)
n, err := protoNumber(p)
if err != nil {
return nil, err
}
name := helperObjName(ct.Helper, p)
if !slices.ContainsFunc(state.Helpers, func(h Helper) bool { return h.Name == name }) {
state.Helpers = append(state.Helpers, Helper{Name: name,
Helper: expr.CtHelper{Name: ct.Helper, L3Proto: unix.NFPROTO_INET, L4Proto: n}})
}
pct := ct
pct.Proto = p
out = append(out, pct)
}
return out, nil
}
func (c *Compiler) compileConntrack(state *FirewallState) error { func (c *Compiler) compileConntrack(state *FirewallState) error {
fwZone := c.cfg.FirewallZone()
for i, ct := range c.cfg.Conntrack { for i, ct := range c.cfg.Conntrack {
tag := fmt.Sprintf("conntrack:%d", i) tag := fmt.Sprintf("conntrack:%d", i)
if config.HasZoneExclusion(ct.Source) || config.HasZoneExclusion(ct.Dest) {
chains := []string{"prerouting"} return fmt.Errorf("conntrack[%d]: zone exclusions are not supported in conntrack entries", i)
switch ct.Chain {
case config.ConntrackOutput:
chains = []string{"output"}
case config.ConntrackBoth:
chains = []string{"prerouting", "output"}
} }
cts := []config.ConntrackRule{ct}
matches, err := l4Matches(ct.Proto, ct.DPort, nil) if ct.Action == config.ConntrackHelper {
if err != nil { if ct.Chain == "" {
ct.Chain = config.ConntrackBoth
}
var err error
if cts, err = expandHelper(state, ct); err != nil {
return fmt.Errorf("conntrack[%d]: %w", i, err) return fmt.Errorf("conntrack[%d]: %w", i, err)
} }
}
srcs, dsts := c.zoneSpecs(ct.Source), c.zoneSpecs(ct.Dest)
if len(srcs) == 0 {
srcs = []config.ZoneSpec{{}}
}
if len(dsts) == 0 {
dsts = []config.ZoneSpec{{}}
}
for _, src := range srcs {
if src.Zone == "all" || src.Zone == "any" {
src.Zone = ""
}
chains := []string{"raw_prerouting"}
switch {
case ct.Chain == config.ConntrackOutput && src.Zone != fwZone && src.Zone != "":
return fmt.Errorf("conntrack[%d]: chain output needs SOURCE %s, got %q", i, fwZone, src.Zone)
case ct.Chain == config.ConntrackPrerouting && src.Zone == fwZone:
return fmt.Errorf("conntrack[%d]: SOURCE %s cannot use chain prerouting", i, fwZone)
case ct.Chain != config.ConntrackPrerouting && src.Zone == fwZone, ct.Chain == config.ConntrackOutput:
chains = []string{"raw_output"}
case ct.Chain == config.ConntrackBoth && src.Zone == "":
chains = []string{"raw_prerouting", "raw_output"}
}
for _, srcAddr := range splitAddrs(src.Addr) {
for _, dst := range dsts {
for _, dstAddr := range splitAddrs(dst.Addr) {
for _, chain := range chains { for _, chain := range chains {
for _, m := range matches { for _, pct := range cts {
exprs := append([]expr.Any{}, m.exprs...) if err := c.compileConntrackPair(state, tag, chain, pct, src.Zone, srcAddr, dst.Zone, dstAddr); err != nil {
return fmt.Errorf("conntrack[%d]: %w", i, err)
}
}
}
}
}
}
}
}
return nil
}
// compileConntrackPair matches iif of the source zone in raw_prerouting and oif of the dest zone in raw_output.
// Helpers are assigned after conntrack (-200) has created the entry, so they go to the mangle-priority
// helper_* chain instead; a raw-priority assignment is a no-op.
func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule,
srcZone, srcAddr, dstZone, dstAddr string) error {
if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" &&
(dstAddr == "" || strings.HasPrefix(dstAddr, "!")) {
return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone)
}
srcIfaces, dstIfaces := c.resolveZoneInterfaces(srcZone, srcAddr), []string{""}
if chain == "raw_prerouting" && c.resolveZoneInterfaces(dstZone, dstAddr) == nil {
return nil
}
if chain == "raw_output" {
srcIfaces, dstIfaces = []string{""}, c.resolveZoneInterfaces(dstZone, dstAddr)
}
out := chain
if ct.Action == config.ConntrackHelper {
out = "helper_" + strings.TrimPrefix(chain, "raw_")
}
for _, srcIface := range srcIfaces {
for _, dstIface := range dstIfaces {
matches, err := c.buildMatchExprs(srcIface, dstIface, chain, ct.Proto, ct.DPort, ct.SPort, srcAddr, dstAddr)
if err != nil {
return err
}
for _, m := range matches {
exprs := m.exprs
switch ct.Action { switch ct.Action {
case config.ConntrackNoTrack: case config.ConntrackNoTrack:
exprs = append(exprs, &expr.Notrack{}) exprs = append(exprs, &expr.Notrack{})
case config.ConntrackHelper:
continue
case config.ConntrackDrop: case config.ConntrackDrop:
exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop})
case config.ConntrackHelper:
exprs = append(exprs, &expr.Objref{Type: unix.NFT_OBJECT_CT_HELPER, Name: helperObjName(ct.Helper, ct.Proto)})
} }
state.Rules[out] = append(state.Rules[out], ManagedRule{
state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: out,
Chain: chain,
Exprs: exprs, Exprs: exprs,
Tag: tag + ":" + chain, Tag: tag + ":" + out,
}) })
} }
} }
@@ -276,14 +406,14 @@ func (c *Compiler) compileRules(state *FirewallState) error {
if err != nil { if err != nil {
return fmt.Errorf("rule[%d]: %w", i, err) return fmt.Errorf("rule[%d]: %w", i, err)
} }
if len(matches)*specCount(rule.Source, rule.Dest, rule.Action) > 1 { if len(matches)*c.specCount(rule.Source, rule.Dest, rule.OrigDest, fwZone, rule.Action) > 1 {
return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists (each expanded rule would get its own limiter)", i) return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists (each expanded rule would get its own limiter)", i)
} }
} }
if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, if err := c.compileOneRule(state, tag, rule.Source, rule.Dest,
proto, dports, sport, proto, dports, sport,
rule.Action, rule.Log, rule.Dest, fwZone, rule.Section); err != nil { rule.Action, rule.Log, rule.Dest, rule.OrigDest, fwZone, rule.Section); err != nil {
return fmt.Errorf("rule[%d]: %w", i, err) return fmt.Errorf("rule[%d]: %w", i, err)
} }
@@ -313,24 +443,24 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
break break
} }
var extra []expr.Any var match, extra []expr.Any
var replaceVerdict []expr.Any var replaceVerdict []expr.Any
if rule.User != "" { if rule.User != "" {
extra = append(extra, matchUID(rule.User)...) match = append(match, matchUID(rule.User)...)
} }
if rule.Mark != "" { if rule.Mark != "" {
extra = append(extra, matchMark(rule.Mark)...) match = append(match, matchMark(rule.Mark)...)
}
if rule.ConnLimit != "" {
match = append(match, matchConnLimit(rule.ConnLimit)...)
}
if rule.Time != nil {
match = append(match, matchTime(rule.Time)...)
} }
if rule.RateLimit != "" { if rule.RateLimit != "" {
extra = append(extra, parseRateLimit(rule.RateLimit)...) extra = append(extra, parseRateLimit(rule.RateLimit)...)
} }
if rule.ConnLimit != "" {
extra = append(extra, matchConnLimit(rule.ConnLimit)...)
}
if rule.Time != nil {
extra = append(extra, matchTime(rule.Time)...)
}
if rule.SetMark != "" { if rule.SetMark != "" {
extra = append(extra, setMarkExprs(rule.SetMark)...) extra = append(extra, setMarkExprs(rule.SetMark)...)
} }
@@ -338,21 +468,27 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}} replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}}
} }
if len(extra) > 0 || len(replaceVerdict) > 0 { if len(match) > 0 || len(extra) > 0 || len(replaceVerdict) > 0 {
existingExprs := rules[idx].Exprs var pre, post, verdict []expr.Any
var verdict []expr.Any for _, e := range rules[idx].Exprs {
var nonVerdict []expr.Any switch e.(type) {
for _, e := range existingExprs { case *expr.Verdict:
if _, ok := e.(*expr.Verdict); ok {
verdict = append(verdict, e) verdict = append(verdict, e)
case *expr.Log:
post = append(post, e)
default:
if len(post) > 0 {
post = append(post, e)
} else { } else {
nonVerdict = append(nonVerdict, e) pre = append(pre, e)
}
} }
} }
if len(replaceVerdict) > 0 { if len(replaceVerdict) > 0 {
verdict = replaceVerdict verdict = replaceVerdict
} }
rules[idx].Exprs = append(append(nonVerdict, extra...), verdict...) // matches go before the log so it only fires for packets the rule matches
rules[idx].Exprs = slices.Concat(pre, match, post, extra, verdict)
} }
} }
state.Rules[chain] = rules state.Rules[chain] = rules
@@ -360,19 +496,32 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, proto string, func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, proto string,
dports, sports config.PortSpec, action config.RuleAction, logLevel string, dports, sports config.PortSpec, action config.RuleAction, logLevel string,
dnatDest string, fwZone string, section config.RuleSection) error { dnatDest, origDest string, fwZone string, section config.RuleSection) error {
for _, src := range zoneSpecs(srcSpec) {
for _, srcAddr := range splitAddrs(src.Addr) {
if action == config.RuleDNAT || action == config.RuleRedirect { if action == config.RuleDNAT || action == config.RuleRedirect {
if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, dstSpec, proto, dports, action, logLevel); err != nil { for _, src := range c.dnatSourceSpecs(srcSpec, fwZone) {
return err if dnatSkipsIntrazone(srcSpec, src.Zone, dstSpec) {
}
continue continue
} }
for _, dst := range zoneSpecs(dstSpec) { for _, srcAddr := range splitAddrs(src.Addr) {
if err := c.compileDNATAccept(state, tag+":accept", src.Zone, srcAddr, dstSpec, proto, dports, sports, action, fwZone, section); err != nil {
return err
}
for _, od := range splitAddrs(origDest) {
if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, od, dstSpec, proto, dports, sports, action, logLevel); err != nil {
return err
}
}
}
}
return nil
}
for _, p := range c.zonePairs(srcSpec, dstSpec, fwZone) {
src, dst := p[0], p[1]
for _, srcAddr := range splitAddrs(src.Addr) {
for _, od := range splitAddrs(origDest) {
for _, dstAddr := range splitAddrs(dst.Addr) { for _, dstAddr := range splitAddrs(dst.Addr) {
if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, proto, if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, od, proto,
dports, sports, action, logLevel, fwZone, section); err != nil { dports, sports, action, logLevel, fwZone, section); err != nil {
return err return err
} }
@@ -383,27 +532,135 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p
return nil return nil
} }
// specCount is how many zone/address combinations compileOneRule expands src and dst into. // dnatSourceSpecs hooks DNAT per source zone like shorewall: all/any expand to every zone but fw (prerouting never sees fw traffic).
func specCount(srcSpec, dstSpec string, action config.RuleAction) int { func (c *Compiler) dnatSourceSpecs(spec, fwZone string) []config.ZoneSpec {
count := func(spec string) (n int) { zone, addr := splitZoneSpec(spec)
for _, z := range zoneSpecs(spec) { base, _, _ := strings.Cut(zone, "!")
n += len(splitAddrs(z.Addr)) if base = strings.TrimSuffix(base, "+"); base != "all" && base != "any" {
return c.zoneSpecs(spec)
} }
return n var out []config.ZoneSpec
for _, z := range c.expandZoneRef(zone) {
if z != fwZone {
out = append(out, config.ZoneSpec{Zone: z, Addr: addr})
} }
if action == config.RuleDNAT || action == config.RuleRedirect {
return count(srcSpec)
} }
return count(srcSpec) * count(dstSpec) return out
} }
// zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. // dnatSkipsIntrazone mirrors shorewall: a zone list or all/any source never pairs a zone with itself unless marked "+".
func zoneSpecs(spec string) []config.ZoneSpec { func dnatSkipsIntrazone(srcSpec, srcZone, dstSpec string) bool {
zones, _, _ := strings.Cut(srcSpec, ":")
base, _, _ := strings.Cut(zones, "!")
wild := base == "all" || base == "any" || strings.Contains(base, ",")
dstZone, _, _ := strings.Cut(dstSpec, ":")
return wild && srcZone == dstZone
}
// compileDNATAccept emits the filter ACCEPT implied by DNAT/REDIRECT (shorewall's DNAT-/REDIRECT- omit it) for the translated flow.
func (c *Compiler) compileDNATAccept(state *FirewallState, tag, srcZone, srcAddr, dstSpec, proto string,
dports, sports config.PortSpec, action config.RuleAction, fwZone string, section config.RuleSection) error {
parts := strings.SplitN(dstSpec, ":", 3)
if len(parts) < 2 {
return fmt.Errorf("DNAT dest must be zone:address or zone:address:port")
}
dstZone, dstAddr := parts[0], parts[1]
if action == config.RuleRedirect {
dstZone, dstAddr = fwZone, ""
}
if len(parts) == 3 {
dports = config.PortSpec{parts[2]}
}
chain := c.selectChain(srcZone, dstZone, fwZone)
n := len(state.Rules[chain])
if err := c.compileZonePair(state, tag, srcZone, srcAddr, dstZone, dstAddr, "", proto,
dports, sports, config.RuleAccept, "", fwZone, section); err != nil {
return err
}
for i := n; i < len(state.Rules[chain]); i++ {
e := state.Rules[chain][i].Exprs
last := len(e) - 1
state.Rules[chain][i].Exprs = append(append(e[:last:last], matchCtBits(expr.CtKeySTATUS, ctStatusDNAT)...), e[last])
}
return nil
}
// specCount is how many zone/address combinations compileOneRule expands src and dst into.
func (c *Compiler) specCount(srcSpec, dstSpec, origDest, fwZone string, action config.RuleAction) int {
n := 0
if action == config.RuleDNAT || action == config.RuleRedirect {
for _, src := range c.dnatSourceSpecs(srcSpec, fwZone) {
n += len(splitAddrs(src.Addr))
}
return n * len(splitAddrs(origDest))
}
for _, p := range c.zonePairs(srcSpec, dstSpec, fwZone) {
n += len(splitAddrs(p[0].Addr)) * len(splitAddrs(p[1].Addr))
}
return n * len(splitAddrs(origDest))
}
// zonePairs is the src/dst zone expansion of a non-DNAT rule, with fw added beside all/any.
func (c *Compiler) zonePairs(srcSpec, dstSpec, fwZone string) [][2]config.ZoneSpec {
srcs, srcGlobal := withFirewall(c.zoneSpecs(srcSpec), fwZone)
dsts, dstGlobal := withFirewall(c.zoneSpecs(dstSpec), fwZone)
var out [][2]config.ZoneSpec
for _, src := range srcs {
for _, dst := range dsts {
if src.Zone == fwZone && dst.Zone == fwZone && (srcGlobal || dstGlobal) {
continue
}
// Exclusion expansion never pairs fw with itself, and pairs a zone with itself only for "all+".
if src.Zone == dst.Zone && (isZoneExclusion(srcSpec) || isZoneExclusion(dstSpec)) &&
(src.Zone == fwZone || !strings.Contains(srcSpec, "+!") && !strings.Contains(dstSpec, "+!")) {
continue
}
out = append(out, [2]config.ZoneSpec{src, dst})
}
}
return out
}
// zoneSpecs expands a comma zone list; "all"/"any" stay global and "all!x,y" becomes every zone but x and y.
func (c *Compiler) zoneSpecs(spec string) []config.ZoneSpec {
zone, addr := splitZoneSpec(spec) zone, addr := splitZoneSpec(spec)
if base, _, _ := strings.Cut(strings.TrimSuffix(zone, "+"), "!"); base == "all" || base == "any" { if !isZoneExclusion(zone) {
if base := strings.TrimSuffix(zone, "+"); base == "all" || base == "any" {
return []config.ZoneSpec{{Zone: zone, Addr: addr}} return []config.ZoneSpec{{Zone: zone, Addr: addr}}
} }
return config.SplitZoneList(spec) return config.SplitZoneList(spec)
}
var out []config.ZoneSpec
for _, z := range c.expandZoneRef(zone) {
out = append(out, config.ZoneSpec{Zone: z, Addr: addr})
}
return out
}
// withFirewall adds the firewall zone beside a global all/any spec, which otherwise only reaches forward.
func withFirewall(specs []config.ZoneSpec, fwZone string) ([]config.ZoneSpec, bool) {
out, global := specs, false
for _, s := range specs {
if fwZone != "" && isGlobalZone(s.Zone) {
global = true
if fw := (config.ZoneSpec{Zone: fwZone, Addr: s.Addr}); !slices.Contains(out, fw) {
out = append(out, fw)
}
}
}
return out, global
}
func isGlobalZone(spec string) bool {
zone, _ := splitZoneSpec(spec)
base := strings.TrimSuffix(zone, "+")
return base == "all" || base == "any"
}
func isZoneExclusion(spec string) bool {
base, _, ok := strings.Cut(spec, "!")
base = strings.TrimSuffix(base, "+")
return ok && (base == "all" || base == "any")
} }
// splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. // splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match.
@@ -414,12 +671,16 @@ func splitAddrs(addr string) []string {
return strings.Split(addr, ",") return strings.Split(addr, ",")
} }
func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, proto string, func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, origDest, proto string,
dports, sports config.PortSpec, action config.RuleAction, logLevel string, dports, sports config.PortSpec, action config.RuleAction, logLevel string,
fwZone string, section config.RuleSection) error { fwZone string, section config.RuleSection) error {
srcIfaces := c.resolveZoneInterfaces(srcZone) srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr)
dstIfaces := c.resolveZoneInterfaces(dstZone) dstIfaces := c.resolveZoneInterfaces(dstZone, dstAddr)
chain := c.selectChain(srcZone, dstZone, fwZone) chain := c.selectChain(srcZone, dstZone, fwZone)
// ponytail: forward daddr is post-DNAT; lift with `ct original daddr` (expr.Ct Direction, google/nftables v0.3.0).
if origDest != "" && chain == "forward" {
return fmt.Errorf("origdest: ORIGDEST on forwarded rules is not supported yet")
}
for _, srcIface := range srcIfaces { for _, srcIface := range srcIfaces {
for _, dstIface := range dstIfaces { for _, dstIface := range dstIfaces {
@@ -427,6 +688,15 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr,
if err != nil { if err != nil {
return err return err
} }
if origDest != "" {
od, err := matchOrigDest(origDest)
if err != nil {
return fmt.Errorf("origdest: %w", err)
}
for i := range matches {
matches[i].exprs = append(matches[i].exprs, od...)
}
}
for _, m := range matches { for _, m := range matches {
exprs := m.exprs exprs := m.exprs
@@ -455,8 +725,8 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr,
return nil return nil
} }
func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, dstSpec, proto string, func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, origDest, dstSpec, proto string,
dports config.PortSpec, action config.RuleAction, logLevel string) error { dports, sports config.PortSpec, action config.RuleAction, logLevel string) error {
chain := "prerouting" chain := "prerouting"
parts := strings.SplitN(dstSpec, ":", 3) parts := strings.SplitN(dstSpec, ":", 3)
@@ -474,9 +744,17 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
dnatPort = uint16(p) dnatPort = uint16(p)
} }
srcIfaces := c.resolveZoneInterfaces(srcZone) srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr)
matches, err := l4Matches(proto, dports, nil) var odExprs []expr.Any
if origDest != "" {
var err error
if odExprs, err = matchOrigDest(origDest); err != nil {
return fmt.Errorf("origdest: %w", err)
}
}
matches, err := l4Matches(proto, dports, sports)
if err != nil { if err != nil {
return err return err
} }
@@ -501,6 +779,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
} }
exprs = append(exprs, src...) exprs = append(exprs, src...)
} }
exprs = append(exprs, odExprs...)
exprs = append(exprs, m.exprs...) exprs = append(exprs, m.exprs...)
@@ -538,6 +817,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
) )
natExpr.RegProtoMin = 2 natExpr.RegProtoMin = 2
natExpr.RegProtoMax = 2 natExpr.RegProtoMax = 2
natExpr.Specified = true
} }
exprs = append(exprs, natExpr) exprs = append(exprs, natExpr)
} else { } else {
@@ -558,6 +838,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
) )
natExpr.RegProtoMin = 2 natExpr.RegProtoMin = 2
natExpr.RegProtoMax = 2 natExpr.RegProtoMax = 2
natExpr.Specified = true
} }
exprs = append(exprs, natExpr) exprs = append(exprs, natExpr)
} }
@@ -589,8 +870,8 @@ func (c *Compiler) compilePolicies(state *FirewallState) error {
} }
chain := c.selectChain(sz, dz, fwZone) chain := c.selectChain(sz, dz, fwZone)
srcIfaces := c.resolveZoneInterfaces(sz) srcIfaces := c.resolveZoneInterfaces(sz, "")
dstIfaces := c.resolveZoneInterfaces(dz) dstIfaces := c.resolveZoneInterfaces(dz, "")
for _, si := range srcIfaces { for _, si := range srcIfaces {
for _, di := range dstIfaces { for _, di := range dstIfaces {
@@ -599,7 +880,7 @@ func (c *Compiler) compilePolicies(state *FirewallState) error {
if si != "" { if si != "" {
exprs = append(exprs, matchIfaceName(true, si)...) exprs = append(exprs, matchIfaceName(true, si)...)
} }
if di != "" && chain == "forward" { if di != "" && chain != "input" {
exprs = append(exprs, matchIfaceName(false, di)...) exprs = append(exprs, matchIfaceName(false, di)...)
} }
@@ -932,15 +1213,26 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string {
return "forward" return "forward"
} }
func (c *Compiler) resolveZoneInterfaces(zone string) []string { // resolveZoneInterfaces returns nil (fail closed) for an unknown zone, or one with no interfaces unless a non-negated address match narrows the rule.
if zone == "all" || zone == "" { func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string {
switch zone {
case "", "all", "all+", "any", "any+":
return []string{""} return []string{""}
} }
ifaces := c.cfg.ZoneInterfaces(zone) z, ok := c.cfg.Zones[zone]
if len(ifaces) > 0 { if !ok {
slog.Warn("compiler: unknown zone, skipping its rules", "zone", zone)
return nil
}
if z.Type == config.ZoneFirewall {
return []string{""}
}
if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 {
return ifaces return ifaces
} }
if z, ok := c.cfg.Zones[zone]; ok && z.Type == config.ZoneIP && !c.zoneHasHosts(zone) { if addr != "" && !strings.HasPrefix(addr, "!") {
return []string{""}
}
if !c.warned[zone] { if !c.warned[zone] {
if c.warned == nil { if c.warned == nil {
c.warned = map[string]bool{} c.warned = map[string]bool{}
@@ -949,17 +1241,6 @@ func (c *Compiler) resolveZoneInterfaces(zone string) []string {
slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone)
} }
return nil return nil
}
return []string{""}
}
func (c *Compiler) zoneHasHosts(zone string) bool {
for _, h := range c.cfg.Hosts {
if h.Zone == zone {
return true
}
}
return false
} }
func (c *Compiler) expandZoneRef(ref string) []string { func (c *Compiler) expandZoneRef(ref string) []string {
@@ -977,7 +1258,7 @@ func (c *Compiler) expandZoneRef(ref string) []string {
} }
} }
if base == "all" || base == "all+" { if base == "all" || base == "all+" || base == "any" || base == "any+" {
var zones []string var zones []string
for name := range c.cfg.Zones { for name := range c.cfg.Zones {
if excluded != nil && excluded[name] { if excluded != nil && excluded[name] {
@@ -997,7 +1278,7 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor
if srcIface != "" { if srcIface != "" {
exprs = append(exprs, matchIfaceName(true, srcIface)...) exprs = append(exprs, matchIfaceName(true, srcIface)...)
} }
if dstIface != "" && chain == "forward" { if dstIface != "" && chain != "input" {
exprs = append(exprs, matchIfaceName(false, dstIface)...) exprs = append(exprs, matchIfaceName(false, dstIface)...)
} }
@@ -1386,6 +1667,34 @@ func matchDestCIDR(cidr string) ([]expr.Any, error) {
return matchAddrCIDR(cidr, false) return matchAddrCIDR(cidr, false)
} }
// matchOrigDest guards the daddr match with the address's nfproto so it is family-correct in the inet table.
func matchOrigDest(addr string) ([]expr.Any, error) {
var proto byte
for i, a := range strings.Split(strings.TrimPrefix(addr, "!"), ",") {
a, _, _ = strings.Cut(a, "/")
p := byte(unix.NFPROTO_IPV6)
if ip := net.ParseIP(a); ip != nil && ip.To4() != nil {
p = unix.NFPROTO_IPV4
}
if i > 0 && p != proto {
return nil, fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr)
}
proto = p
}
dst, err := matchDestCIDR(addr)
if err != nil {
return nil, err
}
return append(matchNFProto(proto), dst...), nil
}
func matchNFProto(proto byte) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}},
}
}
func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) { func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) {
negated := false negated := false
if strings.HasPrefix(cidr, "!") { if strings.HasPrefix(cidr, "!") {
@@ -1471,13 +1780,18 @@ const (
ctStateRelated = 4 ctStateRelated = 4
ctStateNew = 8 ctStateNew = 8
ctStateUntracked = 64 ctStateUntracked = 64
ctStatusDNAT = 32
) )
func matchCtState(stateMask uint32) []expr.Any { func matchCtState(stateMask uint32) []expr.Any {
return matchCtBits(expr.CtKeySTATE, stateMask)
}
func matchCtBits(key expr.CtKey, stateMask uint32) []expr.Any {
stateBytes := make([]byte, 4) stateBytes := make([]byte, 4)
binary.NativeEndian.PutUint32(stateBytes, stateMask) binary.NativeEndian.PutUint32(stateBytes, stateMask)
return []expr.Any{ return []expr.Any{
&expr.Ct{Key: expr.CtKeySTATE, Register: 1}, &expr.Ct{Key: key, Register: 1},
&expr.Bitwise{ &expr.Bitwise{
SourceRegister: 1, SourceRegister: 1,
DestRegister: 1, DestRegister: 1,
@@ -1783,6 +2097,13 @@ func rejectExprs(proto byte, family config.AddressFamily) []expr.Any {
Code: 0, Code: 0,
}} }}
} }
// icmpx is inet-only; ip and ip6 tables silently drop on it.
switch family {
case config.FamilyIP:
return []expr.Any{&expr.Reject{Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 3}} // port-unreachable
case config.FamilyIP6:
return []expr.Any{&expr.Reject{Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 4}} // port-unreachable
}
return []expr.Any{&expr.Reject{ return []expr.Any{&expr.Reject{
Type: unix.NFT_REJECT_ICMPX_UNREACH, Type: unix.NFT_REJECT_ICMPX_UNREACH,
Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH, Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH,
File diff suppressed because it is too large Load Diff
+48 -3
View File
@@ -20,15 +20,25 @@ type ManagedRule struct {
type FirewallState struct { type FirewallState struct {
Rules map[string][]ManagedRule Rules map[string][]ManagedRule
// Helpers are the ct helper objects in kernel (insertion) order.
Helpers []Helper
}
// Helper is a named ct helper object.
type Helper struct {
Name string `json:"name"`
Helper expr.CtHelper `json:"helper"`
} }
type ChangeSet struct { type ChangeSet struct {
Add []ManagedRule Add []ManagedRule
Remove []ManagedRule Remove []ManagedRule
AddHelpers []Helper
RemoveHelpers []string
} }
func (cs *ChangeSet) Empty() bool { func (cs *ChangeSet) Empty() bool {
return len(cs.Add) == 0 && len(cs.Remove) == 0 return len(cs.Add) == 0 && len(cs.Remove) == 0 && len(cs.AddHelpers) == 0 && len(cs.RemoveHelpers) == 0
} }
func (cs *ChangeSet) Summary() string { func (cs *ChangeSet) Summary() string {
@@ -49,6 +59,12 @@ func (cs *ChangeSet) Summary() string {
fmt.Fprintf(&b, " - [%s] %s (handle %d)\n", r.Chain, r.Tag, r.Handle) fmt.Fprintf(&b, " - [%s] %s (handle %d)\n", r.Chain, r.Tag, r.Handle)
} }
} }
for _, h := range cs.AddHelpers {
fmt.Fprintf(&b, " + ct helper %q\n", h.Name)
}
for _, n := range cs.RemoveHelpers {
fmt.Fprintf(&b, " - ct helper %q\n", n)
}
return b.String() return b.String()
} }
@@ -56,7 +72,7 @@ func (cs *ChangeSet) Summary() string {
// each chain, replace the middle, and insert the new rules before the first kept // each chain, replace the middle, and insert the new rules before the first kept
// suffix rule (or append when there is none). // suffix rule (or append when there is none).
func computeDiff(current, desired *FirewallState) *ChangeSet { func computeDiff(current, desired *FirewallState) *ChangeSet {
cs := &ChangeSet{} cs := diffHelpers(current, desired)
chains := make([]string, 0, len(current.Rules)+len(desired.Rules)) chains := make([]string, 0, len(current.Rules)+len(desired.Rules))
for c := range current.Rules { for c := range current.Rules {
@@ -109,7 +125,10 @@ func ruleEqual(a, b ManagedRule) bool {
// restoreChangeSet replaces every managed rule in current with the snapshot's, // restoreChangeSet replaces every managed rule in current with the snapshot's,
// in snapshot order, so a restore cannot reorder rules. // in snapshot order, so a restore cannot reorder rules.
func restoreChangeSet(current, snap *FirewallState) *ChangeSet { func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
cs := &ChangeSet{} cs := &ChangeSet{AddHelpers: snap.Helpers}
for _, h := range current.Helpers {
cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name)
}
for _, rules := range current.Rules { for _, rules := range current.Rules {
for _, r := range rules { for _, r := range rules {
if r.Tag != "" { if r.Tag != "" {
@@ -131,3 +150,29 @@ func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
} }
return cs return cs
} }
// diffHelpers replaces any ct helper object that is missing or differs.
// L3Proto is ignored: the kernel narrows inet to ip/ip6 for single-family helpers such as pptp.
func diffHelpers(current, desired *FirewallState) *ChangeSet {
cs := &ChangeSet{}
same := func(a, b expr.CtHelper) bool { return a.Name == b.Name && a.L4Proto == b.L4Proto }
find := func(hs []Helper, name string) (expr.CtHelper, bool) {
for _, h := range hs {
if h.Name == name {
return h.Helper, true
}
}
return expr.CtHelper{}, false
}
for _, h := range current.Helpers {
if want, ok := find(desired.Helpers, h.Name); !ok || !same(want, h.Helper) {
cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name)
}
}
for _, h := range desired.Helpers {
if have, ok := find(current.Helpers, h.Name); !ok || !same(have, h.Helper) {
cs.AddHelpers = append(cs.AddHelpers, h)
}
}
return cs
}
+40
View File
@@ -5,6 +5,7 @@ import (
"testing" "testing"
"github.com/google/nftables/expr" "github.com/google/nftables/expr"
"golang.org/x/sys/unix"
) )
func TestRestoreChangeSet(t *testing.T) { func TestRestoreChangeSet(t *testing.T) {
@@ -64,3 +65,42 @@ func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) {
t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add)) t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add))
} }
} }
func TestDiffHelpers(t *testing.T) {
h := func(name, typ string, l3 uint16, l4 uint8) Helper {
return Helper{Name: name, Helper: expr.CtHelper{Name: typ, L3Proto: l3, L4Proto: l4}}
}
ftp := h("ftp", "ftp", unix.NFPROTO_INET, unix.IPPROTO_TCP)
tftp := h("tftp", "tftp", unix.NFPROTO_INET, unix.IPPROTO_UDP)
sipUDP := h("sip", "sip", unix.NFPROTO_INET, unix.IPPROTO_UDP)
sipTCP := h("sip", "sip", unix.NFPROTO_INET, unix.IPPROTO_TCP)
current := &FirewallState{Helpers: []Helper{tftp, ftp, sipUDP}}
desired := &FirewallState{Helpers: []Helper{ftp, sipTCP}}
cs := computeDiff(current, desired)
if !reflect.DeepEqual(cs.RemoveHelpers, []string{"tftp", "sip"}) {
t.Errorf("remove = %v", cs.RemoveHelpers)
}
if !reflect.DeepEqual(cs.AddHelpers, []Helper{sipTCP}) {
t.Errorf("add = %v", cs.AddHelpers)
}
// Restore recreates every helper so kernel listing order matches the snapshot.
cs = restoreChangeSet(desired, current)
if !reflect.DeepEqual(cs.RemoveHelpers, []string{"ftp", "sip"}) || !reflect.DeepEqual(cs.AddHelpers, current.Helpers) {
t.Errorf("restore = -%v +%v", cs.RemoveHelpers, cs.AddHelpers)
}
live := &FirewallState{Helpers: []Helper{h("pptp", "pptp", unix.NFPROTO_IPV4, unix.IPPROTO_TCP)}}
want := &FirewallState{Helpers: []Helper{h("pptp", "pptp", unix.NFPROTO_INET, unix.IPPROTO_TCP)}}
if cs := computeDiff(live, want); !cs.Empty() {
t.Errorf("kernel-narrowed l3proto should not diff: %+v", cs)
}
if cs := computeDiff(desired, desired); !cs.Empty() {
t.Errorf("identical helpers should be empty: %+v", cs)
}
if cs := computeDiff(&FirewallState{}, desired); cs.Empty() {
t.Error("missing helpers should not be empty")
}
}
+194 -18
View File
@@ -4,6 +4,9 @@ import (
"fmt" "fmt"
"github.com/google/nftables" "github.com/google/nftables"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
) )
@@ -14,16 +17,101 @@ type Engine struct {
} }
func NewEngine(cfg *config.Config) (*Engine, error) { func NewEngine(cfg *config.Config) (*Engine, error) {
conn, err := nftables.New() conn, err := nftables.New(nftables.WithSockOptions(largeBuffers))
if err != nil { if err != nil {
return nil, fmt.Errorf("connecting to nftables: %w", err) return nil, fmt.Errorf("connecting to nftables: %w", err)
} }
return &Engine{cfg: cfg, conn: conn}, nil return &Engine{cfg: cfg, conn: conn}, nil
} }
var tableFamilies = map[config.AddressFamily]nftables.TableFamily{
config.FamilyINET: nftables.TableFamilyINet,
config.FamilyIP: nftables.TableFamilyIPv4,
config.FamilyIP6: nftables.TableFamilyIPv6,
}
func (e *Engine) family() nftables.TableFamily {
if f, ok := tableFamilies[e.cfg.Settings.AddressFamily]; ok {
return f
}
return nftables.TableFamilyINet
}
func addressFamily(tf nftables.TableFamily) config.AddressFamily {
for f, t := range tableFamilies {
if t == tf {
return f
}
}
return config.FamilyINET
}
// withFamily is the engine for the same table name in another address family.
func (e *Engine) withFamily(f config.AddressFamily) *Engine {
cfg := *e.cfg
cfg.Settings.AddressFamily = f
return &Engine{cfg: &cfg, conn: e.conn}
}
// overlaps reports whether tables of families a and b filter the same traffic:
// inet covers both ip and ip6, which do not overlap each other.
func overlaps(a, b nftables.TableFamily) bool {
if a == b {
return false
}
return (a == nftables.TableFamilyINet && (b == nftables.TableFamilyIPv4 || b == nftables.TableFamilyIPv6)) ||
(b == nftables.TableFamilyINet && (a == nftables.TableFamilyIPv4 || a == nftables.TableFamilyIPv6))
}
// staleTables are our-named tables left by a different address_family that
// would still filter the traffic this family now owns.
func (e *Engine) staleTables() ([]*nftables.Table, error) {
tables, err := e.conn.ListTables()
if err != nil {
return nil, fmt.Errorf("listing tables: %w", err)
}
var stale []*nftables.Table
for _, t := range tables {
if t.Name == e.cfg.Settings.TableName && overlaps(e.family(), t.Family) {
stale = append(stale, t)
}
}
return stale, nil
}
// batchBufSize bounds one batch: the kernel rejects a batch larger than the
// send buffer (EMSGSIZE) and drops ACKs beyond the receive buffer (ENOBUFS)
// after committing it.
// ponytail: fixed cap of tens of thousands of rules; size per batch if exceeded.
const batchBufSize = 64 << 20
// largeBuffers raises both socket buffers, ignoring rmem_max/wmem_max when
// CAP_NET_ADMIN allows it and falling back to the capped sizes otherwise.
func largeBuffers(c *netlink.Conn) error {
rc, err := c.SyscallConn()
if err != nil {
return err
}
var serr error
err = rc.Control(func(fd uintptr) {
for _, o := range [][2]int{{unix.SO_SNDBUFFORCE, unix.SO_SNDBUF}, {unix.SO_RCVBUFFORCE, unix.SO_RCVBUF}} {
if unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[0], batchBufSize) == nil {
continue
}
if serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[1], batchBufSize); serr != nil {
return
}
}
})
if err != nil {
return err
}
return serr
}
func (e *Engine) ensureTable() *nftables.Table { func (e *Engine) ensureTable() *nftables.Table {
return e.conn.AddTable(&nftables.Table{ return e.conn.AddTable(&nftables.Table{
Family: nftables.TableFamilyINet, Family: e.family(),
Name: e.cfg.Settings.TableName, Name: e.cfg.Settings.TableName,
}) })
} }
@@ -69,6 +157,36 @@ func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftable
Hooknum: nftables.ChainHookPrerouting, Hooknum: nftables.ChainHookPrerouting,
Priority: nftables.ChainPriorityNATDest, Priority: nftables.ChainPriorityNATDest,
}, },
"helper_prerouting": {
Name: "helper_prerouting",
Table: table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookPrerouting,
Priority: nftables.ChainPriorityMangle,
},
"helper_output": {
Name: "helper_output",
Table: table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookOutput,
Priority: nftables.ChainPriorityMangle,
},
"raw_prerouting": {
Name: "raw_prerouting",
Table: table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookPrerouting,
Priority: nftables.ChainPriorityRaw,
Policy: policyPtr(nftables.ChainPolicyAccept),
},
"raw_output": {
Name: "raw_output",
Table: table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookOutput,
Priority: nftables.ChainPriorityRaw,
Policy: policyPtr(nftables.ChainPolicyAccept),
},
} }
for name, chain := range chains { for name, chain := range chains {
@@ -101,6 +219,13 @@ func (e *Engine) Apply(changes *ChangeSet) error {
} }
func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error { func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error {
stale, err := e.staleTables()
if err != nil {
return err
}
for _, t := range stale {
e.conn.DelTable(t)
}
table := e.ensureTable() table := e.ensureTable()
chains := e.ensureChains(table, policies) chains := e.ensureChains(table, policies)
@@ -112,6 +237,13 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol
}) })
} }
for _, n := range changes.RemoveHelpers {
e.conn.DeleteObject(helperObj(table, n, expr.CtHelper{}))
}
for _, h := range changes.AddHelpers {
e.conn.AddObj(helperObj(table, h.Name, h.Helper))
}
for _, r := range changes.Add { for _, r := range changes.Add {
chain, ok := chains[r.Chain] chain, ok := chains[r.Chain]
if !ok { if !ok {
@@ -135,18 +267,24 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol
} }
func (e *Engine) Flush() error { func (e *Engine) Flush() error {
tables, err := e.conn.ListTables() tables, err := e.staleTables()
if err != nil { if err != nil {
return fmt.Errorf("listing tables: %w", err) return err
} }
own, err := e.findTable()
for _, t := range tables { if err != nil {
if t.Name == e.cfg.Settings.TableName { return err
e.conn.DelTable(t)
return e.conn.Flush()
} }
if own != nil {
tables = append(tables, own)
} }
if len(tables) == 0 {
return nil return nil
}
for _, t := range tables {
e.conn.DelTable(t)
}
return e.conn.Flush()
} }
func (e *Engine) findTable() (*nftables.Table, error) { func (e *Engine) findTable() (*nftables.Table, error) {
@@ -155,7 +293,7 @@ func (e *Engine) findTable() (*nftables.Table, error) {
return nil, fmt.Errorf("listing tables: %w", err) return nil, fmt.Errorf("listing tables: %w", err)
} }
for _, t := range tables { for _, t := range tables {
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { if t.Name == e.cfg.Settings.TableName && t.Family == e.family() {
return t, nil return t, nil
} }
} }
@@ -172,7 +310,19 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
return state, err return state, err
} }
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) objs, err := e.conn.GetNamedObjects(ourTable)
if err != nil {
return nil, fmt.Errorf("listing objects: %w", err)
}
for _, o := range objs {
if no, ok := o.(*nftables.NamedObj); ok && no.Type == nftables.ObjTypeCtHelper {
if h, ok := no.Obj.(*expr.CtHelper); ok {
state.Helpers = append(state.Helpers, Helper{Name: no.Name, Helper: *h})
}
}
}
chains, err := e.conn.ListChainsOfTableFamily(e.family())
if err != nil { if err != nil {
return nil, fmt.Errorf("listing chains: %w", err) return nil, fmt.Errorf("listing chains: %w", err)
} }
@@ -202,9 +352,11 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
// survives the process that took it. // survives the process that took it.
type Snapshot struct { type Snapshot struct {
Table string `json:"table"` Table string `json:"table"`
Family config.AddressFamily `json:"family,omitempty"`
Present bool `json:"present"` Present bool `json:"present"`
Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"` Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"`
Rules map[string][]SnapshotRule `json:"rules,omitempty"` Rules map[string][]SnapshotRule `json:"rules,omitempty"`
Helpers []Helper `json:"helpers,omitempty"`
} }
// SnapshotRule is a managed rule with its expressions in netlink wire format. // SnapshotRule is a managed rule with its expressions in netlink wire format.
@@ -213,16 +365,30 @@ type SnapshotRule struct {
Exprs [][]byte `json:"exprs"` Exprs [][]byte `json:"exprs"`
} }
// Snapshot captures the live tomswall table so Restore can roll back to it. // Snapshot captures the live tomswall table so Restore can roll back to it,
// falling back to the overlapping table of another family that apply replaces.
func (e *Engine) Snapshot() (*Snapshot, error) { func (e *Engine) Snapshot() (*Snapshot, error) {
snap := &Snapshot{Table: e.cfg.Settings.TableName}
t, err := e.findTable() t, err := e.findTable()
if err != nil || t == nil { if err != nil {
return snap, err return nil, err
}
if t == nil {
stale, err := e.staleTables()
if err != nil {
return nil, err
}
// ponytail: captures one stale table; an inet config replacing both ip and ip6 restores only the first.
if len(stale) > 0 {
return e.withFamily(addressFamily(stale[0].Family)).Snapshot()
}
}
snap := &Snapshot{Table: e.cfg.Settings.TableName, Family: addressFamily(e.family())}
if t == nil {
return snap, nil
} }
snap.Present = true snap.Present = true
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) chains, err := e.conn.ListChainsOfTableFamily(e.family())
if err != nil { if err != nil {
return nil, fmt.Errorf("listing chains: %w", err) return nil, fmt.Errorf("listing chains: %w", err)
} }
@@ -237,10 +403,11 @@ func (e *Engine) Snapshot() (*Snapshot, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
snap.Rules, err = encodeState(state) snap.Rules, err = encodeState(state, byte(e.family()))
if err != nil { if err != nil {
return nil, err return nil, err
} }
snap.Helpers = state.Helpers
return snap, nil return snap, nil
} }
@@ -250,13 +417,18 @@ func (e *Engine) Restore(s *Snapshot) error {
if s.Table != e.cfg.Settings.TableName { if s.Table != e.cfg.Settings.TableName {
return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName) return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName)
} }
// Snapshots predating the family field are of the inet table.
if f := addressFamily(tableFamilies[s.Family]); f != addressFamily(e.family()) {
return e.withFamily(f).Restore(s)
}
if !s.Present { if !s.Present {
return e.Flush() return e.Flush()
} }
want, err := decodeState(s.Rules) want, err := decodeState(s.Rules, byte(e.family()))
if err != nil { if err != nil {
return err return err
} }
want.Helpers = s.Helpers
current, err := e.readCurrentState() current, err := e.readCurrentState()
if err != nil { if err != nil {
return err return err
@@ -264,6 +436,10 @@ func (e *Engine) Restore(s *Snapshot) error {
return e.apply(restoreChangeSet(current, want), s.Policies) return e.apply(restoreChangeSet(current, want), s.Policies)
} }
func helperObj(table *nftables.Table, name string, h expr.CtHelper) *nftables.NamedObj {
return &nftables.NamedObj{Table: table, Name: name, Type: nftables.ObjTypeCtHelper, Obj: &h}
}
func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy {
return &p return &p
} }
+59
View File
@@ -0,0 +1,59 @@
package nftables
import (
"os"
"strconv"
"strings"
"testing"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
)
func TestLargeBuffersRaisesSocketBuffers(t *testing.T) {
c, err := netlink.Dial(unix.NETLINK_NETFILTER, nil)
if err != nil {
t.Skipf("netlink unavailable: %v", err)
}
defer c.Close()
if err := largeBuffers(c); err != nil {
t.Fatal(err)
}
// Without CAP_NET_ADMIN the kernel caps at the sysctl max; it doubles either way.
for opt, sysctl := range map[int]string{unix.SO_RCVBUF: "rmem_max", unix.SO_SNDBUF: "wmem_max"} {
want := 2 * min(batchBufSize, procInt(t, "/proc/sys/net/core/"+sysctl))
if got := sockBuf(t, c, opt); got < want {
t.Errorf("%s-bounded buffer = %d, want >= %d", sysctl, got, want)
}
}
}
func procInt(t *testing.T, path string) int {
t.Helper()
b, err := os.ReadFile(path)
if err != nil {
t.Skipf("reading %s: %v", path, err)
}
v, err := strconv.Atoi(strings.TrimSpace(string(b)))
if err != nil {
t.Fatal(err)
}
return v
}
func sockBuf(t *testing.T, c *netlink.Conn, opt int) int {
t.Helper()
rc, err := c.SyscallConn()
if err != nil {
t.Fatal(err)
}
var v int
var serr error
if err := rc.Control(func(fd uintptr) { v, serr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt) }); err != nil {
t.Fatal(err)
}
if serr != nil {
t.Fatal(serr)
}
return v
}
+113
View File
@@ -0,0 +1,113 @@
package nftables
import (
"testing"
"github.com/google/nftables"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
type sentTable struct {
msg int
family nftables.TableFamily
}
// familyEngine fakes a kernel holding tomswall tables of the given families
// and records table creations/deletions.
func familyEngine(t *testing.T, af config.AddressFamily, live ...nftables.TableFamily) (*Engine, *[]sentTable) {
var sent []sentTable
e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) {
var out []netlink.Message
for _, m := range req {
switch m.Header.Type {
case nftType(unix.NFT_MSG_GETTABLE):
for _, f := range live {
attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}})
out = append(out, netlink.Message{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append([]byte{byte(f), 0, 0, 0}, attrs...)})
}
case nftType(unix.NFT_MSG_NEWTABLE):
sent = append(sent, sentTable{unix.NFT_MSG_NEWTABLE, nftables.TableFamily(m.Data[0])})
case nftType(unix.NFT_MSG_DELTABLE):
sent = append(sent, sentTable{unix.NFT_MSG_DELTABLE, nftables.TableFamily(m.Data[0])})
}
}
return out, nil
})
e.cfg.Settings.AddressFamily = af
return e, &sent
}
func TestApplyIPFamilyReplacesInetTable(t *testing.T) {
e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyINet, nftables.TableFamilyIPv6)
if err := e.Apply(&ChangeSet{}); err != nil {
t.Fatal(err)
}
want := []sentTable{{unix.NFT_MSG_DELTABLE, nftables.TableFamilyINet}, {unix.NFT_MSG_NEWTABLE, nftables.TableFamilyIPv4}}
if len(*sent) != len(want) || (*sent)[0] != want[0] || (*sent)[1] != want[1] {
t.Errorf("got %+v, want %+v (the ip6 table must survive)", *sent, want)
}
}
func TestApplyInetFamilyReplacesIPTables(t *testing.T) {
e, sent := familyEngine(t, config.FamilyINET, nftables.TableFamilyIPv4, nftables.TableFamilyIPv6)
if err := e.Apply(&ChangeSet{}); err != nil {
t.Fatal(err)
}
var dels int
for _, s := range *sent {
if s.msg == unix.NFT_MSG_DELTABLE {
dels++
}
}
if dels != 2 {
t.Errorf("want ip and ip6 tables deleted, got %+v", *sent)
}
}
func TestFlushIPFamilyKeepsIP6Table(t *testing.T) {
e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyIPv6, nftables.TableFamilyIPv4)
if err := e.Flush(); err != nil {
t.Fatal(err)
}
if len(*sent) != 1 || (*sent)[0] != (sentTable{unix.NFT_MSG_DELTABLE, nftables.TableFamilyIPv4}) {
t.Errorf("got %+v, want only the ip table deleted", *sent)
}
}
func TestSnapshotFallsBackToReplacedTable(t *testing.T) {
e, _ := familyEngine(t, config.FamilyIP, nftables.TableFamilyINet)
snap, err := e.Snapshot()
if err != nil {
t.Fatal(err)
}
if !snap.Present || snap.Family != config.FamilyINET {
t.Fatalf("want present inet snapshot, got %+v", snap)
}
// Reverting to the inet snapshot drops the tried ip table.
e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyIPv4)
if err := e.Restore(&Snapshot{Table: "tomswall", Family: config.FamilyINET, Present: true}); err != nil {
t.Fatal(err)
}
want := []sentTable{{unix.NFT_MSG_DELTABLE, nftables.TableFamilyIPv4}, {unix.NFT_MSG_NEWTABLE, nftables.TableFamilyINet}}
if len(*sent) != 2 || (*sent)[0] != want[0] || (*sent)[1] != want[1] {
t.Errorf("got %+v, want %+v", *sent, want)
}
}
func TestRejectExprsFamily(t *testing.T) {
for af, want := range map[config.AddressFamily]expr.Reject{
config.FamilyINET: {Type: unix.NFT_REJECT_ICMPX_UNREACH, Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH},
config.FamilyIP: {Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 3},
config.FamilyIP6: {Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 4},
} {
got := rejectExprs(unix.IPPROTO_UDP, af)[0].(*expr.Reject)
if *got != want {
t.Errorf("%s: got %+v, want %+v", af, *got, want)
}
}
}
+7 -10
View File
@@ -4,14 +4,11 @@ import (
"encoding/binary" "encoding/binary"
"fmt" "fmt"
"github.com/google/nftables"
"github.com/google/nftables/expr" "github.com/google/nftables/expr"
"github.com/mdlayher/netlink" "github.com/mdlayher/netlink"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
const inet = byte(nftables.TableFamilyINet)
// exprByName mirrors the expression types google/nftables can parse back from the kernel. // exprByName mirrors the expression types google/nftables can parse back from the kernel.
var exprByName = map[string]func() expr.Any{ var exprByName = map[string]func() expr.Any{
"ct": func() expr.Any { return &expr.Ct{} }, "ct": func() expr.Any { return &expr.Ct{} },
@@ -40,13 +37,13 @@ var exprByName = map[string]func() expr.Any{
"notrack": func() expr.Any { return &expr.Notrack{} }, "notrack": func() expr.Any { return &expr.Notrack{} },
} }
func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) { func encodeState(state *FirewallState, fam byte) (map[string][]SnapshotRule, error) {
out := make(map[string][]SnapshotRule, len(state.Rules)) out := make(map[string][]SnapshotRule, len(state.Rules))
for chain, rules := range state.Rules { for chain, rules := range state.Rules {
for _, r := range rules { for _, r := range rules {
sr := SnapshotRule{Tag: r.Tag} sr := SnapshotRule{Tag: r.Tag}
for _, e := range r.Exprs { for _, e := range r.Exprs {
b, err := expr.Marshal(inet, e) b, err := expr.Marshal(fam, e)
if err != nil { if err != nil {
return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err) return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err)
} }
@@ -58,13 +55,13 @@ func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) {
return out, nil return out, nil
} }
func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) { func decodeState(rules map[string][]SnapshotRule, fam byte) (*FirewallState, error) {
state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))} state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))}
for chain, rs := range rules { for chain, rs := range rules {
for _, sr := range rs { for _, sr := range rs {
r := ManagedRule{Chain: chain, Tag: sr.Tag} r := ManagedRule{Chain: chain, Tag: sr.Tag}
for _, b := range sr.Exprs { for _, b := range sr.Exprs {
e, err := decodeExpr(b) e, err := decodeExpr(b, fam)
if err != nil { if err != nil {
return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err) return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err)
} }
@@ -77,7 +74,7 @@ func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) {
} }
// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules. // decodeExpr reverses expr.Marshal, as google/nftables does when reading rules.
func decodeExpr(b []byte) (expr.Any, error) { func decodeExpr(b []byte, fam byte) (expr.Any, error) {
ad, err := netlink.NewAttributeDecoder(b) ad, err := netlink.NewAttributeDecoder(b)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -104,13 +101,13 @@ func decodeExpr(b []byte) (expr.Any, error) {
if name == "notrack" { if name == "notrack" {
return e, nil return e, nil
} }
if err := expr.Unmarshal(inet, data, e); err != nil { if err := expr.Unmarshal(fam, data, e); err != nil {
return nil, err return nil, err
} }
// A verdict is an immediate into the verdict register with no data. // A verdict is an immediate into the verdict register with no data.
if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 { if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 {
v := &expr.Verdict{} v := &expr.Verdict{}
if err := expr.Unmarshal(inet, data, v); err != nil { if err := expr.Unmarshal(fam, data, v); err != nil {
return nil, err return nil, err
} }
return v, nil return v, nil
+19 -5
View File
@@ -29,7 +29,7 @@ func TestSnapshotRulesRoundTrip(t *testing.T) {
"input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}}, "input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}},
}} }}
rules, err := encodeState(state) rules, err := encodeState(state, byte(nftables.TableFamilyINet))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -45,7 +45,7 @@ func TestSnapshotRulesRoundTrip(t *testing.T) {
if snap.Policies["input"] != nftables.ChainPolicyAccept { if snap.Policies["input"] != nftables.ChainPolicyAccept {
t.Errorf("policy lost: %v", snap.Policies) t.Errorf("policy lost: %v", snap.Policies)
} }
got, err := decodeState(snap.Rules) got, err := decodeState(snap.Rules, byte(nftables.TableFamilyINet))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -79,7 +79,7 @@ func TestSnapshotAndRestoreAbsentTable(t *testing.T) {
for _, m := range req { for _, m := range req {
sent = append(sent, m.Header.Type) sent = append(sent, m.Header.Type)
if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent { if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent {
data := []byte{inet, 0, 0, 0} data := []byte{byte(nftables.TableFamilyINet), 0, 0, 0}
attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}}) attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}})
return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil
} }
@@ -116,7 +116,7 @@ func TestRestorePresentTable(t *testing.T) {
{Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}}, {Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}},
{Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}, {Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}},
} { } {
enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}) enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}, byte(nftables.TableFamilyINet))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -132,7 +132,7 @@ func TestRestorePresentTable(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
return append([]byte{inet, 0, 0, 0}, b...) return append([]byte{byte(nftables.TableFamilyINet), 0, 0, 0}, b...)
} }
handle := make([]byte, 8) handle := make([]byte, 8)
binary.BigEndian.PutUint64(handle, 7) binary.BigEndian.PutUint64(handle, 7)
@@ -243,3 +243,17 @@ func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, e
} }
return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn} return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn}
} }
func TestEnsureChainsRawPriority(t *testing.T) {
e := testEngine(t, nil)
chains := e.ensureChains(e.ensureTable(), nil)
for name, hook := range map[string]*nftables.ChainHook{"raw_prerouting": nftables.ChainHookPrerouting, "raw_output": nftables.ChainHookOutput} {
c, ok := chains[name]
if !ok {
t.Fatalf("%s chain not declared", name)
}
if *c.Priority != *nftables.ChainPriorityRaw || *c.Hooknum != *hook || c.Type != nftables.ChainTypeFilter || *c.Policy != nftables.ChainPolicyAccept {
t.Errorf("%s: got type %s hook %d prio %d", name, c.Type, *c.Hooknum, *c.Priority)
}
}
}
+23 -1
View File
@@ -2,6 +2,7 @@ package shorewall
import ( import (
"fmt" "fmt"
"log/slog"
"strconv" "strconv"
"strings" "strings"
@@ -99,6 +100,15 @@ func convertDir(dir string, ipv6 bool) (*config.Config, error) {
return cfg, nil return cfg, nil
} }
// disposition maps a shorewall *_DISPOSITION value; unset means CONTINUE and A_ (audit) variants map to their base action.
func disposition(v string) config.PolicyAction {
v = strings.TrimPrefix(strings.ToLower(v), "a_")
if v == "" {
return config.PolicyContinue
}
return config.PolicyAction(v)
}
func subst(s string, params map[string]string) string { func subst(s string, params map[string]string) string {
if !strings.Contains(s, "$") { if !strings.Contains(s, "$") {
return s return s
@@ -115,6 +125,8 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6
if err != nil { if err != nil {
return err return err
} }
cfg.Settings.InvalidDisposition = disposition(conf["INVALID_DISPOSITION"])
cfg.Settings.UntrackedDisposition = disposition(conf["UNTRACKED_DISPOSITION"])
if conf == nil { if conf == nil {
return nil return nil
} }
@@ -130,13 +142,23 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6
} else { } else {
cfg.Settings.LogLevel = "info" cfg.Settings.LogLevel = "info"
} }
if v := conf["LOGLIMIT"]; v != "" {
if strings.HasPrefix(v, "s:") || strings.HasPrefix(v, "d:") {
slog.Warn("shorewall: per-address LOGLIMIT is not supported, limiting each log site globally", "loglimit", v)
v = v[2:]
}
if name, rest, ok := strings.Cut(v, ":"); ok && !strings.Contains(name, "/") {
slog.Warn("shorewall: named LOGLIMIT is not supported, dropping the name", "loglimit", conf["LOGLIMIT"])
v = rest
}
cfg.Settings.LogLimit = v
}
if v, ok := conf["IP_FORWARDING"]; ok { if v, ok := conf["IP_FORWARDING"]; ok {
cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep" cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep"
} }
if v, ok := conf["IMPLICIT_CONTINUE"]; ok { if v, ok := conf["IMPLICIT_CONTINUE"]; ok {
cfg.Settings.ImplicitContinue = v == "Yes" cfg.Settings.ImplicitContinue = v == "Yes"
} }
return nil return nil
} }
+51
View File
@@ -725,3 +725,54 @@ func TestIsIPv6Dir(t *testing.T) {
} }
}) })
} }
func TestConvert_Dispositions(t *testing.T) {
cases := []struct {
conf string
invalid, untracked config.PolicyAction
}{
{"", config.PolicyContinue, config.PolicyContinue},
{"IP_FORWARDING=Yes", config.PolicyContinue, config.PolicyContinue},
{"INVALID_DISPOSITION=CONTINUE\nUNTRACKED_DISPOSITION=ACCEPT", config.PolicyContinue, config.PolicyAccept},
{"INVALID_DISPOSITION=DROP\nUNTRACKED_DISPOSITION=A_DROP", config.PolicyDrop, config.PolicyDrop},
{"INVALID_DISPOSITION=A_REJECT", config.PolicyReject, config.PolicyContinue},
}
for _, tc := range cases {
dir := t.TempDir()
writeFile(t, dir, "shorewall.conf", tc.conf)
writeFile(t, dir, "zones", "fw firewall\nnet ipv4\n")
writeFile(t, dir, "interfaces", "net eth0 -\n")
writeFile(t, dir, "policy", "all all DROP\n")
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("Convert(%q): %v", tc.conf, err)
}
if cfg.Settings.InvalidDisposition != tc.invalid || cfg.Settings.UntrackedDisposition != tc.untracked {
t.Errorf("%q: got invalid=%q untracked=%q, want %q/%q", tc.conf,
cfg.Settings.InvalidDisposition, cfg.Settings.UntrackedDisposition, tc.invalid, tc.untracked)
}
if err := cfg.Validate(); err != nil {
t.Errorf("%q: Validate: %v", tc.conf, err)
}
}
}
func TestConvert_LogLimit(t *testing.T) {
for in, want := range map[string]string{
`LOGLIMIT="s:1/sec:10"`: "1/sec:10",
`LOGLIMIT=2/min`: "2/min",
`LOGLIMIT=name:1/sec:5`: "1/sec:5",
`LOGLIMIT=s:name:1/sec:5`: "1/sec:5",
`LOGLIMIT=`: "",
} {
dir := minimalShorewallDir(t)
writeFile(t, dir, "shorewall.conf", "LOG_LEVEL=info\n"+in+"\n")
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("%s: %v", in, err)
}
if cfg.Settings.LogLimit != want {
t.Errorf("%s: log_limit = %q, want %q", in, cfg.Settings.LogLimit, want)
}
}
}
+61 -26
View File
@@ -25,15 +25,15 @@ const Unit = "tomswall-try-revert"
var ( var (
// Dir holds the lock and the pending snapshot. // Dir holds the lock and the pending snapshot.
Dir = "/var/lib/tomswall" Dir = "/var/lib/tomswall"
// run executes a systemd command; replaced in tests. // Run executes a systemd command. Test hook; production code must not reassign.
run = func(name string, args ...string) error { Run = func(name string, args ...string) error {
if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { if out, err := exec.Command(name, args...).CombinedOutput(); err != nil {
return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out)))
} }
return nil return nil
} }
// restore rolls the live table back to a snapshot; replaced in tests. // Restore rolls the live table back to a snapshot. Test hook; production code must not reassign.
restore = func(s *nftables.Snapshot) error { Restore = func(s *nftables.Snapshot) error {
engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}})
if err != nil { if err != nil {
return err return err
@@ -94,23 +94,7 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error)
if err != nil { if err != nil {
return "", err return "", err
} }
f, err := os.CreateTemp(Dir, ".try-snapshot-*") if err := WriteFile(snapshotPath(), b); err != nil {
if err != nil {
return "", err
}
defer os.Remove(f.Name())
if _, err := f.Write(b); err != nil {
f.Close()
return "", err
}
if err := f.Sync(); err != nil {
f.Close()
return "", err
}
if err := f.Close(); err != nil {
return "", err
}
if err := os.Rename(f.Name(), snapshotPath()); err != nil {
return "", err return "", err
} }
@@ -119,13 +103,46 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error)
return "", discardWith(err) return "", discardWith(err)
} }
_ = disarm() // a leftover timer from an earlier try would block the unit name _ = disarm() // a leftover timer from an earlier try would block the unit name
if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, if err := Run("systemd-run", "--quiet", "--collect", "--unit", Unit,
fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil { fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil {
return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) return "", discardWith(fmt.Errorf("arming revert timer: %w", err))
} }
return id, nil return id, nil
} }
// WriteFile durably replaces path with b: temp file, fsync, rename, fsync the directory.
func WriteFile(path string, b []byte) error {
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
f, err := os.CreateTemp(dir, "."+filepath.Base(path)+"-*")
if err != nil {
return err
}
defer os.Remove(f.Name())
if _, err := f.Write(b); err != nil {
f.Close()
return err
}
if err := f.Sync(); err != nil {
f.Close()
return err
}
if err := f.Close(); err != nil {
return err
}
if err := os.Rename(f.Name(), path); err != nil {
return err
}
d, err := os.Open(dir)
if err != nil {
return err
}
defer d.Close()
return d.Sync()
}
// Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. // Discard drops the pending snapshot and timer without restoring. The caller must hold the lock.
func Discard() error { func Discard() error {
_ = disarm() _ = disarm()
@@ -143,7 +160,7 @@ func discardWith(err error) error {
} }
func disarm() error { func disarm() error {
return run("systemctl", "stop", Unit+".timer") return Run("systemctl", "stop", Unit+".timer")
} }
// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e. // Confirm keeps the tried ruleset. ok is false when no try was pending, i.e.
@@ -175,10 +192,28 @@ func Revert(id string) (reverted bool, err error) {
if err != nil || p == nil || (id != "" && p.ID != id) { if err != nil || p == nil || (id != "" && p.ID != id) {
return false, err return false, err
} }
if err := restore(p.Snapshot); err != nil { return true, restorePending(p)
return false, fmt.Errorf("restoring snapshot: %w", err) }
// Abort restores the pending snapshot after a failed apply, which may have
// committed partially. A failed restore keeps the snapshot and timer so the
// timer still reverts. The caller must hold the lock.
func Abort() error {
p, err := load()
if err != nil {
return err
} }
return true, Discard() if p == nil {
return errors.New("no pending try to abort")
}
return restorePending(p)
}
func restorePending(p *pending) error {
if err := Restore(p.Snapshot); err != nil {
return fmt.Errorf("restoring snapshot: %w", err)
}
return Discard()
} }
func load() (*pending, error) { func load() (*pending, error) {
+45 -7
View File
@@ -15,12 +15,12 @@ func setup(t *testing.T) *[]string {
t.Helper() t.Helper()
Dir = t.TempDir() Dir = t.TempDir()
var cmds []string var cmds []string
orig := run orig := Run
run = func(name string, args ...string) error { Run = func(name string, args ...string) error {
cmds = append(cmds, name+" "+strings.Join(args, " ")) cmds = append(cmds, name+" "+strings.Join(args, " "))
return nil return nil
} }
t.Cleanup(func() { run = orig }) t.Cleanup(func() { Run = orig })
return &cmds return &cmds
} }
@@ -91,7 +91,7 @@ func TestAcquireRefusesWhilePending(t *testing.T) {
func TestArmFailureDiscardsSnapshot(t *testing.T) { func TestArmFailureDiscardsSnapshot(t *testing.T) {
setup(t) setup(t)
run = func(name string, args ...string) error { Run = func(name string, args ...string) error {
if name == "systemd-run" { if name == "systemd-run" {
return errors.New("no systemd") return errors.New("no systemd")
} }
@@ -145,12 +145,12 @@ func TestConfirmAfterRevertFails(t *testing.T) {
func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot { func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot {
t.Helper() t.Helper()
var got []*nftables.Snapshot var got []*nftables.Snapshot
orig := restore orig := Restore
restore = func(s *nftables.Snapshot) error { Restore = func(s *nftables.Snapshot) error {
got = append(got, s) got = append(got, s)
return err return err
} }
t.Cleanup(func() { restore = orig }) t.Cleanup(func() { Restore = orig })
return &got return &got
} }
@@ -209,3 +209,41 @@ func TestRevertStaleIDIgnored(t *testing.T) {
t.Errorf("newer try's snapshot removed: %v", err) t.Errorf("newer try's snapshot removed: %v", err)
} }
} }
func TestAbortRestoresAndDisarms(t *testing.T) {
cmds := setup(t)
restored := stubRestore(t, nil)
snap := &nftables.Snapshot{Table: "tomswall", Present: true}
arm(t, snap)
if err := Abort(); err != nil {
t.Fatal(err)
}
if len(*restored) != 1 || !reflect.DeepEqual((*restored)[0], snap) {
t.Errorf("restored %+v, want the armed snapshot", *restored)
}
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
t.Error("snapshot not removed")
}
if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" {
t.Errorf("timer not stopped, last command %q", last)
}
}
func TestAbortFailureKeepsSnapshotAndTimer(t *testing.T) {
cmds := setup(t)
boom := errors.New("netlink down")
stubRestore(t, boom)
arm(t, &nftables.Snapshot{Table: "tomswall"})
armed := len(*cmds)
if err := Abort(); !errors.Is(err, boom) {
t.Fatalf("Abort error = %v, want %v", err, boom)
}
if _, err := os.Stat(snapshotPath()); err != nil {
t.Fatalf("snapshot gone after failed abort: %v", err)
}
if len(*cmds) != armed {
t.Errorf("revert timer touched after failed abort: %v", (*cmds)[armed:])
}
}
+11
View File
@@ -47,6 +47,17 @@ contents:
file_info: file_info:
mode: 0640 mode: 0640
# systemd unit + environment file for applying a local config at boot.
- src: packaging/tomswall.service
dst: /usr/lib/systemd/system/tomswall.service
file_info:
mode: 0644
- src: packaging/tomswall.env
dst: /etc/tomswall/tomswall.env
type: config|noreplace
file_info:
mode: 0644
# Shell completions (generated by scripts/build-rpm.sh before packaging). # Shell completions (generated by scripts/build-rpm.sh before packaging).
- src: dist/completions/tomswall.bash - src: dist/completions/tomswall.bash
dst: /usr/share/bash-completion/completions/tomswall dst: /usr/share/bash-completion/completions/tomswall
+1
View File
@@ -3,6 +3,7 @@ Description=tomswall control-plane agent (pull and apply firewall config)
Documentation=https://git.unkin.net/unkin/tomswall Documentation=https://git.unkin.net/unkin/tomswall
After=network-online.target After=network-online.target
Wants=network-online.target Wants=network-online.target
Conflicts=tomswall.service
[Service] [Service]
Type=simple Type=simple
+3
View File
@@ -0,0 +1,3 @@
# Config applied by tomswall.service: a tomswall YAML file or a shorewall directory.
TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
#TOMSWALL_CONFIG=/etc/shorewall
+25
View File
@@ -0,0 +1,25 @@
[Unit]
Description=tomswall firewall (apply local config at boot)
Documentation=https://git.unkin.net/unkin/tomswall
DefaultDependencies=no
Wants=network-pre.target
Before=network-pre.target shutdown.target
After=local-fs.target systemd-sysctl.service
Conflicts=shutdown.target tomswall-agent.service
StartLimitIntervalSec=60
StartLimitBurst=5
[Service]
Type=oneshot
RemainAfterExit=yes
Environment=TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
EnvironmentFile=-/etc/tomswall/tomswall.env
ExecStart=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
ExecReload=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
# Fails open: after StartLimitBurst failures within StartLimitIntervalSec, boot continues without the ruleset.
Restart=on-failure
RestartSec=5
# No ExecStop: stopping the unit leaves the ruleset in place (flush would open the firewall).
[Install]
WantedBy=sysinit.target
+7 -2
View File
@@ -6,8 +6,14 @@ settings:
address_family: inet address_family: inet
ip_forwarding: true ip_forwarding: true
log_level: info log_level: info
# rate limit for every log site (shorewall LOGLIMIT, global form): rate/{sec|min|hour|day}[:burst]; unset logs every hit
log_limit: 1/sec:10
table_name: tomswall table_name: tomswall
implicit_continue: false implicit_continue: false
# ct state invalid/untracked verdict: accept, drop, reject, continue (pass to rules)
# defaults: invalid drop, untracked continue (migrate defaults both to continue, as shorewall)
invalid_disposition: drop
untracked_disposition: continue
# Named port groups — reusable port+protocol combos referenced in rules # Named port groups — reusable port+protocol combos referenced in rules
portgroups: portgroups:
@@ -191,13 +197,12 @@ snat:
# conntrack: # conntrack:
# - action: notrack # - action: notrack
# source: net # source: net
# dest: fw # dest: fw:203.0.113.1
# proto: udp # proto: udp
# dport: [53] # dport: [53]
# comment: "Skip conntrack for DNS" # comment: "Skip conntrack for DNS"
# - action: helper # - action: helper
# source: loc # source: loc
# dest: net
# proto: tcp # proto: tcp
# dport: [21] # dport: [21]
# helper: ftp # helper: ftp