88 Commits

Author SHA1 Message Date
benvin afa056b454 Merge pull request 'Add boot unit applying a local config' (#34) from benvin/boot-unit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #34
2026-10-05 21:46:58 +11:00
benvin f593f7625d Merge pull request 'Revert agent generations that cut off the control plane' (#35) from benvin/agent-safe-apply into main
Reviewed-on: #35
2026-10-05 21:46:08 +11:00
benvin 7c8bd87ec0 Merge pull request 'ci: use container-rpmbuilder image' (#36) from benvin/rpmbuilder-image into main
Reviewed-on: #36
2026-10-05 21:34:18 +11:00
unkin-agent 9854b0e7b6 ci: use container-rpmbuilder image
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 14:39:45 +11:00
unkin-agent c7e02c089a Apply the cached config without safe-apply and keep reverted generations in memory
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:56:11 +11:00
unkin-agent 9092b463a0 Apply agent generations as a pending try with a revert timer
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:52:32 +11:00
unkin-agent 502d06bdda Cap boot unit restarts so it fails open
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:50 +11:00
unkin-agent 4ad55fc65e Revert agent generations that cut off the control plane
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:20 +11:00
unkin-agent 7fcbb5fad8 Order boot unit before sysinit and retry on failure
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:44:58 +11:00
unkin-agent dc406c4f56 Add boot unit applying a local config
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:41:57 +11:00
benvin 2832dd0e7d Merge pull request 'Rate-limit log sites with LOGLIMIT' (#33) from benvin/loglimit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #33
2026-10-04 18:19:06 +11:00
unkin-agent 15ab32431d Drop the name from named shorewall LOGLIMIT values
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:58:12 +11:00
unkin-agent be391ed385 Keep rule match extras on the limited log rule
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:56:23 +11:00
unkin-agent 2e8d51759d Rate-limit log sites with shorewall LOGLIMIT
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:52:37 +11:00
benvin b8410488d2 Merge pull request 'Create the nftables table in the configured address family' (#32) from benvin/v4-only-family into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #32
2026-10-04 15:43:03 +11:00
benvin 200b4d3bdf Merge pull request 'Honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION' (#31) from benvin/invalid-disposition into main
Reviewed-on: #31
2026-10-04 15:42:40 +11:00
unkin-agent d6dfeb62b6 Merge remote-tracking branch 'origin/main' into benvin/v4-only-family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/engine.go
2026-10-04 15:40:58 +11:00
benvin 11fafc0e6a Merge pull request 'Fix large-batch netlink errors and revert failed try applies' (#30) from benvin/try-apply-error into main
Reviewed-on: #30
2026-10-04 15:40:08 +11:00
unkin-agent d949fc4772 Create the nftables table in the configured address family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:36:03 +11:00
unkin-agent 2ed5b958b4 default dispositions to continue before the nil-conf return; test bad disposition values
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:35:58 +11:00
unkin-agent c3049e3ed4 honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:34:11 +11:00
unkin-agent fe689e99ed Raise netlink buffers for large batches; restore snapshot when try apply fails
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
2026-10-04 15:34:08 +11:00
benvin 1e32bde555 Merge pull request 'Emit the implied ACCEPT for DNAT and REDIRECT rules' (#29) from benvin/dnat-implied-accept into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #29
2026-10-04 15:28:05 +11:00
unkin-agent 87be6bc6af Merge remote-tracking branch 'origin/main' into benvin/dnat-implied-accept
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
2026-10-04 15:22:41 +11:00
benvin 599c483cd8 Merge pull request 'Accept DHCP on dhcp interfaces as shorewall does' (#27) from benvin/dhcp-option into main
Reviewed-on: #27
2026-10-04 15:21:51 +11:00
benvin 110b109d97 Merge pull request 'Include firewall zone in all/any rule expansion' (#28) from benvin/all-includes-fw into main
Reviewed-on: #28
2026-10-04 15:21:24 +11:00
unkin-agent ac63f65f2f Keep rule extras off the DNAT implied accept, expand all/any DNAT sources and match SPORT
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:12:29 +11:00
unkin-agent e3130b6c3b Expand all/any firewall pairs per zone and dedupe fw
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:10:49 +11:00
unkin-agent df9326330a Count firewall-expanded all/any specs for limiter guard
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:08:48 +11:00
unkin-agent b29ed2446e Emit the implied ACCEPT for DNAT and REDIRECT rules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:08:11 +11:00
unkin-agent 0551120eec Scope DHCP accept rules to IPv4
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:07:41 +11:00
unkin-agent ff4c9b63e8 Include firewall zone in all/any rule expansion
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:06:24 +11:00
unkin-agent e48e9079bd Accept DHCP on dhcp interfaces as shorewall does
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:06:11 +11:00
benvin 2dd408c54e Merge pull request 'Create a Gitea release on tag' (#26) from benvin/gitea-releases into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #26
2026-10-04 15:00:18 +11:00
unkin-agent 6fc8ae9256 Run release step on prebuilt tea image
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 14:51:14 +11:00
benvin b849626a8c Merge pull request 'Attach conntrack helpers via ct helper objects' (#25) from benvin/ct-helpers into main
Reviewed-on: #25
2026-10-04 14:17:23 +11:00
unkin-agent da32729556 Create a Gitea release with binary, RPM and checksums on tag
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 13:59:34 +11:00
unkin-agent b85ba61fe6 Merge remote-tracking branch 'origin/main' into benvin/ct-helpers
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
#	internal/nftables/engine.go
2026-10-04 13:58:25 +11:00
benvin 9abaa85d6a Merge pull request 'Compile conntrack rules into raw-priority chains' (#22) from benvin/notrack-raw into main
Reviewed-on: #22
2026-10-04 13:55:36 +11:00
unkin-agent 9078f410ee Reject conntrack fw DEST without an address in prerouting
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:56:29 +10:00
unkin-agent da65c5d0a8 Treat conntrack SOURCE all/any as global like an omitted source
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:54:37 +10:00
unkin-agent 90d2875301 Reject zone exclusions in conntrack entries
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:53:09 +10:00
unkin-agent c557c4b78a Attach conntrack helpers via ct helper objects
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-03 23:50:31 +10:00
unkin-agent bf235291d0 Fail closed on unknown/none conntrack DEST and pair all+ symmetrically
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:50:18 +10:00
unkin-agent 1c379bf5f1 Merge remote-tracking branch 'origin/main' into benvin/notrack-raw 2026-10-03 23:49:40 +10:00
unkin-agent b6d67897de Expand all!zone exclusions and fail closed on unknown zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-03 23:48:08 +10:00
unkin-agent 9432bb05c9 Reject unknown and excluded zone names in rule, blrule and conntrack specs 2026-10-03 23:48:08 +10:00
benvin 9a7d1ab35b Merge pull request 'Bump google/nftables to v0.3.0' (#23) from benvin/nftables-v0.3.0 into main
Reviewed-on: #23
2026-10-03 23:45:34 +10:00
benvin 3c5d23a811 Merge pull request 'Match source/dest zones on conntrack rules' (#24) from benvin/conntrack-zones into benvin/notrack-raw
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Reviewed-on: #24
2026-10-03 23:44:11 +10:00
unkin-agent 337490d995 Reject fw source on prerouting conntrack and pin global omitted-zone entries
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 23:01:29 +10:00
unkin-agent 6af17a4c02 Restrict raw_output to fw sources and reject unmatched prerouting dest zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:59:14 +10:00
unkin-agent 3c5f1cacd4 Match source/dest zones and addresses on conntrack rules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:56:26 +10:00
unkin-agent 445d14c61e bump google/nftables to v0.3.0
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
Set NAT.Specified on DNAT with a port so compiled exprs match kernel readback.
2026-10-03 22:56:20 +10:00
unkin-agent c9adff32f6 Compile conntrack rules into raw-priority chains
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:54:34 +10:00
benvin 63c46fec81 Merge pull request 'Match dest zone oif on output-chain rules and policies' (#20) from benvin/output-oif into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #20
2026-10-03 22:51:28 +10:00
unkin-agent 5238886c6e Merge remote-tracking branch 'origin/main' into benvin/output-oif
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler_test.go
2026-10-03 22:46:09 +10:00
benvin 532bd80a8a Merge pull request 'Match ORIGDEST in DNAT and filter rules' (#21) from benvin/origdest into main
Reviewed-on: #21
2026-10-03 22:44:37 +10:00
unkin-agent 036021d726 Treat only non-negated addresses as scoping interfaceless zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:20:41 +10:00
unkin-agent ff5a52b9e2 Reject ORIGDEST on forwarded rules and mixed-family lists
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:18:48 +10:00
unkin-agent 44e1ba852e Skip rules and policies for zones without interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:18:10 +10:00
unkin-agent 34cf6dc9ab Match ORIGDEST in DNAT and filter rules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:16:40 +10:00
unkin-agent b5be665902 Match dest zone oif on output-chain rules and policies
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-03 22:15:24 +10:00
benvin 78afe7c242 Merge pull request 'fix: stop differential apply rewriting unchanged rules' (#18) from benvin/diff-expr-equality into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #18
2026-10-03 21:24:48 +10:00
unkin-agent 1ad025a4b7 Merge remote-tracking branch 'origin/main' into benvin/diff-expr-equality
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/diff.go
2026-10-03 21:16:39 +10:00
unkin-agent aa0d3e10c2 Merge remote-tracking branch 'origin/main' into benvin/diff-expr-equality
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-03 21:15:27 +10:00
benvin 3687baabe5 Merge pull request 'Fix MSS clamp loading the option as an IPv6 exthdr' (#19) from benvin/mss-clamp-tcpopt into main
Reviewed-on: #19
2026-10-03 21:15:19 +10:00
unkin-agent abbf434af5 Merge remote-tracking branch 'origin/main' into benvin/mss-clamp-tcpopt
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_test.go
2026-10-03 21:14:10 +10:00
benvin 937556abeb Merge pull request 'Add try/confirm safe-apply with out-of-process revert' (#17) from benvin/safe-apply into main
Reviewed-on: #17
2026-10-03 21:13:57 +10:00
benvin a7be035456 Merge pull request 'Match any listed port or protocol in a rule' (#15) from benvin/multiport-multiproto into main
Reviewed-on: #15
2026-10-03 21:13:10 +10:00
benvin e6244d6bf0 Merge pull request 'Expand comma zone lists in rule source and dest' (#16) from benvin/comma-zones into benvin/multiport-multiproto
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Reviewed-on: #16
2026-10-03 21:12:52 +10:00
unkin-agent 0b92f1c2f3 Refuse mutating commands during a try and scope reverts to the try ID
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
apply, flush and purge take the try lock and fail with ErrPending, which now
names 'tomswall revert' as the recovery after a failed automatic revert. Each
try gets an ID passed to the timer's 'revert --id', so a stale timer cannot
revert a newer try. Adds tests for Revert success/failure/stale ID and for
restoring a present table at the engine level.
2026-10-03 20:56:52 +10:00
unkin-agent 6ac03e1012 Persist try snapshot and arm a systemd revert timer
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 20:52:42 +10:00
unkin-agent 7f9c010e1a Drop agent auto-revert; skip agent apply while a try is pending 2026-10-03 20:52:42 +10:00
unkin-agent df7ebb8efe fix: preserve rule order in diff and insert replacements in place
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 20:52:01 +10:00
unkin-agent 457056d58a Pick reject type by resolved protocol number
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 20:51:25 +10:00
unkin-agent 8efed72c96 Match exact all/any zone tokens, skip only interface-less ip zones, keep DNAT free of rule extras, reject '!' inside address 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 20:50:36 +10:00
unkin-agent ed3209681d Load MSS option via tcpopt exthdr op in clamp rule
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 20:50:13 +10:00
unkin-agent 852d6bf2ca Resolve common IANA protocol names and reject ports on portless protocols
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
2026-10-03 20:49:54 +10:00
unkin-agent f522589a83 Merge remote-tracking branch 'origin/benvin/multiport-multiproto' into benvin/comma-zones 2026-10-03 20:48:31 +10:00
unkin-agent e16d95fb63 fix: compare rule expressions by value in diff
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 20:48:29 +10:00
unkin-agent 9c30f1fa54 Reject unknown protocols and dports on mixed ICMP proto 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 20:47:46 +10:00
unkin-agent ad2dd9e3e7 Merge benvin/multiport-multiproto; reject limits on zone and address lists
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 20:46:37 +10:00
unkin-agent cc12c4a43a Add try/confirm safe-apply and agent auto-revert
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 20:45:57 +10:00
unkin-agent 211bacd507 Expand comma zone lists in rule source and dest
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 20:45:32 +10:00
unkin-agent 30758855ff Reject empty list elements, unknown ICMP types and limits on list 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 20:45:16 +10:00
unkin-agent 9976bc9190 Match any listed port or protocol instead of AND-ing them
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 20:41:42 +10:00
benvin 410109515e Merge pull request 'ci: switch Go toolchain steps to gobuilder + shared S3 cache' (#14) from benvin/gocache into main
Reviewed-on: #14
2026-10-02 23:56:29 +10:00
unkin-agent 27504777f8 ci: switch Go toolchain steps to gobuilder + shared S3 cache
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
Cut CI time/network by caching go build/vet/test artifacts in S3.
- bump build/test/pre-commit/release(test,build) to
  artifactapi.../gobuilder:0.1.2-alma9
- wire GOCACHEPROG via go-cache-plugin with an absolute cache dir
- add GOCACHE_* env (ci-tomswall prefix) and AWS creds from org secrets
- leave rpm package/upload steps untouched (no Go toolchain)
2026-10-02 23:53:13 +10:00
42 changed files with 5555 additions and 533 deletions
+13 -1
View File
@@ -3,9 +3,21 @@ when:
steps: steps:
- name: build - name: build
image: golang:1.23 image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/gobuilder:0.1.2-alma9
commands: commands:
- go build ./... - go build ./...
environment:
GOCACHE_S3_BUCKET: gocache
GOCACHE_S3_REGION: us-east-1
GOCACHE_S3_ENDPOINT_URL: https://s3.ceph.unkin.net
GOCACHE_S3_PATH_STYLE: "true"
GOCACHE_KEY_PREFIX: ci-tomswall
GOCACHE_METRICS: "true"
GOCACHEPROG: go-cache-plugin --cache-dir=/tmp/gocache
AWS_ACCESS_KEY_ID:
from_secret: GOCACHE_AWS_ACCESS_KEY_ID
AWS_SECRET_ACCESS_KEY:
from_secret: GOCACHE_AWS_SECRET_ACCESS_KEY
backend_options: backend_options:
kubernetes: kubernetes:
serviceAccountName: default serviceAccountName: default
+13 -1
View File
@@ -3,9 +3,21 @@ when:
steps: steps:
- name: pre-commit - name: pre-commit
image: git.unkin.net/unkin/almalinux9-gobuilder:20260606 image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/gobuilder:0.1.2-alma9
commands: commands:
- uvx pre-commit run --all-files - uvx pre-commit run --all-files
environment:
GOCACHE_S3_BUCKET: gocache
GOCACHE_S3_REGION: us-east-1
GOCACHE_S3_ENDPOINT_URL: https://s3.ceph.unkin.net
GOCACHE_S3_PATH_STYLE: "true"
GOCACHE_KEY_PREFIX: ci-tomswall
GOCACHE_METRICS: "true"
GOCACHEPROG: go-cache-plugin --cache-dir=/tmp/gocache
AWS_ACCESS_KEY_ID:
from_secret: GOCACHE_AWS_ACCESS_KEY_ID
AWS_SECRET_ACCESS_KEY:
from_secret: GOCACHE_AWS_SECRET_ACCESS_KEY
backend_options: backend_options:
kubernetes: kubernetes:
serviceAccountName: default serviceAccountName: default
+69 -3
View File
@@ -4,9 +4,21 @@ when:
steps: steps:
- name: test - name: test
image: golang:1.23 image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/gobuilder:0.1.2-alma9
commands: commands:
- go test ./... - go test ./...
environment:
GOCACHE_S3_BUCKET: gocache
GOCACHE_S3_REGION: us-east-1
GOCACHE_S3_ENDPOINT_URL: https://s3.ceph.unkin.net
GOCACHE_S3_PATH_STYLE: "true"
GOCACHE_KEY_PREFIX: ci-tomswall
GOCACHE_METRICS: "true"
GOCACHEPROG: go-cache-plugin --cache-dir=/tmp/gocache
AWS_ACCESS_KEY_ID:
from_secret: GOCACHE_AWS_ACCESS_KEY_ID
AWS_SECRET_ACCESS_KEY:
from_secret: GOCACHE_AWS_SECRET_ACCESS_KEY
backend_options: backend_options:
kubernetes: kubernetes:
serviceAccountName: default serviceAccountName: default
@@ -19,10 +31,22 @@ steps:
cpu: 2 cpu: 2
- name: build - name: build
image: git.unkin.net/unkin/almalinux9-gobuilder:20260606 image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/gobuilder:0.1.2-alma9
commands: commands:
- make dist-build VERSION=${CI_COMMIT_TAG} - make dist-build VERSION=${CI_COMMIT_TAG}
depends_on: [test] depends_on: [test]
environment:
GOCACHE_S3_BUCKET: gocache
GOCACHE_S3_REGION: us-east-1
GOCACHE_S3_ENDPOINT_URL: https://s3.ceph.unkin.net
GOCACHE_S3_PATH_STYLE: "true"
GOCACHE_KEY_PREFIX: ci-tomswall
GOCACHE_METRICS: "true"
GOCACHEPROG: go-cache-plugin --cache-dir=/tmp/gocache
AWS_ACCESS_KEY_ID:
from_secret: GOCACHE_AWS_ACCESS_KEY_ID
AWS_SECRET_ACCESS_KEY:
from_secret: GOCACHE_AWS_SECRET_ACCESS_KEY
backend_options: backend_options:
kubernetes: kubernetes:
serviceAccountName: default serviceAccountName: default
@@ -35,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]
@@ -80,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
+13 -1
View File
@@ -3,10 +3,22 @@ when:
steps: steps:
- name: test - name: test
image: golang:1.23 image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/gobuilder:0.1.2-alma9
commands: commands:
- go vet ./... - go vet ./...
- go test ./... - go test ./...
environment:
GOCACHE_S3_BUCKET: gocache
GOCACHE_S3_REGION: us-east-1
GOCACHE_S3_ENDPOINT_URL: https://s3.ceph.unkin.net
GOCACHE_S3_PATH_STYLE: "true"
GOCACHE_KEY_PREFIX: ci-tomswall
GOCACHE_METRICS: "true"
GOCACHEPROG: go-cache-plugin --cache-dir=/tmp/gocache
AWS_ACCESS_KEY_ID:
from_secret: GOCACHE_AWS_ACCESS_KEY_ID
AWS_SECRET_ACCESS_KEY:
from_secret: GOCACHE_AWS_SECRET_ACCESS_KEY
backend_options: backend_options:
kubernetes: kubernetes:
serviceAccountName: default serviceAccountName: default
+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.`,
+30
View File
@@ -0,0 +1,30 @@
package main
import (
"errors"
"os"
"path/filepath"
"testing"
"github.com/spf13/cobra"
"git.unkin.net/unkin/tomswall/internal/tryapply"
)
func TestMutatingCommandsRefuseWhileTryPending(t *testing.T) {
orig := tryapply.Dir
tryapply.Dir = t.TempDir()
t.Cleanup(func() { tryapply.Dir = orig })
if err := os.WriteFile(filepath.Join(tryapply.Dir, "try-snapshot.json"), []byte("{}"), 0o600); err != nil {
t.Fatal(err)
}
configPath = "../../tomswall.example.yaml"
for _, cmd := range []*cobra.Command{applyCmd(), flushCmd(), purgeCmd()} {
t.Run(cmd.Use, func(t *testing.T) {
if err := cmd.RunE(cmd, nil); !errors.Is(err, tryapply.ErrPending) {
t.Errorf("got %v, want ErrPending", err)
}
})
}
}
+24
View File
@@ -12,6 +12,7 @@ import (
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables" "git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/shorewall" "git.unkin.net/unkin/tomswall/internal/shorewall"
"git.unkin.net/unkin/tomswall/internal/tryapply"
) )
var configPath string var configPath string
@@ -33,6 +34,9 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`,
root.AddCommand( root.AddCommand(
applyCmd(), applyCmd(),
tryCmd(),
confirmCmd(),
revertCmd(),
planCmd(), planCmd(),
validateCmd(), validateCmd(),
statusCmd(), statusCmd(),
@@ -86,6 +90,12 @@ The firewall is never torn down — existing connections are preserved.`,
return err return err
} }
unlock, err := tryapply.Acquire()
if err != nil {
return err
}
defer unlock()
engine, err := nftables.NewEngine(cfg) engine, err := nftables.NewEngine(cfg)
if err != nil { if err != nil {
return fmt.Errorf("initializing nftables: %w", err) return fmt.Errorf("initializing nftables: %w", err)
@@ -227,6 +237,14 @@ func purgeCmd() *cobra.Command {
return err return err
} }
if !dryRun {
unlock, err := tryapply.Acquire()
if err != nil {
return err
}
defer unlock()
}
engine, err := nftables.NewEngine(cfg) engine, err := nftables.NewEngine(cfg)
if err != nil { if err != nil {
return fmt.Errorf("initializing nftables: %w", err) return fmt.Errorf("initializing nftables: %w", err)
@@ -272,6 +290,12 @@ func flushCmd() *cobra.Command {
return err return err
} }
unlock, err := tryapply.Acquire()
if err != nil {
return err
}
defer unlock()
engine, err := nftables.NewEngine(cfg) engine, err := nftables.NewEngine(cfg)
if err != nil { if err != nil {
return fmt.Errorf("initializing nftables: %w", err) return fmt.Errorf("initializing nftables: %w", err)
+168
View File
@@ -0,0 +1,168 @@
package main
import (
"fmt"
"os"
"os/signal"
"strings"
"syscall"
"time"
"github.com/spf13/cobra"
"git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/tryapply"
)
// revertGrace lets the in-process revert win before the systemd fallback fires.
const revertGrace = 30 * time.Second
func tryCmd() *cobra.Command {
var timeout time.Duration
cmd := &cobra.Command{
Use: "try",
Short: "Apply configuration and revert unless confirmed within a timeout",
Long: `Try snapshots the live tomswall table to disk, applies the configuration, then
waits for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is
restored atomically. A transient systemd timer restores it too if this process
dies. Confirm from a new session to prove new connections still work.`,
RunE: func(cmd *cobra.Command, args []string) error {
cfg, err := loadConfig()
if err != nil {
return err
}
// Registered before apply so a confirm or hangup cannot be missed.
signal.Ignore(syscall.SIGPIPE)
confirm := make(chan os.Signal, 1)
signal.Notify(confirm, syscall.SIGUSR1)
abort := make(chan os.Signal, 1)
signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM)
id, err := tryApply(cfg, timeout+revertGrace)
if err != nil || id == "" {
return err
}
fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout)
msg, err := confirmOrRevert(confirm, abort, timeout, func() (bool, error) { return tryapply.Revert(id) })
if err != nil {
return err
}
fmt.Println(msg)
return nil
},
}
cmd.Flags().DurationVar(&timeout, "timeout", 60*time.Second, "time to wait for confirmation before reverting")
return cmd
}
// tryApply applies cfg under a pending try and returns its ID, or "" when
// there was nothing to change.
func tryApply(cfg *config.Config, fallback time.Duration) (string, error) {
unlock, err := tryapply.Acquire()
if err != nil {
return "", err
}
defer unlock()
engine, err := nftables.NewEngine(cfg)
if err != nil {
return "", fmt.Errorf("initializing nftables: %w", err)
}
changes, err := engine.Plan()
if err != nil {
return "", fmt.Errorf("computing changes: %w", err)
}
if changes.Empty() {
fmt.Println("No changes needed — firewall is up to date.")
return "", nil
}
fmt.Println(changes.Summary())
snap, err := engine.Snapshot()
if err != nil {
return "", fmt.Errorf("snapshotting ruleset: %w", err)
}
id, err := tryapply.Arm(snap, os.Getpid(), fallback)
if err != nil {
return "", err
}
if err := engine.Apply(changes); err != nil {
if aerr := tryapply.Abort(); aerr != nil {
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: previous ruleset restored", err)
}
return id, nil
}
// confirmOrRevert waits for confirm; on abort or timeout it runs revert.
// A confirm already delivered wins over a simultaneous abort or timeout.
func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() (bool, error)) (string, error) {
reason := "not confirmed within " + timeout.String()
select {
case <-confirm:
return "Confirmed.", nil
case s := <-abort:
reason = "interrupted by " + s.String()
case <-time.After(timeout):
}
select {
case <-confirm:
return "Confirmed.", nil
default:
}
reverted, err := revert()
if err != nil {
return "", fmt.Errorf("%s; revert failed: %w", reason, err)
}
if !reverted {
return "Already resolved by 'tomswall confirm' or 'tomswall revert'.", nil
}
return "", fmt.Errorf("%s: previous ruleset restored", reason)
}
func confirmCmd() *cobra.Command {
return &cobra.Command{
Use: "confirm",
Short: "Keep the configuration applied by a pending 'tomswall try'",
RunE: func(cmd *cobra.Command, args []string) error {
pid, ok, err := tryapply.Confirm()
if err != nil {
return err
}
if !ok {
return fmt.Errorf("no pending 'tomswall try': it was already reverted or confirmed")
}
// A recycled PID must not be signalled.
if comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)); err == nil && strings.TrimSpace(string(comm)) == "tomswall" {
_ = syscall.Kill(pid, syscall.SIGUSR1)
}
fmt.Println("Confirmed.")
return nil
},
}
}
func revertCmd() *cobra.Command {
var id string
cmd := &cobra.Command{
Use: "revert",
Short: "Restore the ruleset saved by a pending 'tomswall try'",
RunE: func(cmd *cobra.Command, args []string) error {
reverted, err := tryapply.Revert(id)
if err != nil {
return err
}
if !reverted {
fmt.Println("No pending 'tomswall try'.")
return nil
}
fmt.Println("Previous ruleset restored.")
return nil
},
}
cmd.Flags().StringVar(&id, "id", "", "only revert the try with this ID (used by the revert timer)")
return cmd
}
+63
View File
@@ -0,0 +1,63 @@
package main
import (
"errors"
"os"
"syscall"
"testing"
"time"
)
func TestConfirmOrRevert(t *testing.T) {
tests := []struct {
name string
confirm bool
abort bool
resolved bool
revertErr error
wantRevert bool
wantErr bool
wantResolved bool
}{
{name: "confirmed keeps ruleset", confirm: true},
{name: "confirm wins over simultaneous abort", confirm: true, abort: true},
{name: "timeout reverts", wantRevert: true, wantErr: true},
{name: "hangup reverts", abort: true, wantRevert: true, wantErr: true},
{name: "revert failure surfaces", revertErr: errors.New("boom"), wantRevert: true, wantErr: true},
{name: "resolved elsewhere", resolved: true, wantRevert: true, wantResolved: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
for i := 0; i < 20; i++ { // select between ready channels is random
confirm := make(chan os.Signal, 1)
abort := make(chan os.Signal, 1)
if tt.confirm {
confirm <- syscall.SIGUSR1
}
if tt.abort {
abort <- syscall.SIGHUP
}
called := false
msg, err := confirmOrRevert(confirm, abort, 10*time.Millisecond, func() (bool, error) {
called = true
return !tt.resolved, tt.revertErr
})
if called != tt.wantRevert {
t.Fatalf("revert called = %v, want %v", called, tt.wantRevert)
}
if (err != nil) != tt.wantErr {
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
}
if tt.revertErr != nil && !errors.Is(err, tt.revertErr) {
t.Fatalf("revert error not wrapped: %v", err)
}
if tt.wantResolved && msg == "" {
t.Fatal("expected already-resolved message")
}
if !(tt.confirm && tt.abort) {
break
}
}
})
}
}
+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.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/netlink v1.7.2 // 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=
+233 -24
View File
@@ -2,20 +2,32 @@ 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"
"git.unkin.net/unkin/tomswall/internal/nftables" "git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/tryapply"
) )
// 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
@@ -24,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
@@ -61,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)
@@ -81,40 +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
} }
// revertGeneration marks generation as reverted before restoring, so a failed
// restore can never lead to re-applying it, then reports status. It is also kept
// in memory in case persisting fails. A failed
// restore leaves the snapshot and timer armed and is reported as failed.
func (a *Agent) revertGeneration(ctx context.Context, generation int64, status string, cause error, revert func() error) error {
rv := &reverted{Generation: generation, Status: status, Error: cause.Error()}
a.lastReverted = rv
if err := a.writeReverted(rv); err != nil {
slog.Error("agent: persisting reverted generation failed", "err", err)
}
suffix := ""
if rerr := revert(); rerr != nil {
suffix = fmt.Sprintf("; restore: %v; revert timer pending", rerr)
rv.Status = StatusFailed
rv.Error += suffix
if werr := a.writeReverted(rv); werr != nil {
slog.Error("agent: persisting reverted generation failed", "err", werr)
}
}
a.reportReverted(ctx, rv)
return fmt.Errorf("generation %d %s: %w%s", generation, rv.Status, cause, suffix)
}
var (
verifyAttempts = 3
verifyDelay = 2 * time.Second
verifyTimeout = 5 * time.Second
)
// errUnreachable means the API could not be reached through the new ruleset.
var errUnreachable = errors.New("control plane unreachable after apply")
// confirm reports generation as applied over a fresh connection, which proves
// the API is reachable through the new ruleset. Any HTTP response counts as
// reachable; only repeated transport failures return errUnreachable.
func (a *Agent) confirm(ctx context.Context, generation int64) error {
var err error
for i := 0; i < verifyAttempts; i++ {
if i > 0 {
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(verifyDelay):
}
}
actx, cancel := context.WithTimeout(ctx, verifyTimeout)
err = a.Client.ReportStatus(actx, Status{Status: StatusApplied, Generation: generation})
cancel()
if ctx.Err() != nil {
return ctx.Err()
}
var uerr *url.Error
if !errors.As(err, &uerr) {
if err != nil {
slog.Warn("agent: reporting status failed", "err", err)
}
return nil
}
}
return fmt.Errorf("%w: %v", errUnreachable, err)
}
// reverted is a generation rolled back after a failed apply or for severing the
// API; persisted so it is not re-applied until a newer generation arrives.
type reverted struct {
Generation int64 `json:"generation"`
Status string `json:"status,omitempty"`
Error string `json:"error,omitempty"`
Reported bool `json:"reported"`
}
func (a *Agent) revertedPath() string {
return filepath.Join(filepath.Dir(a.Cache.Path), "reverted.json")
}
func (a *Agent) readReverted() (*reverted, error) {
b, err := os.ReadFile(a.revertedPath())
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
var rv reverted
if err := json.Unmarshal(b, &rv); err != nil {
return nil, fmt.Errorf("parsing %s: %w", a.revertedPath(), err)
}
return &rv, nil
}
func (a *Agent) writeReverted(rv *reverted) error {
b, err := json.Marshal(rv)
if err != nil {
return err
}
return tryapply.WriteFile(a.revertedPath(), b)
}
// reportReverted reports rv until the control plane accepts it.
func (a *Agent) reportReverted(ctx context.Context, rv *reverted) {
if rv.Reported {
return
}
status := rv.Status
if status == "" {
status = StatusReverted
}
if err := a.Client.ReportStatus(ctx, Status{Status: status, Generation: rv.Generation, Error: rv.Error}); err != nil {
slog.Warn("agent: reporting reverted generation failed, retrying next cycle", "generation", rv.Generation, "err", err)
return
}
rv.Reported = true
if err := a.writeReverted(rv); err != nil {
slog.Warn("agent: persisting reverted generation failed", "err", err)
}
}
// EngineApplier applies via the real nftables differential engine. // EngineApplier applies via the real nftables differential engine.
type EngineApplier struct{} type EngineApplier struct{}
// Apply computes and applies the differential change set for cfg. // Apply computes and applies the differential change set for cfg, with safe
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { // 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 -13
View File
@@ -55,20 +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)
srcZone := zoneFromSpec(r.Source)
if _, ok := c.Zones[srcZone]; !ok {
return fmt.Errorf("blrules[%d]: source zone %q not defined", i, srcZone)
}
} }
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!") {
dstZone := zoneFromSpec(r.Dest)
if _, ok := c.Zones[dstZone]; !ok {
return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, dstZone)
}
} }
} }
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
} }
+75
View File
@@ -3,6 +3,7 @@ package config
import ( import (
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"strings" "strings"
"testing" "testing"
) )
@@ -854,6 +855,32 @@ func TestValidateRules(t *testing.T) {
}, },
wantErr: "source zone \"missing\" not defined", wantErr: "source zone \"missing\" not defined",
}, },
{
name: "comma zone lists are valid",
rules: []Rule{
{Action: RuleAccept, Source: "fw,loc", Dest: "loc,net:192.0.2.1,198.51.100.1"},
},
},
{
name: "undefined zone in dest list",
rules: []Rule{
{Action: RuleAccept, Source: "loc", Dest: "net,missing"},
},
wantErr: "dest zone \"missing\" not defined",
},
{
name: "negation prefixing the whole address list is valid",
rules: []Rule{
{Action: RuleAccept, Source: "net:!192.0.2.1,198.51.100.1", Dest: "fw"},
},
},
{
name: "negation inside an address list",
rules: []Rule{
{Action: RuleAccept, Source: "net", Dest: "loc:192.0.2.1,!198.51.100.1"},
},
wantErr: "'!' may only prefix the whole address list",
},
{ {
name: "all keyword is valid source", name: "all keyword is valid source",
rules: []Rule{ rules: []Rule{
@@ -1006,3 +1033,51 @@ func TestValidateSNAT(t *testing.T) {
}) })
} }
} }
func TestSplitZoneList(t *testing.T) {
tests := []struct {
in string
want []ZoneSpec
}{
{"net", []ZoneSpec{{Zone: "net"}}},
{"fw,lan,svr", []ZoneSpec{{Zone: "fw"}, {Zone: "lan"}, {Zone: "svr"}}},
{"svr:192.0.2.17", []ZoneSpec{{Zone: "svr", Addr: "192.0.2.17"}}},
{"net:192.0.2.1,198.51.100.1", []ZoneSpec{{Zone: "net", Addr: "192.0.2.1,198.51.100.1"}}},
{"lan,svr:192.0.2.17", []ZoneSpec{{Zone: "lan"}, {Zone: "svr", Addr: "192.0.2.17"}}},
{"net:2001:db8::1", []ZoneSpec{{Zone: "net", Addr: "2001:db8::1"}}},
}
for _, tt := range tests {
if got := SplitZoneList(tt.in); !reflect.DeepEqual(got, tt.want) {
t.Errorf("SplitZoneList(%q) = %+v, want %+v", tt.in, got, tt.want)
}
}
}
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",
+60 -19
View File
@@ -1,6 +1,9 @@
package config package config
import "fmt" import (
"fmt"
"strings"
)
type RuleAction string type RuleAction string
@@ -170,25 +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 _, srcPart := range splitZones(r.Source) {
srcZone := zoneFromSpec(srcPart)
if _, ok := c.Zones[srcZone]; !ok {
return fmt.Errorf("rule[%d]: source zone %q not defined", i, srcZone)
}
}
} }
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 _, dstPart := range splitZones(r.Dest) {
dstZone := zoneFromSpec(dstPart)
if _, ok := c.Zones[dstZone]; !ok {
return fmt.Errorf("rule[%d]: dest zone %q not defined", i, dstZone)
}
}
} }
} }
@@ -223,6 +214,30 @@ func (c *Config) validateRules() error {
return nil return nil
} }
// ZoneSpec is one zone of a SOURCE/DEST list; Addr is its comma-separated address list, if any.
type ZoneSpec struct{ Zone, Addr string }
// SplitZoneList parses "lan,svr:a,b": commas before the first colon separate zones,
// commas after it separate addresses of the last zone (shorewall semantics).
func SplitZoneList(spec string) []ZoneSpec {
zones, addr, _ := strings.Cut(spec, ":")
var out []ZoneSpec
for _, z := range strings.Split(zones, ",") {
if z = strings.TrimSpace(z); z != "" {
out = append(out, ZoneSpec{Zone: z})
}
}
if len(out) > 0 {
out[len(out)-1].Addr = addr
}
return out
}
// validAddrList reports whether '!' appears only at the start, negating the whole list.
func validAddrList(addr string) bool {
return !strings.Contains(strings.TrimPrefix(addr, "!"), "!")
}
// zoneFromSpec extracts the zone name from a zone spec like "net" or "net:192.168.1.0/24". // zoneFromSpec extracts the zone name from a zone spec like "net" or "net:192.168.1.0/24".
func zoneFromSpec(spec string) string { func zoneFromSpec(spec string) string {
for i, c := range spec { for i, c := range spec {
@@ -233,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
+117 -48
View File
@@ -2,6 +2,8 @@ package nftables
import ( import (
"fmt" "fmt"
"reflect"
"sort"
"strings" "strings"
"github.com/google/nftables/expr" "github.com/google/nftables/expr"
@@ -12,19 +14,31 @@ type ManagedRule struct {
Handle uint64 Handle uint64
Exprs []expr.Any Exprs []expr.Any
Tag string Tag string
// Before is the handle of the existing rule an added rule is inserted ahead of; 0 appends.
Before uint64
} }
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 {
@@ -32,7 +46,11 @@ func (cs *ChangeSet) Summary() string {
if len(cs.Add) > 0 { if len(cs.Add) > 0 {
fmt.Fprintf(&b, " + %d rule(s) to add\n", len(cs.Add)) fmt.Fprintf(&b, " + %d rule(s) to add\n", len(cs.Add))
for _, r := range cs.Add { for _, r := range cs.Add {
fmt.Fprintf(&b, " + [%s] %s\n", r.Chain, r.Tag) if r.Before != 0 {
fmt.Fprintf(&b, " + [%s] %s (before handle %d)\n", r.Chain, r.Tag, r.Before)
} else {
fmt.Fprintf(&b, " + [%s] %s\n", r.Chain, r.Tag)
}
} }
} }
if len(cs.Remove) > 0 { if len(cs.Remove) > 0 {
@@ -41,69 +59,120 @@ 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()
} }
// Rules are first-match, so order matters: keep the common prefix and suffix of
// each chain, replace the middle, and insert the new rules before the first kept
// 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)
currentByTag := make(map[string][]ManagedRule) chains := make([]string, 0, len(current.Rules)+len(desired.Rules))
for _, rules := range current.Rules { for c := range current.Rules {
for _, r := range rules { chains = append(chains, c)
}
for c := range desired.Rules {
if _, ok := current.Rules[c]; !ok {
chains = append(chains, c)
}
}
sort.Strings(chains)
for _, chain := range chains {
var cur []ManagedRule
for _, r := range current.Rules[chain] {
if r.Tag != "" { if r.Tag != "" {
currentByTag[r.Tag] = append(currentByTag[r.Tag], r) cur = append(cur, r)
} }
} }
} want := desired.Rules[chain]
desiredByTag := make(map[string][]ManagedRule) pre := 0
for _, rules := range desired.Rules { for pre < len(cur) && pre < len(want) && ruleEqual(cur[pre], want[pre]) {
for _, r := range rules { pre++
desiredByTag[r.Tag] = append(desiredByTag[r.Tag], r) }
suf := 0
for suf < len(cur)-pre && suf < len(want)-pre &&
ruleEqual(cur[len(cur)-1-suf], want[len(want)-1-suf]) {
suf++
} }
}
for tag, desiredRules := range desiredByTag { var before uint64
currentRules, exists := currentByTag[tag] if suf > 0 {
if !exists { before = cur[len(cur)-suf].Handle
cs.Add = append(cs.Add, desiredRules...)
continue
} }
if !rulesMatch(currentRules, desiredRules) { cs.Remove = append(cs.Remove, cur[pre:len(cur)-suf]...)
cs.Remove = append(cs.Remove, currentRules...) for _, r := range want[pre : len(want)-suf] {
cs.Add = append(cs.Add, desiredRules...) r.Before = before
} cs.Add = append(cs.Add, r)
}
for tag, currentRules := range currentByTag {
if _, exists := desiredByTag[tag]; !exists {
cs.Remove = append(cs.Remove, currentRules...)
} }
} }
return cs return cs
} }
func rulesMatch(a, b []ManagedRule) bool { func ruleEqual(a, b ManagedRule) bool {
if len(a) != len(b) { return a.Chain == b.Chain && a.Tag == b.Tag && reflect.DeepEqual(a.Exprs, b.Exprs)
return false
}
for i := range a {
if a[i].Chain != b[i].Chain {
return false
}
if !exprsEqual(a[i].Exprs, b[i].Exprs) {
return false
}
}
return true
} }
func exprsEqual(a, b []expr.Any) bool { // restoreChangeSet replaces every managed rule in current with the snapshot's,
if len(a) != len(b) { // in snapshot order, so a restore cannot reorder rules.
return false func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
cs := &ChangeSet{AddHelpers: snap.Helpers}
for _, h := range current.Helpers {
cs.RemoveHelpers = append(cs.RemoveHelpers, h.Name)
} }
as := fmt.Sprintf("%v", a) for _, rules := range current.Rules {
bs := fmt.Sprintf("%v", b) for _, r := range rules {
return as == bs if r.Tag != "" {
cs.Remove = append(cs.Remove, r)
}
}
}
chains := make([]string, 0, len(snap.Rules))
for c := range snap.Rules {
chains = append(chains, c)
}
sort.Strings(chains)
for _, c := range chains {
for _, r := range snap.Rules[c] {
if r.Tag != "" {
cs.Add = append(cs.Add, r)
}
}
}
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
} }
+106
View File
@@ -0,0 +1,106 @@
package nftables
import (
"reflect"
"testing"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
)
func TestRestoreChangeSet(t *testing.T) {
accept := []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}}
drop := []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
snap := &FirewallState{Rules: map[string][]ManagedRule{
"input": {
{Chain: "input", Tag: "ssh", Exprs: accept, Handle: 4},
{Chain: "input", Tag: "web", Exprs: accept, Handle: 5},
{Chain: "input", Tag: "", Exprs: drop, Handle: 6},
},
"forward": {{Chain: "forward", Tag: "fwd", Exprs: accept, Handle: 7}},
}}
current := &FirewallState{Rules: map[string][]ManagedRule{
"input": {
{Chain: "input", Tag: "web", Exprs: accept, Handle: 10},
{Chain: "input", Tag: "ssh", Exprs: drop, Handle: 11},
{Chain: "input", Tag: "", Exprs: drop, Handle: 12},
},
}}
cs := restoreChangeSet(current, snap)
var removed []uint64
for _, r := range cs.Remove {
removed = append(removed, r.Handle)
}
if len(removed) != 2 || removed[0] != 10 || removed[1] != 11 {
t.Errorf("expected managed handles [10 11] removed, untagged kept; got %v", removed)
}
var added []string
for _, r := range cs.Add {
added = append(added, r.Tag)
}
want := []string{"fwd", "ssh", "web"}
if len(added) != len(want) {
t.Fatalf("added %v, want %v", added, want)
}
for i := range want {
if added[i] != want[i] {
t.Fatalf("added %v, want %v (snapshot order per chain)", added, want)
}
}
if !reflect.DeepEqual(cs.Add[1].Exprs, accept) {
t.Error("ssh not restored to its snapshot exprs")
}
}
func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) {
current := &FirewallState{Rules: map[string][]ManagedRule{
"input": {{Chain: "input", Tag: "x", Handle: 1}},
}}
cs := restoreChangeSet(current, &FirewallState{Rules: map[string][]ManagedRule{}})
if len(cs.Remove) != 1 || len(cs.Add) != 0 {
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")
}
}
+284 -27
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,21 +17,107 @@ 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,
}) })
} }
func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain { // ensureChains declares the base chains; policies overrides their default policy.
func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftables.ChainPolicy) map[string]*nftables.Chain {
chains := map[string]*nftables.Chain{ chains := map[string]*nftables.Chain{
"input": { "input": {
Name: "input", Name: "input",
@@ -68,9 +157,42 @@ func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain
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 {
if p, ok := policies[name]; ok {
chain.Policy = policyPtr(p)
}
chains[name] = e.conn.AddChain(chain) chains[name] = e.conn.AddChain(chain)
} }
return chains return chains
@@ -93,8 +215,19 @@ func (e *Engine) Plan() (*ChangeSet, error) {
} }
func (e *Engine) Apply(changes *ChangeSet) error { func (e *Engine) Apply(changes *ChangeSet) error {
return e.apply(changes, nil)
}
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) chains := e.ensureChains(table, policies)
for _, r := range changes.Remove { for _, r := range changes.Remove {
e.conn.DelRule(&nftables.Rule{ e.conn.DelRule(&nftables.Rule{
@@ -104,35 +237,67 @@ func (e *Engine) Apply(changes *ChangeSet) error {
}) })
} }
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 {
return fmt.Errorf("unknown chain %q", r.Chain) return fmt.Errorf("unknown chain %q", r.Chain)
} }
e.conn.AddRule(&nftables.Rule{ rule := &nftables.Rule{
Table: table, Table: table,
Chain: chain, Chain: chain,
Exprs: r.Exprs, Exprs: r.Exprs,
UserData: []byte(r.Tag), UserData: []byte(r.Tag),
}) }
if r.Before != 0 {
rule.Position = r.Before
e.conn.InsertRule(rule)
} else {
e.conn.AddRule(rule)
}
} }
return e.conn.Flush() return e.conn.Flush()
} }
func (e *Engine) Flush() error { func (e *Engine) Flush() error {
tables, err := e.staleTables()
if err != nil {
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 {
e.conn.DelTable(t)
}
return e.conn.Flush()
}
func (e *Engine) findTable() (*nftables.Table, error) {
tables, err := e.conn.ListTables() tables, err := e.conn.ListTables()
if err != nil { if err != nil {
return 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 { if t.Name == e.cfg.Settings.TableName && t.Family == e.family() {
e.conn.DelTable(t) return t, nil
return e.conn.Flush()
} }
} }
return nil return nil, nil
} }
func (e *Engine) readCurrentState() (*FirewallState, error) { func (e *Engine) readCurrentState() (*FirewallState, error) {
@@ -140,26 +305,26 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
Rules: make(map[string][]ManagedRule), Rules: make(map[string][]ManagedRule),
} }
tables, err := e.conn.ListTables() ourTable, err := e.findTable()
if err != nil { if err != nil || ourTable == nil {
return state, nil return state, err
} }
var ourTable *nftables.Table objs, err := e.conn.GetNamedObjects(ourTable)
for _, t := range tables { if err != nil {
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { return nil, fmt.Errorf("listing objects: %w", err)
ourTable = t }
break 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})
}
} }
} }
if ourTable == nil { chains, err := e.conn.ListChainsOfTableFamily(e.family())
return state, nil
}
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet)
if err != nil { if err != nil {
return state, nil return nil, fmt.Errorf("listing chains: %w", err)
} }
for _, chain := range chains { for _, chain := range chains {
@@ -168,7 +333,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
} }
rules, err := e.conn.GetRules(ourTable, chain) rules, err := e.conn.GetRules(ourTable, chain)
if err != nil { if err != nil {
continue return nil, fmt.Errorf("listing rules of %s: %w", chain.Name, err)
} }
for _, rule := range rules { for _, rule := range rules {
state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{ state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{
@@ -183,6 +348,98 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
return state, nil return state, nil
} }
// Snapshot is the tomswall table as captured live, serialisable so a revert
// survives the process that took it.
type Snapshot struct {
Table string `json:"table"`
Family config.AddressFamily `json:"family,omitempty"`
Present bool `json:"present"`
Policies map[string]nftables.ChainPolicy `json:"policies,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.
type SnapshotRule struct {
Tag string `json:"tag"`
Exprs [][]byte `json:"exprs"`
}
// 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) {
t, err := e.findTable()
if err != nil {
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
chains, err := e.conn.ListChainsOfTableFamily(e.family())
if err != nil {
return nil, fmt.Errorf("listing chains: %w", err)
}
snap.Policies = make(map[string]nftables.ChainPolicy)
for _, c := range chains {
if c.Table.Name == snap.Table && c.Policy != nil {
snap.Policies[c.Name] = *c.Policy
}
}
state, err := e.readCurrentState()
if err != nil {
return nil, err
}
snap.Rules, err = encodeState(state, byte(e.family()))
if err != nil {
return nil, err
}
snap.Helpers = state.Helpers
return snap, nil
}
// Restore atomically returns the tomswall table to the snapshot: rule order
// and chain policies included, or removed if it was absent.
func (e *Engine) Restore(s *Snapshot) error {
if 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 {
return e.Flush()
}
want, err := decodeState(s.Rules, byte(e.family()))
if err != nil {
return err
}
want.Helpers = s.Helpers
current, err := e.readCurrentState()
if err != nil {
return err
}
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)
}
}
}
+116
View File
@@ -0,0 +1,116 @@
package nftables
import (
"encoding/binary"
"fmt"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
)
// exprByName mirrors the expression types google/nftables can parse back from the kernel.
var exprByName = map[string]func() expr.Any{
"ct": func() expr.Any { return &expr.Ct{} },
"range": func() expr.Any { return &expr.Range{} },
"meta": func() expr.Any { return &expr.Meta{} },
"cmp": func() expr.Any { return &expr.Cmp{} },
"counter": func() expr.Any { return &expr.Counter{} },
"objref": func() expr.Any { return &expr.Objref{} },
"payload": func() expr.Any { return &expr.Payload{} },
"lookup": func() expr.Any { return &expr.Lookup{} },
"immediate": func() expr.Any { return &expr.Immediate{} },
"bitwise": func() expr.Any { return &expr.Bitwise{} },
"redir": func() expr.Any { return &expr.Redir{} },
"nat": func() expr.Any { return &expr.NAT{} },
"limit": func() expr.Any { return &expr.Limit{} },
"quota": func() expr.Any { return &expr.Quota{} },
"dynset": func() expr.Any { return &expr.Dynset{} },
"log": func() expr.Any { return &expr.Log{} },
"exthdr": func() expr.Any { return &expr.Exthdr{} },
"connlimit": func() expr.Any { return &expr.Connlimit{} },
"queue": func() expr.Any { return &expr.Queue{} },
"flow_offload": func() expr.Any { return &expr.FlowOffload{} },
"reject": func() expr.Any { return &expr.Reject{} },
"masq": func() expr.Any { return &expr.Masq{} },
"hash": func() expr.Any { return &expr.Hash{} },
"notrack": func() expr.Any { return &expr.Notrack{} },
}
func encodeState(state *FirewallState, fam byte) (map[string][]SnapshotRule, error) {
out := make(map[string][]SnapshotRule, len(state.Rules))
for chain, rules := range state.Rules {
for _, r := range rules {
sr := SnapshotRule{Tag: r.Tag}
for _, e := range r.Exprs {
b, err := expr.Marshal(fam, e)
if err != nil {
return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err)
}
sr.Exprs = append(sr.Exprs, b)
}
out[chain] = append(out[chain], sr)
}
}
return out, nil
}
func decodeState(rules map[string][]SnapshotRule, fam byte) (*FirewallState, error) {
state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))}
for chain, rs := range rules {
for _, sr := range rs {
r := ManagedRule{Chain: chain, Tag: sr.Tag}
for _, b := range sr.Exprs {
e, err := decodeExpr(b, fam)
if err != nil {
return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err)
}
r.Exprs = append(r.Exprs, e)
}
state.Rules[chain] = append(state.Rules[chain], r)
}
}
return state, nil
}
// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules.
func decodeExpr(b []byte, fam byte) (expr.Any, error) {
ad, err := netlink.NewAttributeDecoder(b)
if err != nil {
return nil, err
}
ad.ByteOrder = binary.BigEndian
var name string
var data []byte
for ad.Next() {
switch ad.Type() {
case unix.NFTA_EXPR_NAME:
name = ad.String()
case unix.NFTA_EXPR_DATA:
data = ad.Bytes()
}
}
if err := ad.Err(); err != nil {
return nil, err
}
newExpr, ok := exprByName[name]
if !ok {
return nil, fmt.Errorf("unsupported expression %q", name)
}
e := newExpr()
if name == "notrack" {
return e, nil
}
if err := expr.Unmarshal(fam, data, e); err != nil {
return nil, err
}
// 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 {
v := &expr.Verdict{}
if err := expr.Unmarshal(fam, data, v); err != nil {
return nil, err
}
return v, nil
}
return e, nil
}
+259
View File
@@ -0,0 +1,259 @@
package nftables
import (
"bytes"
"encoding/binary"
"encoding/json"
"reflect"
"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"
)
func TestSnapshotRulesRoundTrip(t *testing.T) {
exprs := []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0, 22}},
&expr.Ct{Register: 1, Key: expr.CtKeySTATE},
&expr.Notrack{},
&expr.Verdict{Kind: expr.VerdictAccept},
}
state := &FirewallState{Rules: map[string][]ManagedRule{
"input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}},
}}
rules, err := encodeState(state, byte(nftables.TableFamilyINet))
if err != nil {
t.Fatal(err)
}
b, err := json.Marshal(&Snapshot{Table: "tomswall", Present: true, Rules: rules,
Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}})
if err != nil {
t.Fatal(err)
}
var snap Snapshot
if err := json.Unmarshal(b, &snap); err != nil {
t.Fatal(err)
}
if snap.Policies["input"] != nftables.ChainPolicyAccept {
t.Errorf("policy lost: %v", snap.Policies)
}
got, err := decodeState(snap.Rules, byte(nftables.TableFamilyINet))
if err != nil {
t.Fatal(err)
}
in := got.Rules["input"]
if len(in) != 2 || in[0].Tag != "ssh" || in[1].Tag != "drop" {
t.Fatalf("rules/order lost: %+v", in)
}
if !reflect.DeepEqual(in[0].Exprs, exprs) {
t.Errorf("exprs changed:\n got %#v\nwant %#v", in[0].Exprs, exprs)
}
if _, ok := in[1].Exprs[0].(*expr.Verdict); !ok {
t.Errorf("verdict decoded as %T", in[1].Exprs[0])
}
}
func TestEnsureChainsPolicyOverride(t *testing.T) {
e := testEngine(t, nil)
chains := e.ensureChains(e.ensureTable(), map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept})
if *chains["input"].Policy != nftables.ChainPolicyAccept {
t.Error("input policy not overridden")
}
if *chains["forward"].Policy != nftables.ChainPolicyDrop {
t.Error("forward policy should keep its default")
}
}
func TestSnapshotAndRestoreAbsentTable(t *testing.T) {
tablePresent := false
var sent []netlink.HeaderType
e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) {
for _, m := range req {
sent = append(sent, m.Header.Type)
if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent {
data := []byte{byte(nftables.TableFamilyINet), 0, 0, 0}
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 nil, nil
})
snap, err := e.Snapshot()
if err != nil {
t.Fatal(err)
}
if snap.Present || snap.Table != "tomswall" {
t.Fatalf("want absent tomswall snapshot, got %+v", snap)
}
// The try created the table; restoring the absent snapshot deletes it.
tablePresent = true
sent = nil
if err := e.Restore(snap); err != nil {
t.Fatal(err)
}
deleted := false
for _, ht := range sent {
deleted = deleted || ht == nftType(unix.NFT_MSG_DELTABLE)
}
if !deleted {
t.Errorf("table not deleted; sent %v", sent)
}
}
func TestRestorePresentTable(t *testing.T) {
want := []SnapshotRule{}
for _, r := range []ManagedRule{
{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}}},
} {
enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}, byte(nftables.TableFamilyINet))
if err != nil {
t.Fatal(err)
}
want = append(want, enc["input"]...)
}
snap := &Snapshot{Table: "tomswall", Present: true,
Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept},
Rules: map[string][]SnapshotRule{"input": want}}
// Live state: the tried ruleset left one managed rule (handle 7) in input.
attrs := func(a ...netlink.Attribute) []byte {
b, err := netlink.MarshalAttributes(a)
if err != nil {
t.Fatal(err)
}
return append([]byte{byte(nftables.TableFamilyINet), 0, 0, 0}, b...)
}
handle := make([]byte, 8)
binary.BigEndian.PutUint64(handle, 7)
var batch []netlink.Message
e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) {
if len(req) == 0 {
return nil, nil
}
reply := func(msg int, data []byte) ([]netlink.Message, error) {
return []netlink.Message{{Header: netlink.Header{Type: nftType(msg), Sequence: req[0].Header.Sequence}, Data: data}}, nil
}
switch req[0].Header.Type {
case nftType(unix.NFT_MSG_GETTABLE):
return reply(unix.NFT_MSG_NEWTABLE, attrs(netlink.Attribute{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}))
case nftType(unix.NFT_MSG_GETCHAIN):
return reply(unix.NFT_MSG_NEWCHAIN, attrs(
netlink.Attribute{Type: unix.NFTA_CHAIN_TABLE, Data: []byte("tomswall\x00")},
netlink.Attribute{Type: unix.NFTA_CHAIN_NAME, Data: []byte("input\x00")}))
case nftType(unix.NFT_MSG_GETRULE):
return reply(unix.NFT_MSG_NEWRULE, attrs(
netlink.Attribute{Type: unix.NFTA_RULE_TABLE, Data: []byte("tomswall\x00")},
netlink.Attribute{Type: unix.NFTA_RULE_CHAIN, Data: []byte("input\x00")},
netlink.Attribute{Type: unix.NFTA_RULE_HANDLE, Data: handle},
netlink.Attribute{Type: unix.NFTA_RULE_USERDATA, Data: []byte("tried")}))
}
batch = append(batch, req...)
return nil, nil
})
if err := e.Restore(snap); err != nil {
t.Fatal(err)
}
var deleted []uint64
var added []SnapshotRule
policy := map[string]uint32{}
for _, m := range batch {
ad, err := netlink.NewAttributeDecoder(m.Data[4:])
if err != nil {
t.Fatal(err)
}
ad.ByteOrder = binary.BigEndian
var name string
var r SnapshotRule
var h uint64
var pol *uint32
for ad.Next() {
switch {
case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_NAME:
name = ad.String()
case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_POLICY:
v := ad.Uint32()
pol = &v
case m.Header.Type == nftType(unix.NFT_MSG_DELRULE) && ad.Type() == unix.NFTA_RULE_HANDLE:
h = ad.Uint64()
case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_USERDATA:
r.Tag = string(ad.Bytes())
case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_EXPRESSIONS:
ad.Nested(func(nad *netlink.AttributeDecoder) error {
for nad.Next() {
r.Exprs = append(r.Exprs, bytes.Clone(nad.Bytes()))
}
return nil
})
}
}
switch m.Header.Type {
case nftType(unix.NFT_MSG_NEWCHAIN):
if pol != nil {
policy[name] = *pol
}
case nftType(unix.NFT_MSG_DELRULE):
deleted = append(deleted, h)
case nftType(unix.NFT_MSG_NEWRULE):
added = append(added, r)
}
}
if !reflect.DeepEqual(deleted, []uint64{7}) {
t.Errorf("deleted handles %v, want [7]", deleted)
}
if !reflect.DeepEqual(added, want) {
t.Errorf("restored rules differ from snapshot:\n got %+v\nwant %+v", added, want)
}
if policy["input"] != uint32(nftables.ChainPolicyAccept) || policy["forward"] != uint32(nftables.ChainPolicyDrop) {
t.Errorf("chain policies %v: input must be restored to accept, forward keep drop", policy)
}
}
func TestRestoreRejectsOtherTable(t *testing.T) {
e := testEngine(t, nil)
if err := e.Restore(&Snapshot{Table: "other"}); err == nil {
t.Error("expected table mismatch error")
}
}
func nftType(msg int) netlink.HeaderType {
return netlink.HeaderType(unix.NFNL_SUBSYS_NFTABLES<<8 | msg)
}
func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, error)) *Engine {
if dial == nil {
dial = func([]netlink.Message) ([]netlink.Message, error) { return nil, nil }
}
conn, err := nftables.New(nftables.WithTestDial(dial))
if err != nil {
t.Fatal(err)
}
return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn}
}
func TestEnsureChainsRawPriority(t *testing.T) {
e := testEngine(t, nil)
chains := e.ensureChains(e.ensureTable(), nil)
for name, hook := range map[string]*nftables.ChainHook{"raw_prerouting": nftables.ChainHookPrerouting, "raw_output": nftables.ChainHookOutput} {
c, ok := chains[name]
if !ok {
t.Fatalf("%s chain not declared", name)
}
if *c.Priority != *nftables.ChainPriorityRaw || *c.Hooknum != *hook || c.Type != nftables.ChainTypeFilter || *c.Policy != nftables.ChainPolicyAccept {
t.Errorf("%s: got type %s hook %d prio %d", name, c.Type, *c.Hooknum, *c.Priority)
}
}
}
+23 -1
View File
@@ -2,6 +2,7 @@ package shorewall
import ( import (
"fmt" "fmt"
"log/slog"
"strconv" "strconv"
"strings" "strings"
@@ -99,6 +100,15 @@ func convertDir(dir string, ipv6 bool) (*config.Config, error) {
return cfg, nil return cfg, nil
} }
// disposition maps a shorewall *_DISPOSITION value; unset means CONTINUE and A_ (audit) variants map to their base action.
func disposition(v string) config.PolicyAction {
v = strings.TrimPrefix(strings.ToLower(v), "a_")
if v == "" {
return config.PolicyContinue
}
return config.PolicyAction(v)
}
func subst(s string, params map[string]string) string { func subst(s string, params map[string]string) string {
if !strings.Contains(s, "$") { if !strings.Contains(s, "$") {
return s return s
@@ -115,6 +125,8 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6
if err != nil { if err != nil {
return err return err
} }
cfg.Settings.InvalidDisposition = disposition(conf["INVALID_DISPOSITION"])
cfg.Settings.UntrackedDisposition = disposition(conf["UNTRACKED_DISPOSITION"])
if conf == nil { if conf == nil {
return nil return nil
} }
@@ -130,13 +142,23 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6
} else { } else {
cfg.Settings.LogLevel = "info" cfg.Settings.LogLevel = "info"
} }
if v := conf["LOGLIMIT"]; v != "" {
if strings.HasPrefix(v, "s:") || strings.HasPrefix(v, "d:") {
slog.Warn("shorewall: per-address LOGLIMIT is not supported, limiting each log site globally", "loglimit", v)
v = v[2:]
}
if name, rest, ok := strings.Cut(v, ":"); ok && !strings.Contains(name, "/") {
slog.Warn("shorewall: named LOGLIMIT is not supported, dropping the name", "loglimit", conf["LOGLIMIT"])
v = rest
}
cfg.Settings.LogLimit = v
}
if v, ok := conf["IP_FORWARDING"]; ok { if v, ok := conf["IP_FORWARDING"]; ok {
cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep" cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep"
} }
if v, ok := conf["IMPLICIT_CONTINUE"]; ok { if v, ok := conf["IMPLICIT_CONTINUE"]; ok {
cfg.Settings.ImplicitContinue = v == "Yes" cfg.Settings.ImplicitContinue = v == "Yes"
} }
return nil return nil
} }
+51
View File
@@ -725,3 +725,54 @@ func TestIsIPv6Dir(t *testing.T) {
} }
}) })
} }
func TestConvert_Dispositions(t *testing.T) {
cases := []struct {
conf string
invalid, untracked config.PolicyAction
}{
{"", config.PolicyContinue, config.PolicyContinue},
{"IP_FORWARDING=Yes", config.PolicyContinue, config.PolicyContinue},
{"INVALID_DISPOSITION=CONTINUE\nUNTRACKED_DISPOSITION=ACCEPT", config.PolicyContinue, config.PolicyAccept},
{"INVALID_DISPOSITION=DROP\nUNTRACKED_DISPOSITION=A_DROP", config.PolicyDrop, config.PolicyDrop},
{"INVALID_DISPOSITION=A_REJECT", config.PolicyReject, config.PolicyContinue},
}
for _, tc := range cases {
dir := t.TempDir()
writeFile(t, dir, "shorewall.conf", tc.conf)
writeFile(t, dir, "zones", "fw firewall\nnet ipv4\n")
writeFile(t, dir, "interfaces", "net eth0 -\n")
writeFile(t, dir, "policy", "all all DROP\n")
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("Convert(%q): %v", tc.conf, err)
}
if cfg.Settings.InvalidDisposition != tc.invalid || cfg.Settings.UntrackedDisposition != tc.untracked {
t.Errorf("%q: got invalid=%q untracked=%q, want %q/%q", tc.conf,
cfg.Settings.InvalidDisposition, cfg.Settings.UntrackedDisposition, tc.invalid, tc.untracked)
}
if err := cfg.Validate(); err != nil {
t.Errorf("%q: Validate: %v", tc.conf, err)
}
}
}
func TestConvert_LogLimit(t *testing.T) {
for in, want := range map[string]string{
`LOGLIMIT="s:1/sec:10"`: "1/sec:10",
`LOGLIMIT=2/min`: "2/min",
`LOGLIMIT=name:1/sec:5`: "1/sec:5",
`LOGLIMIT=s:name:1/sec:5`: "1/sec:5",
`LOGLIMIT=`: "",
} {
dir := minimalShorewallDir(t)
writeFile(t, dir, "shorewall.conf", "LOG_LEVEL=info\n"+in+"\n")
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("%s: %v", in, err)
}
if cfg.Settings.LogLimit != want {
t.Errorf("%s: log_limit = %q, want %q", in, cfg.Settings.LogLimit, want)
}
}
}
+235
View File
@@ -0,0 +1,235 @@
// Package tryapply keeps the state of a pending 'tomswall try' on disk so the
// revert survives the try process, backed by a transient systemd timer.
package tryapply
import (
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
"syscall"
"time"
"git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables"
)
// Unit is the transient systemd unit that reverts an unconfirmed try.
const Unit = "tomswall-try-revert"
var (
// Dir holds the lock and the pending snapshot.
Dir = "/var/lib/tomswall"
// Run executes a systemd command. Test hook; production code must not reassign.
Run = func(name string, args ...string) error {
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 nil
}
// Restore rolls the live table back to a snapshot. Test hook; production code must not reassign.
Restore = func(s *nftables.Snapshot) error {
engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}})
if err != nil {
return err
}
return engine.Restore(s)
}
)
// ErrPending means a try awaits confirmation; nothing else may apply meanwhile.
var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' to keep it or 'tomswall revert' to restore the previous ruleset (also the recovery if an automatic revert failed)")
type pending struct {
ID string `json:"id"`
PID int `json:"pid"`
Snapshot *nftables.Snapshot `json:"snapshot"`
}
func snapshotPath() string { return filepath.Join(Dir, "try-snapshot.json") }
// Acquire takes the exclusive try lock, failing with ErrPending while a try is unconfirmed.
func Acquire() (unlock func(), err error) {
unlock, err = lock()
if err != nil {
return nil, err
}
if _, err := os.Stat(snapshotPath()); err == nil {
unlock()
return nil, ErrPending
}
return unlock, nil
}
func lock() (func(), error) {
if err := os.MkdirAll(Dir, 0o755); err != nil {
return nil, err
}
f, err := os.OpenFile(filepath.Join(Dir, "try.lock"), os.O_CREATE|os.O_RDWR, 0o600)
if err != nil {
return nil, err
}
if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX); err != nil {
f.Close()
return nil, fmt.Errorf("locking %s: %w", f.Name(), err)
}
return func() { f.Close() }, nil
}
// Arm persists snap and schedules an out-of-process revert after delay,
// returning the try ID that scopes later reverts to this try.
// The caller must hold the lock from Acquire.
func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error) {
raw := make([]byte, 8)
if _, err := rand.Read(raw); err != nil {
return "", err
}
id := hex.EncodeToString(raw)
b, err := json.Marshal(pending{ID: id, PID: pid, Snapshot: snap})
if err != nil {
return "", err
}
if err := WriteFile(snapshotPath(), b); err != nil {
return "", err
}
exe, err := os.Executable()
if err != nil {
return "", discardWith(err)
}
_ = disarm() // a leftover timer from an earlier try would block the unit name
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 {
return "", discardWith(fmt.Errorf("arming revert timer: %w", err))
}
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.
func Discard() error {
_ = disarm()
if err := os.Remove(snapshotPath()); err != nil && !os.IsNotExist(err) {
return err
}
return nil
}
func discardWith(err error) error {
if derr := Discard(); derr != nil {
return fmt.Errorf("%w (discarding snapshot: %v)", err, derr)
}
return err
}
func disarm() error {
return Run("systemctl", "stop", Unit+".timer")
}
// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e.
// it was already reverted; pid is the waiting try process, if any.
func Confirm() (pid int, ok bool, err error) {
unlock, err := lock()
if err != nil {
return 0, false, err
}
defer unlock()
p, err := load()
if err != nil || p == nil {
return 0, false, err
}
return p.PID, true, Discard()
}
// Revert restores the pending snapshot. A non-empty id only reverts that try,
// so a stale timer cannot undo a newer one. reverted is false when nothing
// matching was pending (already confirmed or reverted). A failed restore keeps
// the snapshot so 'tomswall revert' can retry.
func Revert(id string) (reverted bool, err error) {
unlock, err := lock()
if err != nil {
return false, err
}
defer unlock()
p, err := load()
if err != nil || p == nil || (id != "" && p.ID != id) {
return false, err
}
return true, restorePending(p)
}
// 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
}
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) {
b, err := os.ReadFile(snapshotPath())
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
var p pending
if err := json.Unmarshal(b, &p); err != nil {
return nil, fmt.Errorf("parsing %s: %w", snapshotPath(), err)
}
if p.Snapshot == nil {
return nil, fmt.Errorf("%s has no snapshot", snapshotPath())
}
return &p, nil
}
+249
View File
@@ -0,0 +1,249 @@
package tryapply
import (
"errors"
"os"
"reflect"
"strings"
"testing"
"time"
"git.unkin.net/unkin/tomswall/internal/nftables"
)
func setup(t *testing.T) *[]string {
t.Helper()
Dir = t.TempDir()
var cmds []string
orig := Run
Run = func(name string, args ...string) error {
cmds = append(cmds, name+" "+strings.Join(args, " "))
return nil
}
t.Cleanup(func() { Run = orig })
return &cmds
}
func arm(t *testing.T, snap *nftables.Snapshot) string {
t.Helper()
unlock, err := Acquire()
if err != nil {
t.Fatal(err)
}
defer unlock()
id, err := Arm(snap, 4242, 90*time.Second)
if err != nil {
t.Fatal(err)
}
return id
}
func TestArmPersistsSnapshotAndTimer(t *testing.T) {
cmds := setup(t)
snap := &nftables.Snapshot{Table: "tomswall", Present: true,
Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1, 2, 3}}}}}}
id := arm(t, snap)
info, err := os.Stat(snapshotPath())
if err != nil {
t.Fatal(err)
}
if info.Mode().Perm() != 0o600 {
t.Errorf("snapshot mode %v, want 0600", info.Mode().Perm())
}
p, err := load()
if err != nil {
t.Fatal(err)
}
if p.ID != id || id == "" || p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) {
t.Errorf("round trip mismatch: %+v", p)
}
last := (*cmds)[len(*cmds)-1]
if !strings.HasPrefix(last, "systemd-run ") || !strings.Contains(last, "--unit "+Unit) ||
!strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert --id "+id) {
t.Errorf("unexpected arm command %q", last)
}
}
func TestAbsentTableSnapshotRoundTrip(t *testing.T) {
setup(t)
arm(t, &nftables.Snapshot{Table: "tomswall"})
p, err := load()
if err != nil {
t.Fatal(err)
}
if p.Snapshot.Present || p.Snapshot.Table != "tomswall" {
t.Errorf("absent table not preserved: %+v", p.Snapshot)
}
}
func TestAcquireRefusesWhilePending(t *testing.T) {
setup(t)
arm(t, &nftables.Snapshot{Table: "tomswall"})
_, err := Acquire()
if !errors.Is(err, ErrPending) {
t.Fatalf("second try: got %v, want ErrPending", err)
}
if !strings.Contains(err.Error(), "tomswall revert") {
t.Errorf("ErrPending does not name the recovery: %v", err)
}
}
func TestArmFailureDiscardsSnapshot(t *testing.T) {
setup(t)
Run = func(name string, args ...string) error {
if name == "systemd-run" {
return errors.New("no systemd")
}
return nil
}
unlock, err := Acquire()
if err != nil {
t.Fatal(err)
}
defer unlock()
if _, err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil {
t.Fatal("expected arm error")
}
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
t.Error("snapshot left behind without a revert timer")
}
}
func TestConfirmPendingDisarms(t *testing.T) {
cmds := setup(t)
arm(t, &nftables.Snapshot{Table: "tomswall"})
pid, ok, err := Confirm()
if err != nil || !ok || pid != 4242 {
t.Fatalf("Confirm = %d, %v, %v", pid, ok, err)
}
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)
}
unlock, err := Acquire()
if err != nil {
t.Fatalf("new try refused after confirm: %v", err)
}
unlock()
}
func TestConfirmAfterRevertFails(t *testing.T) {
setup(t)
_, ok, err := Confirm()
if err != nil || ok {
t.Fatalf("Confirm with nothing pending = %v, %v; want not ok", ok, err)
}
reverted, err := Revert("")
if err != nil || reverted {
t.Fatalf("Revert with nothing pending = %v, %v", reverted, err)
}
}
func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot {
t.Helper()
var got []*nftables.Snapshot
orig := Restore
Restore = func(s *nftables.Snapshot) error {
got = append(got, s)
return err
}
t.Cleanup(func() { Restore = orig })
return &got
}
func TestRevertPendingRestoresAndDisarms(t *testing.T) {
cmds := setup(t)
restored := stubRestore(t, nil)
snap := &nftables.Snapshot{Table: "tomswall", Present: true,
Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1}}}}}}
id := arm(t, snap)
reverted, err := Revert(id)
if err != nil || !reverted {
t.Fatalf("Revert = %v, %v", reverted, 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 TestRevertFailureKeepsSnapshot(t *testing.T) {
setup(t)
boom := errors.New("netlink down")
stubRestore(t, boom)
id := arm(t, &nftables.Snapshot{Table: "tomswall"})
if _, err := Revert(id); !errors.Is(err, boom) {
t.Fatalf("Revert error = %v, want %v", err, boom)
}
if _, err := os.Stat(snapshotPath()); err != nil {
t.Fatalf("snapshot gone after failed revert: %v", err)
}
if _, err := Acquire(); !errors.Is(err, ErrPending) {
t.Errorf("failed revert must keep the try pending, got %v", err)
}
}
func TestRevertStaleIDIgnored(t *testing.T) {
setup(t)
restored := stubRestore(t, nil)
arm(t, &nftables.Snapshot{Table: "tomswall"})
reverted, err := Revert("stale-try")
if err != nil || reverted {
t.Fatalf("stale Revert = %v, %v; want no-op", reverted, err)
}
if len(*restored) != 0 {
t.Error("stale timer restored a newer try's snapshot")
}
if _, err := os.Stat(snapshotPath()); err != nil {
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