78 Commits

Author SHA1 Message Date
unkin-agent 78dbb6ad18 Merge remote-tracking branch 'origin/main' into benvin/nft-decodable
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
# Conflicts:
#	internal/nftables/compiler_test.go
2026-10-10 01:10:37 +11:00
benvin 9aee3ad7eb Merge pull request 'Exclude sub-zone hosts from wildcard parent interfaces' (#39) from benvin/wildcard-subzone-exclusion into main
Reviewed-on: #39
2026-10-10 01:08:33 +11:00
unkin-agent 9bfdf292ad Count limiter expansion after table-family filtering
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 23:52:02 +11:00
unkin-agent bc65312647 Cover per-family expansion across rules, NAT, tunnels and blrules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 23:48:53 +11:00
unkin-agent 9a65973d7f Expand address matches per family, fail compile on conflicting guards 2026-10-09 23:48:53 +11:00
unkin-agent 460eb20db5 Carve sub-zone host interfaces out of wildcard parent matches
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:41:34 +11:00
unkin-agent 42a4dab6a3 Emit one family guard per rule, none in ip/ip6 tables
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:39:23 +11:00
unkin-agent 190ff72643 Exclude sub-zone hosts from wildcard parent 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-09 23:35:32 +11:00
benvin b8ad59b053 Merge pull request 'Match zones defined by hosts entries' (#38) from benvin/hosts-zones into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #38
2026-10-09 23:28:54 +11:00
unkin-agent 3174eabd94 Honour hosts routeback for same-interface intra-zone pairs
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:15:50 +11:00
unkin-agent 0b68110220 Merge remote-tracking branch 'origin/main' into benvin/hosts-zones
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.go
#	internal/nftables/compiler_test.go
2026-10-09 23:12:36 +11:00
benvin 70df237121 Merge pull request 'Accept intra-zone traffic between different interfaces' (#37) from benvin/intrazone-multi-iface into main
Reviewed-on: #37
2026-10-09 23:10:48 +11:00
unkin-agent 799c7f3524 Guard zone host exclusions with the 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-09 22:54:07 +11:00
unkin-agent ecc349cb6f Skip fw->fw policies and treat dest-side + as intra-zone override
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:48:33 +11:00
unkin-agent 695869c80b Match zones defined by hosts 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-09 22:48:12 +11:00
unkin-agent 96a1ba8351 Accept intra-zone traffic between different interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 22:45:53 +11:00
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
38 changed files with 4383 additions and 452 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=
+231 -29
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,46 +111,219 @@ 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
// Report the FIB so the control plane can scope router enforcement. if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) {
if fib := CollectFIB(ctx); len(fib) > 0 { slog.Warn("agent: clearing reverted generation failed", "err", err)
if err := a.Client.ReportRoutes(ctx, fib); err != nil { }
slog.Warn("agent: reporting routes failed", "err", err) // Report the FIB so the control plane can scope router enforcement.
} if fib := CollectFIB(ctx); len(fib) > 0 {
if err := a.Client.ReportRoutes(ctx, fib); err != nil {
slog.Warn("agent: reporting routes failed", "err", err)
} }
} }
return nil return nil
} }
// EngineApplier applies via the real nftables differential engine. // revertGeneration marks generation as reverted before restoring, so a failed
type EngineApplier struct{} // 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)
}
// Apply computes and applies the differential change set for cfg. It refuses var (
// while a 'tomswall try' awaits confirmation. verifyAttempts = 3
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { verifyDelay = 2 * time.Second
unlock, err := tryapply.Acquire() 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 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 { if err != nil {
return err return err
} }
defer unlock() 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.
type EngineApplier struct{}
// Apply computes and applies the differential change set for cfg, with safe
// under a pending try as 'tomswall try' does. The caller holds the try lock.
func (EngineApplier) Apply(_ context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error) {
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 err := c.validateZoneRef(r.Dest); err != nil {
if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && return fmt.Errorf("blrules[%d]: dest %w", i, err)
!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)
}
}
} }
} }
return nil return nil
+24 -1
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"`
TableName string `yaml:"table_name"` // 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"`
// 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
} }
+50
View File
@@ -473,6 +473,20 @@ func TestValidateHosts(t *testing.T) {
}, },
wantErr: "interface \"eth99\" not defined in interfaces", wantErr: "interface \"eth99\" not defined in interfaces",
}, },
{
name: "host interface matched by wildcard",
zones: map[string]Zone{
"fw": {Type: ZoneFirewall},
"net": {Type: ZoneIP},
"lan": {Type: ZoneIP, Parents: []string{"net"}},
},
interfaces: []Interface{
{Zone: "net", Interface: "enp+"},
},
hosts: []Host{
{Zone: "lan", Interface: "enp2s0", Addresses: []string{"192.0.2.0/24"}},
},
},
{ {
name: "zone not defined", name: "zone not defined",
zones: map[string]Zone{ zones: map[string]Zone{
@@ -531,6 +545,13 @@ func TestValidateHosts(t *testing.T) {
}, },
wantErr: "interface required", wantErr: "interface required",
}, },
{
name: "invalid exclusion",
zones: map[string]Zone{"fw": {Type: ZoneFirewall}, "net": {Type: ZoneIP}, "loc": {Type: ZoneIP}},
interfaces: []Interface{{Zone: "net", Interface: "eth0"}},
hosts: []Host{{Zone: "loc", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}, Exclusions: []string{"192.0.2.0/24!192.0.2.7"}}},
wantErr: "invalid address",
},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -1052,3 +1073,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:], "+!")
}
+46 -5
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",
+15 -2
View File
@@ -1,6 +1,11 @@
package config package config
import "fmt" import (
"fmt"
"net/netip"
"slices"
"strings"
)
type Host struct { type Host struct {
Zone string `yaml:"zone"` Zone string `yaml:"zone"`
@@ -40,7 +45,8 @@ func (c *Config) validateHosts() error {
ifaceFound := false ifaceFound := false
for _, iface := range c.Interfaces { for _, iface := range c.Interfaces {
if iface.Interface == h.Interface || iface.PhysicalName() == h.Interface { prefix, wild := strings.CutSuffix(iface.PhysicalName(), "+")
if iface.Interface == h.Interface || iface.PhysicalName() == h.Interface || (wild && strings.HasPrefix(h.Interface, prefix)) {
ifaceFound = true ifaceFound = true
break break
} }
@@ -52,6 +58,13 @@ func (c *Config) validateHosts() error {
if !h.Dynamic && len(h.Addresses) == 0 { if !h.Dynamic && len(h.Addresses) == 0 {
return fmt.Errorf("host[%d]: at least one address required (or set dynamic: true)", i) return fmt.Errorf("host[%d]: at least one address required (or set dynamic: true)", i)
} }
for _, a := range slices.Concat(h.Addresses, h.Exclusions) {
if _, err := netip.ParsePrefix(a); err != nil {
if _, err := netip.ParseAddr(a); err != nil {
return fmt.Errorf("host[%d]: invalid address %q", i, a)
}
}
}
} }
return nil return nil
} }
+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)
} }
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+50 -5
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")
}
}
+195 -19
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()
if err != nil {
return err
}
if own != nil {
tables = append(tables, own)
}
if len(tables) == 0 {
return nil
} }
for _, t := range tables { for _, t := range tables {
if t.Name == e.cfg.Settings.TableName { e.conn.DelTable(t)
e.conn.DelTable(t)
return e.conn.Flush()
}
} }
return nil 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)
}
}
}
+239
View File
@@ -0,0 +1,239 @@
package nftables
import (
"os"
"os/exec"
"reflect"
"strings"
"testing"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
func guardCfg(af config.AddressFamily) *config.Config {
return hostsCfg(func(c *config.Config) {
c.Settings.AddressFamily = af
c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64")
c.Hosts[0].Exclusions = []string{"192.0.2.9", "2001:db8::9"}
c.Interfaces[0].Options.NoSmurfs = true
c.Rules = append(c.Rules,
config.Rule{Action: config.RuleDrop, Source: "net:203.0.113.7", Dest: "lan:192.0.2.5"},
config.Rule{Action: config.RuleDrop, Source: "net:2001:db8:1::7", Dest: "lan:2001:db8::5"},
config.Rule{Action: config.RuleDrop, Source: "vpn", Dest: "net:!192.0.2.1"},
config.Rule{Action: config.RuleDrop, Source: "vpn", Dest: "net:!192.0.2.1,2001:db8::1"},
config.Rule{Action: config.RuleDrop, Source: "vpn:203.0.113.7", Dest: "net:!2001:db8::5"},
config.Rule{Action: config.RuleDNAT, Source: "vpn", Dest: "lan:192.0.2.10", Proto: "tcp",
DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,2001:db8::5"},
config.Rule{Action: config.RuleDrop, Source: "vpn:192.0.2.77,2001:db8:7::7", Dest: "fw"})
c.Blrules = []config.BlruleRule{{Action: config.BlruleDrop, Source: "vpn:!192.0.2.1,2001:db8::1", Dest: "fw"}}
c.SNAT = []config.SNATRule{
{Action: config.SNATMasquerade, Dest: "wlo1", Source: "!192.0.2.0/24,2001:db8::/48"},
{Action: config.SNATAddress, Address: "203.0.113.1", Dest: "wlo1", Source: "!192.0.2.9,2001:db8::9"},
{Action: config.SNATAddress, Address: "203.0.113.1", Dest: "wlo1", Source: "2001:db8::/48"},
}
c.Tunnels = []config.Tunnel{{Type: "gre", Zone: "vpn", Gateways: []string{"203.0.113.50", "2001:db8:5::1"}}}
c.StaticNAT = []config.StaticNAT{
{External: "203.0.113.60", Interface: "wlo1", Internal: "192.0.2.60"},
{External: "2001:db8:6::1", Interface: "wlo1", Internal: "2001:db8::60"},
}
})
}
func TestCompile_FamilyGuardsDecodable(t *testing.T) {
addrLen := map[config.AddressFamily]uint32{config.FamilyIP: 4, config.FamilyIP6: 16}
for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} {
t.Run(string(af), func(t *testing.T) {
for chain, rules := range mustCompile(t, guardCfg(af)).Rules {
for _, r := range rules {
guards, l3 := 0, false
for i, e := range r.Exprs {
if nfprotoGuard(r.Exprs, i) != 0 {
guards++
if l3 {
t.Errorf("%s %s: guard after a network payload: %s", chain, r.Tag, describeRule(r))
}
}
if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseNetworkHeader {
l3 = true
if n := addrLen[af]; n != 0 && (p.Len == 4 || p.Len == 16) && p.Len != n {
t.Errorf("%s %s: other-family address in %s table: %s", chain, r.Tag, af, describeRule(r))
}
}
}
if max := map[bool]int{true: 1, false: 0}[af == config.FamilyINET]; guards > max {
t.Errorf("%s %s: %d family guards in %s table: %s", chain, r.Tag, guards, af, describeRule(r))
}
}
}
})
}
}
func TestCompile_FamilyGuardsPerFamily(t *testing.T) {
state := mustCompile(t, guardCfg(config.FamilyINET))
for _, tt := range []struct {
chain, tag string
want []string
}{
{"forward", "rule:3", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 !daddr=192.0.2.1",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip6 !daddr=2001:db8::/64",
"iif=tun0 oif=wlo1 ip6 daddr=2001:db8::9",
"iif=tun0 oif=enp2s0 ip6",
}},
{"forward", "rule:4", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 !daddr=192.0.2.1",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip6 !daddr=2001:db8::/64 !daddr=2001:db8::1",
"iif=tun0 oif=wlo1 ip6 daddr=2001:db8::9 !daddr=2001:db8::1",
"iif=tun0 oif=enp2s0 ip6 !daddr=2001:db8::1",
}},
{"forward", "rule:5", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 saddr=203.0.113.7",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 saddr=203.0.113.7",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 saddr=203.0.113.7",
}},
{"prerouting", "rule:6", []string{"iif=tun0 ip4 !daddr=203.0.113.5"}},
{"forward", "rule:6:accept", []string{"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.0/24 !daddr=192.0.2.9 daddr=192.0.2.10"}},
{"input", "rule:7", []string{"iif=tun0 ip4 saddr=192.0.2.77", "iif=tun0 ip6 saddr=2001:db8:7::7"}},
{"input", "blrule:0", []string{"iif=tun0 ip4 !saddr=192.0.2.1", "iif=tun0 ip6 !saddr=2001:db8::1"}},
{"postrouting", "snat:0", []string{"oif=wlo1 ip4 !saddr=192.0.2.0/24", "oif=wlo1 ip6 !saddr=2001:db8::/48"}},
{"postrouting", "snat:1", []string{"oif=wlo1 ip4 !saddr=192.0.2.9"}},
{"postrouting", "snat:2", nil},
{"input", "tunnel:0", []string{"ip4 saddr=203.0.113.50", "ip6 saddr=2001:db8:5::1"}},
{"prerouting", "staticnat:dnat:0", []string{"iif=wlo1 ip4 daddr=203.0.113.60"}},
{"postrouting", "staticnat:snat:1", []string{"oif=wlo1 ip6 saddr=2001:db8::60"}},
} {
if got := describeTagged(state, tt.chain, tt.tag); !reflect.DeepEqual(got, tt.want) {
t.Errorf("%s %s = %q, want %q", tt.chain, tt.tag, got, tt.want)
}
}
}
// TestCompile_NegatedV4DropKeepsV6 checks a single-family table keeps its family's half of a
// negated DROP: "everything except 192.0.2.1" still drops all IPv6.
func TestCompile_NegatedV4DropKeepsV6(t *testing.T) {
for af, want := range map[config.AddressFamily][]string{
config.FamilyIP: {"iif=tun0 oif=wlo1 !daddr=192.0.2.0/24 !daddr=192.0.2.1", "iif=tun0 oif=wlo1 daddr=192.0.2.9 !daddr=192.0.2.1", "iif=tun0 oif=enp2s0 !daddr=198.51.100.0/24 !daddr=192.0.2.1"},
config.FamilyIP6: {"iif=tun0 oif=wlo1 !daddr=2001:db8::/64", "iif=tun0 oif=wlo1 daddr=2001:db8::9", "iif=tun0 oif=enp2s0"},
} {
if got := describeTagged(mustCompile(t, guardCfg(af)), "forward", "rule:3"); !reflect.DeepEqual(got, want) {
t.Errorf("%s rule:3 = %q, want %q", af, got, want)
}
}
}
func TestSplitAddrs(t *testing.T) {
for in, want := range map[string][]string{
"": {""},
"192.0.2.1,2001:db8::1": {"192.0.2.1", "2001:db8::1"},
"!192.0.2.1": {"!192.0.2.1", "::/0"},
"!2001:db8::1": {"0.0.0.0/0", "!2001:db8::1"},
"!192.0.2.1,2001:db8::1,198.51.100.0/24": {"!192.0.2.1,198.51.100.0/24", "!2001:db8::1"},
} {
if got := splitAddrs(in); !reflect.DeepEqual(got, want) {
t.Errorf("splitAddrs(%q) = %q, want %q", in, got, want)
}
}
}
func TestMatchGuardedCIDR(t *testing.T) {
for in, want := range map[string]string{
"192.0.2.1": "ip4 saddr=192.0.2.1",
"2001:db8::/48": "ip6 saddr=2001:db8::/48",
"!192.0.2.1,198.51.100.0/24": "ip4 !saddr=192.0.2.1 !saddr=198.51.100.0/24",
"!2001:db8::1": "ip6 !saddr=2001:db8::1",
"0.0.0.0/0": "ip4",
"::/0": "ip6",
} {
e, err := matchSourceCIDR(in)
if err != nil {
t.Fatalf("%s: %v", in, err)
}
if got := describeRule(ManagedRule{Exprs: e}); got != want {
t.Errorf("matchSourceCIDR(%q) = %q, want %q", in, got, want)
}
}
if e, _ := matchDestCIDR("!2001:db8::5"); describeRule(ManagedRule{Exprs: e}) != "ip6 !daddr=2001:db8::5" {
t.Errorf("matchDestCIDR(!2001:db8::5) = %q", describeRule(ManagedRule{Exprs: e}))
}
if _, err := matchSourceCIDR("!nonsense"); err == nil {
t.Error("invalid negated address: want error")
}
}
func TestFamilyGuards(t *testing.T) {
v4, _ := matchSourceCIDR("192.0.2.1")
v6, _ := matchDestCIDR("2001:db8::1")
both, _ := matchDestCIDR("198.51.100.1")
state := func(e ...[]expr.Any) *FirewallState {
var r []expr.Any
for _, x := range e {
r = append(r, x...)
}
return &FirewallState{Rules: map[string][]ManagedRule{"input": {{Exprs: r, Tag: "t"}}}}
}
s := state(v4, both)
if err := familyGuards(s, config.FamilyINET); err != nil || describeRule(s.Rules["input"][0]) != "ip4 saddr=192.0.2.1 daddr=198.51.100.1" {
t.Errorf("same-family guards not merged: %v %q", err, describeRule(s.Rules["input"][0]))
}
if err := familyGuards(state(v4, v6), config.FamilyINET); err == nil || !strings.Contains(err.Error(), "conflicting") {
t.Errorf("conflicting guards: want error, got %v", err)
}
s = state(v6)
if err := familyGuards(s, config.FamilyIP); err != nil || len(s.Rules["input"]) != 0 {
t.Errorf("ip table must drop IPv6 rules: %v %v", err, s.Rules["input"])
}
s = state(v6)
if err := familyGuards(s, config.FamilyIP6); err != nil || describeRule(s.Rules["input"][0]) != "daddr=2001:db8::1" {
t.Errorf("ip6 table must strip the guard: %v %q", err, describeRule(s.Rules["input"][0]))
}
if !famsAgree(0, unix.NFPROTO_IPV4, 0, unix.NFPROTO_IPV4) || famsAgree(unix.NFPROTO_IPV4, 0, unix.NFPROTO_IPV6) {
t.Error("famsAgree")
}
}
// TestNetnsNftListDecodes applies each family in a fresh user+net namespace and requires
// nft(8) to list the ruleset and a second plan to be empty. Needs unshare and nft.
func TestNetnsNftListDecodes(t *testing.T) {
if af := os.Getenv("TOMSWALL_NETNS_CHILD"); af != "" {
netnsChild(t, config.AddressFamily(af))
return
}
if os.Getenv("TOMSWALL_NETNS_TEST") == "" {
t.Skip("set TOMSWALL_NETNS_TEST=1 to run (needs unshare and nft)")
}
for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} {
cmd := exec.Command("unshare", "-rn", os.Args[0], "-test.run=^TestNetnsNftListDecodes$", "-test.v")
cmd.Env = append(os.Environ(), "TOMSWALL_NETNS_CHILD="+string(af))
if out, err := cmd.CombinedOutput(); err != nil {
t.Errorf("%s: %v\n%s", af, err, out)
}
}
}
func netnsChild(t *testing.T, af config.AddressFamily) {
e, err := NewEngine(guardCfg(af))
if err != nil {
t.Fatal(err)
}
cs, err := e.Plan()
if err != nil {
t.Fatal(err)
}
if err := e.Apply(cs); err != nil {
t.Fatal(err)
}
if out, err := exec.Command("nft", "list", "ruleset").CombinedOutput(); err != nil {
t.Fatalf("nft list ruleset: %v\n%s", err, out)
}
if cs, err = e.Plan(); err != nil || len(cs.Add)+len(cs.Remove) != 0 {
t.Fatalf("second plan not empty: %d add, %d remove, err %v", len(cs.Add), len(cs.Remove), err)
}
}
+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)
}
}
}
+39 -9
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
} }
@@ -226,6 +248,9 @@ func convertInterfaces(dir string, cfg *config.Config, params map[string]string)
for _, row := range rows { for _, row := range rows {
zone := subst(field(row, 0), params) zone := subst(field(row, 0), params)
iface := subst(field(row, 1), params) iface := subst(field(row, 1), params)
if isDash(zone) {
zone = ""
}
intf := config.Interface{ intf := config.Interface{
Zone: zone, Zone: zone,
@@ -370,12 +395,14 @@ func convertHosts(dir string, cfg *config.Config, params map[string]string) erro
zone := subst(field(row, 0), params) zone := subst(field(row, 0), params)
hostDef := subst(field(row, 1), params) hostDef := subst(field(row, 1), params)
hostDef, excl, _ := strings.Cut(hostDef, "!")
iface, addrs := splitHostDef(hostDef) iface, addrs := splitHostDef(hostDef)
host := config.Host{ host := config.Host{
Zone: zone, Zone: zone,
Interface: iface, Interface: iface,
Addresses: addrs, Addresses: addrs,
Exclusions: splitAddrList(excl),
} }
optsStr := subst(field(row, 2), params) optsStr := subst(field(row, 2), params)
@@ -393,16 +420,19 @@ func splitHostDef(s string) (string, []string) {
if idx < 0 { if idx < 0 {
return s, nil return s, nil
} }
iface := s[:idx] return s[:idx], splitAddrList(s[idx+1:])
addrPart := s[idx+1:] }
// splitAddrList splits a comma address list, unwrapping shorewall6 [addr]/len brackets.
func splitAddrList(s string) []string {
var addrs []string var addrs []string
for _, a := range strings.Split(addrPart, ",") { for _, a := range strings.Split(s, ",") {
a = strings.TrimSpace(a) a = strings.NewReplacer("[", "", "]", "").Replace(strings.TrimSpace(a))
if a != "" { if a != "" {
addrs = append(addrs, a) addrs = append(addrs, a)
} }
} }
return iface, addrs return addrs
} }
func parseHostOptions(s string) config.HostOptions { func parseHostOptions(s string) config.HostOptions {
+79
View File
@@ -3,6 +3,7 @@ package shorewall
import ( import (
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"testing" "testing"
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
@@ -725,3 +726,81 @@ 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)
}
}
}
func TestConvert_HostsExclusions(t *testing.T) {
dir := minimalShorewallDir(t)
writeFile(t, dir, "interfaces", `
net eth0
- eth1
`)
writeFile(t, dir, "hosts", `
loc eth0:192.0.2.0/24,198.51.100.0/24!192.0.2.7,192.0.2.8 routeback
loc eth1:[2001:db8::]/64
`)
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("Convert: %v", err)
}
want := []config.Host{
{Zone: "loc", Interface: "eth0", Addresses: []string{"192.0.2.0/24", "198.51.100.0/24"},
Exclusions: []string{"192.0.2.7", "192.0.2.8"}, Options: config.HostOptions{RouteBack: true}},
{Zone: "loc", Interface: "eth1", Addresses: []string{"2001:db8::/64"}},
}
if !reflect.DeepEqual(cfg.Hosts, want) {
t.Errorf("hosts = %+v, want %+v", cfg.Hosts, want)
}
if err := cfg.Validate(); err != nil {
t.Errorf("Validate: %v", err)
}
}
+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