33 Commits

Author SHA1 Message Date
benvin 9aee3ad7eb Merge pull request 'Exclude sub-zone hosts from wildcard parent interfaces' (#39) from benvin/wildcard-subzone-exclusion into main
Reviewed-on: #39
2026-10-10 01:08:33 +11:00
unkin-agent 460eb20db5 Carve sub-zone host interfaces out of wildcard parent matches
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:41:34 +11:00
unkin-agent 190ff72643 Exclude sub-zone hosts from wildcard parent interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 23:35:32 +11:00
benvin b8ad59b053 Merge pull request 'Match zones defined by hosts entries' (#38) from benvin/hosts-zones into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #38
2026-10-09 23:28:54 +11:00
unkin-agent 3174eabd94 Honour hosts routeback for same-interface intra-zone pairs
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:15:50 +11:00
unkin-agent 0b68110220 Merge remote-tracking branch 'origin/main' into benvin/hosts-zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
2026-10-09 23:12:36 +11:00
benvin 70df237121 Merge pull request 'Accept intra-zone traffic between different interfaces' (#37) from benvin/intrazone-multi-iface into main
Reviewed-on: #37
2026-10-09 23:10:48 +11:00
unkin-agent 799c7f3524 Guard zone host exclusions with the address family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:54:07 +11:00
unkin-agent ecc349cb6f Skip fw->fw policies and treat dest-side + as intra-zone override
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:48:33 +11:00
unkin-agent 695869c80b Match zones defined by hosts entries
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:48:12 +11:00
unkin-agent 96a1ba8351 Accept intra-zone traffic between different interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 22:45:53 +11:00
benvin afa056b454 Merge pull request 'Add boot unit applying a local config' (#34) from benvin/boot-unit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #34
2026-10-05 21:46:58 +11:00
benvin f593f7625d Merge pull request 'Revert agent generations that cut off the control plane' (#35) from benvin/agent-safe-apply into main
Reviewed-on: #35
2026-10-05 21:46:08 +11:00
benvin 7c8bd87ec0 Merge pull request 'ci: use container-rpmbuilder image' (#36) from benvin/rpmbuilder-image into main
Reviewed-on: #36
2026-10-05 21:34:18 +11:00
unkin-agent 9854b0e7b6 ci: use container-rpmbuilder image
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 14:39:45 +11:00
unkin-agent c7e02c089a Apply the cached config without safe-apply and keep reverted generations in memory
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:56:11 +11:00
unkin-agent 9092b463a0 Apply agent generations as a pending try with a revert timer
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:52:32 +11:00
unkin-agent 502d06bdda Cap boot unit restarts so it fails open
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:50 +11:00
unkin-agent 4ad55fc65e Revert agent generations that cut off the control plane
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:20 +11:00
unkin-agent 7fcbb5fad8 Order boot unit before sysinit and retry on failure
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:44:58 +11:00
unkin-agent dc406c4f56 Add boot unit applying a local config
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:41:57 +11:00
benvin 2832dd0e7d Merge pull request 'Rate-limit log sites with LOGLIMIT' (#33) from benvin/loglimit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #33
2026-10-04 18:19:06 +11:00
unkin-agent 15ab32431d Drop the name from named shorewall LOGLIMIT values
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:58:12 +11:00
unkin-agent be391ed385 Keep rule match extras on the limited log rule
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:56:23 +11:00
unkin-agent 2e8d51759d Rate-limit log sites with shorewall LOGLIMIT
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:52:37 +11:00
benvin b8410488d2 Merge pull request 'Create the nftables table in the configured address family' (#32) from benvin/v4-only-family into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #32
2026-10-04 15:43:03 +11:00
benvin 200b4d3bdf Merge pull request 'Honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION' (#31) from benvin/invalid-disposition into main
Reviewed-on: #31
2026-10-04 15:42:40 +11:00
unkin-agent d6dfeb62b6 Merge remote-tracking branch 'origin/main' into benvin/v4-only-family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/engine.go
2026-10-04 15:40:58 +11:00
benvin 11fafc0e6a Merge pull request 'Fix large-batch netlink errors and revert failed try applies' (#30) from benvin/try-apply-error into main
Reviewed-on: #30
2026-10-04 15:40:08 +11:00
unkin-agent d949fc4772 Create the nftables table in the configured address family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:36:03 +11:00
unkin-agent 2ed5b958b4 default dispositions to continue before the nil-conf return; test bad disposition values
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:35:58 +11:00
unkin-agent c3049e3ed4 honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:34:11 +11:00
unkin-agent fe689e99ed Raise netlink buffers for large batches; restore snapshot when try apply fails
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
2026-10-04 15:34:08 +11:00
29 changed files with 2369 additions and 221 deletions
+1 -1
View File
@@ -59,7 +59,7 @@ steps:
cpu: 2 cpu: 2
- name: package - name: package
image: git.unkin.net/unkin/almalinux9-rpmbuilder:latest image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/rpmbuilder:0.1.0-alma9
commands: commands:
- ./scripts/build-rpm.sh ${CI_COMMIT_TAG} - ./scripts/build-rpm.sh ${CI_COMMIT_TAG}
depends_on: [build] depends_on: [build]
+7
View File
@@ -430,6 +430,13 @@ report the generation applied, giving a fleet-wide "converged / N behind" view.
source/dest disables its rule loudly, never opens it. source/dest disables its rule loudly, never opens it.
- **Adds fail closed, the control plane fails open.** Partial rollout blocks new - **Adds fail closed, the control plane fails open.** Partial rollout blocks new
flows until every hop converges; a dead API leaves the last-good posture running. flows until every hop converges; a dead API leaves the last-good posture running.
- **A generation that severs the API is reverted.** The agent applies as a
`tomswall try` does (on-disk snapshot, systemd revert timer, shared lock), then
reports `applied` over a fresh connection. If that fails at the transport level,
or the apply errors, it records the generation in
`/var/lib/tomswall/reverted.json` (skipped until a newer one arrives), restores
the snapshot and reports `reverted`/`failed`. A failed restore leaves the timer
to retry it and is reported `failed`.
--- ---
+2 -1
View File
@@ -29,7 +29,8 @@ func agentCmd() *cobra.Command {
Long: `Agent runs the control-plane pull loop: it fetches this device's compiled Long: `Agent runs the control-plane pull loop: it fetches this device's compiled
config from tomswallapi, differentially applies it, and reports the applied config from tomswallapi, differentially applies it, and reports the applied
generation back. It caches the last known-good config and, if the control plane generation back. It caches the last known-good config and, if the control plane
is unreachable, keeps applying that cache — it never fails closed. is unreachable, keeps applying that cache — it never fails closed. A new
generation that cuts the agent off from the API is reverted and reported as such.
The agent token defaults to the TOMSWALL_AGENT_TOKEN environment variable, and The agent token defaults to the TOMSWALL_AGENT_TOKEN environment variable, and
the device name defaults to the system hostname.`, the device name defaults to the system hostname.`,
+3 -3
View File
@@ -89,10 +89,10 @@ func tryApply(cfg *config.Config, fallback time.Duration) (string, error) {
return "", err return "", err
} }
if err := engine.Apply(changes); err != nil { if err := engine.Apply(changes); err != nil {
if derr := tryapply.Discard(); derr != nil { if aerr := tryapply.Abort(); aerr != nil {
err = fmt.Errorf("%w (discarding snapshot: %v)", err, derr) return "", fmt.Errorf("applying changes: %w; %v; the revert timer restores the previous ruleset within %s", err, aerr, fallback)
} }
return "", fmt.Errorf("applying changes: %w", err) return "", fmt.Errorf("applying changes: %w: previous ruleset restored", err)
} }
return id, nil return id, nil
} }
+226 -24
View File
@@ -2,8 +2,13 @@ package agent
import ( import (
"context" "context"
"encoding/json"
"errors"
"fmt" "fmt"
"log/slog" "log/slog"
"net/url"
"os"
"path/filepath"
"time" "time"
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
@@ -12,11 +17,17 @@ import (
) )
// Applier applies a translated config to the firewall. Abstracted so the run // Applier applies a translated config to the firewall. Abstracted so the run
// loop is testable without touching the kernel. // loop is testable without touching the kernel. With safe, a change is applied
// as a pending try: revert restores the previous ruleset (a failed revert leaves
// the revert timer armed) and keep drops the snapshot. Both are nil when nothing
// changed or safe is false.
type Applier interface { type Applier interface {
Apply(ctx context.Context, cfg *config.Config) error Apply(ctx context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error)
} }
// revertDelay is when the revert timer fires if the agent dies mid-apply.
const revertDelay = time.Minute
// Agent runs the pull-apply-report loop for one device. // Agent runs the pull-apply-report loop for one device.
type Agent struct { type Agent struct {
Client *Client Client *Client
@@ -25,6 +36,9 @@ type Agent struct {
Applier Applier Applier Applier
// Resolver overrides the DNS resolver (tests); nil derives it per-config. // Resolver overrides the DNS resolver (tests); nil derives it per-config.
Resolver *Resolver Resolver *Resolver
// lastReverted covers a reverted generation whose persistence failed.
lastReverted *reverted
} }
// Run loops until ctx is cancelled, applying one cycle per Interval (and once // Run loops until ctx is cancelled, applying one cycle per Interval (and once
@@ -62,16 +76,31 @@ func (a *Agent) RunOnce(ctx context.Context) error {
return fmt.Errorf("control plane unreachable and no cached config: %w", err) return fmt.Errorf("control plane unreachable and no cached config: %w", err)
} }
// Re-apply last known-good; do not report a generation we didn't fetch. // Re-apply last known-good; do not report a generation we didn't fetch.
return a.applyConfig(ctx, cached, false) return a.applyConfig(ctx, cached, nil)
} }
if err := a.Cache.Write(raw); err != nil { rv, err := a.readReverted()
slog.Warn("agent: caching config failed", "err", err) if err != nil {
return err
} }
return a.applyConfig(ctx, rc, true) if a.lastReverted != nil && (rv == nil || a.lastReverted.Generation > rv.Generation) {
rv = a.lastReverted
}
if rv != nil {
a.reportReverted(ctx, rv)
if rc.Generation <= rv.Generation {
slog.Info("agent: generation was reverted, waiting for a newer one", "generation", rc.Generation)
return nil
}
}
return a.applyConfig(ctx, rc, raw)
} }
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error { // applyConfig applies rc. A fetched config (raw != nil) is applied as a pending
// try, verified by reaching the API through the new ruleset and reverted if that
// fails; only then is it cached. The cached config is the last verified-good one,
// so it is applied plainly: there is nothing to verify it against.
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error {
resolver := a.Resolver resolver := a.Resolver
if resolver == nil { if resolver == nil {
resolver = NewResolver(rc.Resolver) resolver = NewResolver(rc.Resolver)
@@ -82,14 +111,65 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
if err != nil { if err != nil {
return fmt.Errorf("translate: %w", err) return fmt.Errorf("translate: %w", err)
} }
if err := a.Applier.Apply(ctx, cfg); err != nil {
return fmt.Errorf("apply: %w", err) unlock, err := tryapply.Acquire()
if errors.Is(err, tryapply.ErrPending) {
slog.Warn("agent: a 'tomswall try' is pending, skipping cycle")
return nil
}
if err != nil {
return err
}
defer unlock()
revert, keep, err := a.Applier.Apply(ctx, cfg, raw != nil)
if err != nil {
err = fmt.Errorf("apply: %w", err)
if raw != nil && revert != nil {
return a.revertGeneration(ctx, rc.Generation, StatusFailed, err, revert)
}
if revert != nil {
if rerr := revert(); rerr != nil {
err = fmt.Errorf("%w; restore: %v; revert timer pending", err, rerr)
}
}
if raw != nil {
if rerr := a.Client.ReportStatus(ctx, Status{Status: StatusFailed, Generation: rc.Generation, Error: err.Error()}); rerr != nil {
slog.Warn("agent: reporting status failed", "err", rerr)
}
}
return err
}
if raw == nil {
slog.Info("agent: applied cached config", "generation", rc.Generation, "rules", len(cfg.Rules))
return nil
}
if err := a.confirm(ctx, rc.Generation); err != nil {
if keep != nil && ctx.Err() == nil {
return a.revertGeneration(ctx, rc.Generation, StatusReverted, err, revert)
}
// Shutdown is not a verdict on the generation: keep it.
if keep != nil {
if kerr := keep(); kerr != nil {
slog.Warn("agent: dropping snapshot failed", "err", kerr)
}
}
return err
}
if keep != nil {
if err := keep(); err != nil {
return fmt.Errorf("dropping snapshot: %w", err)
}
} }
slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules))
if report { if err := a.Cache.Write(raw); err != nil {
if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { slog.Warn("agent: caching config failed", "err", err)
slog.Warn("agent: reporting status failed", "err", err) }
a.lastReverted = nil
if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) {
slog.Warn("agent: clearing reverted generation failed", "err", err)
} }
// Report the FIB so the control plane can scope router enforcement. // Report the FIB so the control plane can scope router enforcement.
if fib := CollectFIB(ctx); len(fib) > 0 { if fib := CollectFIB(ctx); len(fib) > 0 {
@@ -97,31 +177,153 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
slog.Warn("agent: reporting routes failed", "err", err) slog.Warn("agent: reporting routes failed", "err", err)
} }
} }
return nil
}
// revertGeneration marks generation as reverted before restoring, so a failed
// restore can never lead to re-applying it, then reports status. It is also kept
// in memory in case persisting fails. A failed
// restore leaves the snapshot and timer armed and is reported as failed.
func (a *Agent) revertGeneration(ctx context.Context, generation int64, status string, cause error, revert func() error) error {
rv := &reverted{Generation: generation, Status: status, Error: cause.Error()}
a.lastReverted = rv
if err := a.writeReverted(rv); err != nil {
slog.Error("agent: persisting reverted generation failed", "err", err)
}
suffix := ""
if rerr := revert(); rerr != nil {
suffix = fmt.Sprintf("; restore: %v; revert timer pending", rerr)
rv.Status = StatusFailed
rv.Error += suffix
if werr := a.writeReverted(rv); werr != nil {
slog.Error("agent: persisting reverted generation failed", "err", werr)
}
}
a.reportReverted(ctx, rv)
return fmt.Errorf("generation %d %s: %w%s", generation, rv.Status, cause, suffix)
}
var (
verifyAttempts = 3
verifyDelay = 2 * time.Second
verifyTimeout = 5 * time.Second
)
// errUnreachable means the API could not be reached through the new ruleset.
var errUnreachable = errors.New("control plane unreachable after apply")
// confirm reports generation as applied over a fresh connection, which proves
// the API is reachable through the new ruleset. Any HTTP response counts as
// reachable; only repeated transport failures return errUnreachable.
func (a *Agent) confirm(ctx context.Context, generation int64) error {
var err error
for i := 0; i < verifyAttempts; i++ {
if i > 0 {
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(verifyDelay):
}
}
actx, cancel := context.WithTimeout(ctx, verifyTimeout)
err = a.Client.ReportStatus(actx, Status{Status: StatusApplied, Generation: generation})
cancel()
if ctx.Err() != nil {
return ctx.Err()
}
var uerr *url.Error
if !errors.As(err, &uerr) {
if err != nil {
slog.Warn("agent: reporting status failed", "err", err)
} }
return nil return nil
}
}
return fmt.Errorf("%w: %v", errUnreachable, err)
}
// reverted is a generation rolled back after a failed apply or for severing the
// API; persisted so it is not re-applied until a newer generation arrives.
type reverted struct {
Generation int64 `json:"generation"`
Status string `json:"status,omitempty"`
Error string `json:"error,omitempty"`
Reported bool `json:"reported"`
}
func (a *Agent) revertedPath() string {
return filepath.Join(filepath.Dir(a.Cache.Path), "reverted.json")
}
func (a *Agent) readReverted() (*reverted, error) {
b, err := os.ReadFile(a.revertedPath())
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
var rv reverted
if err := json.Unmarshal(b, &rv); err != nil {
return nil, fmt.Errorf("parsing %s: %w", a.revertedPath(), err)
}
return &rv, nil
}
func (a *Agent) writeReverted(rv *reverted) error {
b, err := json.Marshal(rv)
if err != nil {
return err
}
return tryapply.WriteFile(a.revertedPath(), b)
}
// reportReverted reports rv until the control plane accepts it.
func (a *Agent) reportReverted(ctx context.Context, rv *reverted) {
if rv.Reported {
return
}
status := rv.Status
if status == "" {
status = StatusReverted
}
if err := a.Client.ReportStatus(ctx, Status{Status: status, Generation: rv.Generation, Error: rv.Error}); err != nil {
slog.Warn("agent: reporting reverted generation failed, retrying next cycle", "generation", rv.Generation, "err", err)
return
}
rv.Reported = true
if err := a.writeReverted(rv); err != nil {
slog.Warn("agent: persisting reverted generation failed", "err", err)
}
} }
// EngineApplier applies via the real nftables differential engine. // EngineApplier applies via the real nftables differential engine.
type EngineApplier struct{} type EngineApplier struct{}
// Apply computes and applies the differential change set for cfg. It refuses // Apply computes and applies the differential change set for cfg, with safe
// while a 'tomswall try' awaits confirmation. // under a pending try as 'tomswall try' does. The caller holds the try lock.
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { func (EngineApplier) Apply(_ context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error) {
unlock, err := tryapply.Acquire()
if err != nil {
return err
}
defer unlock()
engine, err := nftables.NewEngine(cfg) engine, err := nftables.NewEngine(cfg)
if err != nil { if err != nil {
return fmt.Errorf("initializing nftables: %w", err) return nil, nil, fmt.Errorf("initializing nftables: %w", err)
} }
changes, err := engine.Plan() changes, err := engine.Plan()
if err != nil { if err != nil {
return fmt.Errorf("computing changes: %w", err) return nil, nil, fmt.Errorf("computing changes: %w", err)
} }
if changes.Empty() { if changes.Empty() {
return nil return nil, nil, nil
} }
return engine.Apply(changes) if !safe {
return nil, nil, engine.Apply(changes)
}
snap, err := engine.Snapshot()
if err != nil {
return nil, nil, fmt.Errorf("snapshotting ruleset: %w", err)
}
// PID 0: 'tomswall confirm' must not signal the agent.
if _, err := tryapply.Arm(snap, 0, revertDelay); err != nil {
return nil, nil, err
}
return tryapply.Abort, tryapply.Discard, engine.Apply(changes)
} }
+2 -2
View File
@@ -132,10 +132,10 @@ type fakeApplier struct {
lastGen int lastGen int
} }
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error { func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config, _ bool) (func() error, func() error, error) {
atomic.AddInt32(&f.count, 1) atomic.AddInt32(&f.count, 1)
f.lastGen = len(cfg.Rules) f.lastGen = len(cfg.Rules)
return nil return nil, nil, nil
} }
const renderedYAML = `generation: 7 const renderedYAML = `generation: 7
+4 -10
View File
@@ -2,7 +2,8 @@ package agent
import ( import (
"os" "os"
"path/filepath"
"git.unkin.net/unkin/tomswall/internal/tryapply"
) )
// Cache persists the last known-good rendered config to disk so the agent can // Cache persists the last known-good rendered config to disk so the agent can
@@ -11,16 +12,9 @@ type Cache struct {
Path string Path string
} }
// Write atomically stores the raw config bytes. // Write durably stores the raw config bytes.
func (c Cache) Write(raw []byte) error { func (c Cache) Write(raw []byte) error {
if err := os.MkdirAll(filepath.Dir(c.Path), 0o755); err != nil { return tryapply.WriteFile(c.Path, raw)
return err
}
tmp := c.Path + ".tmp"
if err := os.WriteFile(tmp, raw, 0o600); err != nil {
return err
}
return os.Rename(tmp, c.Path)
} }
// Read returns the cached config, or (nil, nil) when no cache exists yet. // Read returns the cached config, or (nil, nil) when no cache exists yet.
+26 -4
View File
@@ -26,7 +26,9 @@ func NewClient(baseURL, device, token string) *Client {
BaseURL: baseURL, BaseURL: baseURL,
Device: device, Device: device,
Token: token, Token: token,
HTTP: &http.Client{Timeout: 30 * time.Second}, // No keep-alives: every request, the post-apply check included, opens a
// fresh connection that must pass the current ruleset.
HTTP: &http.Client{Timeout: 30 * time.Second, Transport: noKeepAlive()},
} }
} }
@@ -94,10 +96,24 @@ func (c *Client) ReportRoutes(ctx context.Context, prefixes []string) error {
return nil return nil
} }
// ReportStatus tells the control plane which generation this device has applied. // Status values reported to POST /api/v1/devices/{name}/status.
func (c *Client) ReportStatus(ctx context.Context, generation int64) error { const (
StatusApplied = "applied"
StatusReverted = "reverted"
StatusFailed = "failed"
)
// Status is the outcome of applying one generation.
type Status struct {
Status string `json:"status"`
Generation int64 `json:"generation"`
Error string `json:"error,omitempty"`
}
// ReportStatus tells the control plane the outcome of applying a generation.
func (c *Client) ReportStatus(ctx context.Context, st Status) error {
url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device) url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device)
payload, _ := json.Marshal(map[string]int64{"generation": generation}) payload, _ := json.Marshal(st)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload)) req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
if err != nil { if err != nil {
return err return err
@@ -116,3 +132,9 @@ func (c *Client) ReportStatus(ctx context.Context, generation int64) error {
} }
return nil return nil
} }
func noKeepAlive() http.RoundTripper {
t := http.DefaultTransport.(*http.Transport).Clone()
t.DisableKeepAlives = true
return t
}
+396
View File
@@ -0,0 +1,396 @@
package agent
import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/tryapply"
)
func TestMain(m *testing.M) {
dir, err := os.MkdirTemp("", "tomswall-agent-test")
if err != nil {
panic(err)
}
tryapply.Dir = dir
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
return nil
}
verifyDelay = time.Millisecond
verifyTimeout = time.Second
code := m.Run()
os.RemoveAll(dir)
os.Exit(code)
}
// fakeAPI serves a config generation and records status reports; while cut it
// drops connections to the status endpoint, as a severing ruleset would.
type fakeAPI struct {
*httptest.Server
gen atomic.Int64
cut atomic.Bool
code atomic.Int32
mu sync.Mutex
reports []Status
}
func newFakeAPI(t *testing.T, gen int64) *fakeAPI {
f := &fakeAPI{}
f.gen.Store(gen)
f.code.Store(http.StatusNoContent)
f.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v1/devices/fw-a/config":
_, _ = w.Write([]byte(strings.Replace(renderedYAML, "generation: 7", "generation: "+itoa(f.gen.Load()), 1)))
case "/api/v1/devices/fw-a/status":
if f.cut.Load() {
conn, _, _ := w.(http.Hijacker).Hijack()
conn.Close()
return
}
var st Status
_ = json.NewDecoder(r.Body).Decode(&st)
f.mu.Lock()
f.reports = append(f.reports, st)
f.mu.Unlock()
w.WriteHeader(int(f.code.Load()))
default:
w.WriteHeader(http.StatusNotFound)
}
}))
t.Cleanup(f.Close)
return f
}
// timerCmds records the systemd commands tryapply runs.
var timerCmds []string
func itoa(n int64) string { b, _ := json.Marshal(n); return string(b) }
func (f *fakeAPI) last() Status {
f.mu.Lock()
defer f.mu.Unlock()
if len(f.reports) == 0 {
return Status{}
}
return f.reports[len(f.reports)-1]
}
// fakeEngine always changes the ruleset, when safe under a real tryapply pending
// try; onApply simulates its effect and restoreErr fails the restore.
type fakeEngine struct {
applies, plain, restores int
err, restoreErr error
onApply func()
onRestore func()
}
func (f *fakeEngine) Apply(_ context.Context, _ *config.Config, safe bool) (func() error, func() error, error) {
if !safe {
f.plain++
return nil, nil, f.err
}
if _, err := tryapply.Arm(&nftables.Snapshot{Table: "tomswall"}, 0, time.Minute); err != nil {
return nil, nil, err
}
tryapply.Restore = func(*nftables.Snapshot) error {
f.restores++
if f.onRestore != nil {
f.onRestore()
}
return f.restoreErr
}
f.applies++
if f.onApply != nil {
f.onApply()
}
return tryapply.Abort, tryapply.Discard, f.err
}
// pending reports whether a snapshot is still armed and its timer not stopped since.
func pending(t *testing.T) bool {
t.Helper()
_, err := os.Stat(filepath.Join(tryapply.Dir, "try-snapshot.json"))
armed := len(timerCmds) > 0 && timerCmds[len(timerCmds)-1] == "systemd-run"
if (err == nil) != armed {
t.Fatalf("snapshot present=%v but timer armed=%v", err == nil, armed)
}
return armed
}
func newAgent(t *testing.T, api *fakeAPI, eng *fakeEngine) *Agent {
return &Agent{
Client: NewClient(api.URL, "fw-a", "tok"),
Cache: Cache{Path: filepath.Join(t.TempDir(), "rendered.yaml")},
Applier: eng,
}
}
func cachedGen(t *testing.T, a *Agent) int64 {
rc, err := a.Cache.Read()
if err != nil {
t.Fatal(err)
}
if rc == nil {
return 0
}
return rc.Generation
}
func TestSafeApplyReachableApplies(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.restores != 0 || api.last() != (Status{Status: StatusApplied, Generation: 7}) || cachedGen(t, a) != 7 || pending(t) {
t.Fatalf("restores=%d last=%+v cache=%d", eng.restores, api.last(), cachedGen(t, a))
}
}
func TestSafeApplyUnreachableRevertsAndReports(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, onRestore: func() { api.cut.Store(false) }}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
t.Fatalf("want errUnreachable, got %v", err)
}
if eng.restores != 1 || cachedGen(t, a) != 0 || pending(t) {
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
}
if st := api.last(); st.Status != StatusReverted || st.Generation != 7 || st.Error == "" {
t.Fatalf("last report %+v", st)
}
rv, _ := a.readReverted()
if rv == nil || rv.Generation != 7 || !rv.Reported {
t.Fatalf("persisted %+v", rv)
}
}
func TestSafeApplyRevertReportedOnceReachable(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }}
a := newAgent(t, api, eng)
_ = a.RunOnce(context.Background())
if eng.restores != 1 || api.last().Status != "" {
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
}
api.cut.Store(false)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 1 || api.last() != (Status{Status: StatusReverted, Generation: 7, Error: api.last().Error}) {
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
}
}
func TestSafeApplyShutdownDoesNotRevert(t *testing.T) {
api := newFakeAPI(t, 7)
ctx, cancel := context.WithCancel(context.Background())
eng := &fakeEngine{onApply: func() { api.cut.Store(true); cancel() }}
a := newAgent(t, api, eng)
if err := a.RunOnce(ctx); !errors.Is(err, context.Canceled) {
t.Fatalf("want context.Canceled, got %v", err)
}
if rv, _ := a.readReverted(); eng.restores != 0 || rv != nil || cachedGen(t, a) != 0 {
t.Fatalf("restores=%d reverted=%+v cache=%d", eng.restores, rv, cachedGen(t, a))
}
}
func TestSafeApplyHTTPErrorDoesNotRevert(t *testing.T) {
api := newFakeAPI(t, 7)
api.code.Store(http.StatusInternalServerError)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.restores != 0 || cachedGen(t, a) != 7 {
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
}
}
func TestSafeApplyApplyErrorRestoresAndReportsFailed(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{err: errors.New("netlink: boom")}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil {
t.Fatal("want error")
}
if st := api.last(); eng.restores != 1 || st.Status != StatusFailed || !strings.Contains(st.Error, "boom") || pending(t) {
t.Fatalf("restores=%d last=%+v", eng.restores, st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported {
t.Fatalf("persisted %+v", rv)
}
}
func TestSafeApplyApplyErrorRestoreFailsKeepsTimer(t *testing.T) {
t.Cleanup(func() { _ = tryapply.Discard() })
api := newFakeAPI(t, 7)
eng := &fakeEngine{err: errors.New("netlink: boom"), restoreErr: errors.New("netlink: stuck")}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "restore: restoring snapshot: netlink: stuck") {
t.Fatalf("got %v", err)
}
want := "apply: netlink: boom; restore: restoring snapshot: netlink: stuck; revert timer pending"
if st := api.last(); st != (Status{Status: StatusFailed, Generation: 7, Error: want}) || !pending(t) {
t.Fatalf("last=%+v", st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || rv.Status != StatusFailed {
t.Fatalf("persisted %+v", rv)
}
// The next cycle waits for the timer instead of re-applying.
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
t.Fatalf("err=%v applies=%d", err, eng.applies)
}
}
func TestSafeApplyUnreachableRestoreFailsKeepsTimer(t *testing.T) {
t.Cleanup(func() { _ = tryapply.Discard() })
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, restoreErr: errors.New("netlink: stuck")}
eng.onRestore = func() { api.cut.Store(false) }
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) || !strings.Contains(err.Error(), "revert timer pending") {
t.Fatalf("got %v", err)
}
st := api.last()
if st.Status != StatusFailed || st.Generation != 7 || !strings.HasPrefix(st.Error, errUnreachable.Error()) ||
!strings.HasSuffix(st.Error, "; restore: restoring snapshot: netlink: stuck; revert timer pending") || !pending(t) {
t.Fatalf("last=%+v", st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported || cachedGen(t, a) != 0 {
t.Fatalf("persisted %+v cache=%d", rv, cachedGen(t, a))
}
}
func TestSafeApplyRevertedGenerationSkippedAfterRestart(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.writeReverted(&reverted{Generation: 7, Reported: true}); err != nil {
t.Fatal(err)
}
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 0 {
t.Fatalf("reverted generation re-applied")
}
api.gen.Store(8)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if rv, _ := a.readReverted(); eng.applies != 1 || api.last().Generation != 8 || rv != nil {
t.Fatalf("applies=%d last=%+v reverted=%+v", eng.applies, api.last(), rv)
}
}
func TestSafeApplySkipsWhileTryPending(t *testing.T) {
marker := filepath.Join(tryapply.Dir, "try-snapshot.json")
if err := os.WriteFile(marker, []byte("{}"), 0o600); err != nil {
t.Fatal(err)
}
defer os.Remove(marker)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 0 || api.last().Status != "" {
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
}
}
// failArm makes arming the revert timer fail, as without systemd.
func failArm(t *testing.T) {
orig := tryapply.Run
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
if name == "systemd-run" {
return errors.New("no systemd")
}
return nil
}
t.Cleanup(func() { tryapply.Run = orig })
}
func TestSafeApplyArmFailureReportsFailedAndRetries(t *testing.T) {
failArm(t)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "no systemd") {
t.Fatalf("got %v", err)
}
if st := api.last(); eng.applies != 0 || st.Status != StatusFailed || st.Generation != 7 || pending(t) || cachedGen(t, a) != 0 {
t.Fatalf("applies=%d last=%+v", eng.applies, st)
}
if rv, _ := a.readReverted(); rv != nil || a.lastReverted != nil {
t.Fatalf("arm failure marked generation reverted: %+v", rv)
}
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
return nil
}
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 || cachedGen(t, a) != 7 {
t.Fatalf("retry err=%v applies=%d", err, eng.applies)
}
}
func TestCachedConfigAppliesWithoutArm(t *testing.T) {
failArm(t)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.Cache.Write([]byte(renderedYAML)); err != nil {
t.Fatal(err)
}
api.Close()
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.plain != 1 || eng.applies != 0 || pending(t) {
t.Fatalf("plain=%d safe=%d", eng.plain, eng.applies)
}
}
func TestSafeApplyRevertedKeptInMemoryWhenPersistFails(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onRestore: func() { api.cut.Store(false) }}
a := newAgent(t, api, eng)
// A non-empty directory in its place makes persisting reverted.json fail.
eng.onApply = func() {
api.cut.Store(true)
_ = os.MkdirAll(filepath.Join(a.revertedPath(), "x"), 0o755)
}
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
t.Fatalf("want errUnreachable, got %v", err)
}
if err := os.RemoveAll(a.revertedPath()); err != nil {
t.Fatal(err)
}
if eng.restores != 1 || api.last().Status != StatusReverted || pending(t) {
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
}
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
t.Fatalf("reverted generation re-applied: err=%v applies=%d", err, eng.applies)
}
}
+23
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"regexp"
"strings" "strings"
"gopkg.in/yaml.v3" "gopkg.in/yaml.v3"
@@ -55,10 +56,17 @@ type Settings struct {
AddressFamily AddressFamily `yaml:"address_family,omitempty"` AddressFamily AddressFamily `yaml:"address_family,omitempty"`
IPForwarding bool `yaml:"ip_forwarding"` IPForwarding bool `yaml:"ip_forwarding"`
LogLevel string `yaml:"log_level"` LogLevel string `yaml:"log_level"`
// LogLimit rate-limits every log site, shorewall LOGLIMIT syntax rate/unit[:burst]; unset logs every hit.
LogLimit string `yaml:"log_limit,omitempty"`
TableName string `yaml:"table_name"` TableName string `yaml:"table_name"`
// When true, auto-generate CONTINUE policies for sub-zones to their parent zones. // When true, auto-generate CONTINUE policies for sub-zones to their parent zones.
ImplicitContinue bool `yaml:"implicit_continue,omitempty"` ImplicitContinue bool `yaml:"implicit_continue,omitempty"`
// Verdict for ct state invalid/untracked packets; continue passes them to the rules.
// Unset: invalid drops, untracked continues.
InvalidDisposition PolicyAction `yaml:"invalid_disposition,omitempty"`
UntrackedDisposition PolicyAction `yaml:"untracked_disposition,omitempty"`
} }
// Load reads a config file in YAML or JSON format (detected by extension). // Load reads a config file in YAML or JSON format (detected by extension).
@@ -107,6 +115,8 @@ func (c *Config) applyDefaults() {
} }
} }
var logLimitRe = regexp.MustCompile(`^[1-9][0-9]*/(sec|second|min|minute|hour|day)(:[1-9][0-9]*)?$`)
var validAddressFamilies = map[AddressFamily]bool{ var validAddressFamilies = map[AddressFamily]bool{
FamilyINET: true, FamilyIP: true, FamilyIP6: true, FamilyINET: true, FamilyIP: true, FamilyIP6: true,
} }
@@ -115,6 +125,19 @@ func (c *Config) validateSettings() error {
if !validAddressFamilies[c.Settings.AddressFamily] { if !validAddressFamilies[c.Settings.AddressFamily] {
return fmt.Errorf("unknown address_family %q (use inet, ip, or ip6)", c.Settings.AddressFamily) return fmt.Errorf("unknown address_family %q (use inet, ip, or ip6)", c.Settings.AddressFamily)
} }
if l := c.Settings.LogLimit; l != "" && !logLimitRe.MatchString(l) {
return fmt.Errorf("invalid log_limit %q (use rate/{sec|min|hour|day}[:burst]; per-source s:/d: is not supported)", l)
}
for name, d := range map[string]PolicyAction{
"invalid_disposition": c.Settings.InvalidDisposition,
"untracked_disposition": c.Settings.UntrackedDisposition,
} {
switch d {
case "", PolicyAccept, PolicyDrop, PolicyReject, PolicyContinue:
default:
return fmt.Errorf("unknown %s %q (use accept, drop, reject, or continue)", name, d)
}
}
return nil return nil
} }
+50
View File
@@ -473,6 +473,20 @@ func TestValidateHosts(t *testing.T) {
}, },
wantErr: "interface \"eth99\" not defined in interfaces", wantErr: "interface \"eth99\" not defined in interfaces",
}, },
{
name: "host interface matched by wildcard",
zones: map[string]Zone{
"fw": {Type: ZoneFirewall},
"net": {Type: ZoneIP},
"lan": {Type: ZoneIP, Parents: []string{"net"}},
},
interfaces: []Interface{
{Zone: "net", Interface: "enp+"},
},
hosts: []Host{
{Zone: "lan", Interface: "enp2s0", Addresses: []string{"192.0.2.0/24"}},
},
},
{ {
name: "zone not defined", name: "zone not defined",
zones: map[string]Zone{ zones: map[string]Zone{
@@ -531,6 +545,13 @@ func TestValidateHosts(t *testing.T) {
}, },
wantErr: "interface required", wantErr: "interface required",
}, },
{
name: "invalid exclusion",
zones: map[string]Zone{"fw": {Type: ZoneFirewall}, "net": {Type: ZoneIP}, "loc": {Type: ZoneIP}},
interfaces: []Interface{{Zone: "net", Interface: "eth0"}},
hosts: []Host{{Zone: "loc", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}, Exclusions: []string{"192.0.2.0/24!192.0.2.7"}}},
wantErr: "invalid address",
},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -1052,3 +1073,32 @@ func TestSplitZoneList(t *testing.T) {
} }
} }
} }
func TestValidateDispositions(t *testing.T) {
for _, tc := range []struct {
invalid, untracked PolicyAction
wantErr string
}{
{"", "", ""},
{PolicyContinue, PolicyDrop, ""},
{"bogus", "", `unknown invalid_disposition "bogus"`},
{PolicyAccept, "log", `unknown untracked_disposition "log"`},
} {
c := baseConfig()
c.Settings.InvalidDisposition = tc.invalid
c.Settings.UntrackedDisposition = tc.untracked
checkErr(t, c.Validate(), tc.wantErr)
}
}
func TestValidateLogLimit(t *testing.T) {
for v, ok := range map[string]bool{
"": true, "1/sec": true, "1/sec:10": true, "30/minute:5": true, "2/hour": true, "1/day:1": true,
"s:1/sec:10": false, "d:1/sec": false, "1": false, "1/week": false, "0/sec": false, "1/sec:": false,
} {
c := &Config{Settings: Settings{AddressFamily: FamilyINET, LogLimit: v}}
if err := c.validateSettings(); (err == nil) != ok {
t.Errorf("log_limit %q: err = %v, want ok=%v", v, err, ok)
}
}
}
+15 -2
View File
@@ -1,6 +1,11 @@
package config package config
import "fmt" import (
"fmt"
"net/netip"
"slices"
"strings"
)
type Host struct { type Host struct {
Zone string `yaml:"zone"` Zone string `yaml:"zone"`
@@ -40,7 +45,8 @@ func (c *Config) validateHosts() error {
ifaceFound := false ifaceFound := false
for _, iface := range c.Interfaces { for _, iface := range c.Interfaces {
if iface.Interface == h.Interface || iface.PhysicalName() == h.Interface { prefix, wild := strings.CutSuffix(iface.PhysicalName(), "+")
if iface.Interface == h.Interface || iface.PhysicalName() == h.Interface || (wild && strings.HasPrefix(h.Interface, prefix)) {
ifaceFound = true ifaceFound = true
break break
} }
@@ -52,6 +58,13 @@ func (c *Config) validateHosts() error {
if !h.Dynamic && len(h.Addresses) == 0 { if !h.Dynamic && len(h.Addresses) == 0 {
return fmt.Errorf("host[%d]: at least one address required (or set dynamic: true)", i) return fmt.Errorf("host[%d]: at least one address required (or set dynamic: true)", i)
} }
for _, a := range slices.Concat(h.Addresses, h.Exclusions) {
if _, err := netip.ParsePrefix(a); err != nil {
if _, err := netip.ParseAddr(a); err != nil {
return fmt.Errorf("host[%d]: invalid address %q", i, a)
}
}
}
} }
return nil return nil
} }
+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)
} }
+353 -62
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"net" "net"
"net/netip"
"slices" "slices"
"sort" "sort"
"strconv" "strconv"
@@ -65,26 +66,68 @@ func (c *Compiler) Compile() (*FirewallState, error) {
return nil, fmt.Errorf("static-nat: %w", err) return nil, fmt.Errorf("static-nat: %w", err)
} }
c.compileMSSClamp(state) c.compileMSSClamp(state)
limitLogs(state, c.cfg.Settings.LogLimit)
return state, nil return state, nil
} }
// limitLogs puts a limit in front of every log expression. A limit stops the
// whole rule, so like shorewall's separate LOG rule, a log followed by an action
// splits into a limited log-only rule and the same rule without the log. A
// LOG-action rule (nothing but rate limits after the log) keeps only the log rule.
func limitLogs(state *FirewallState, spec string) {
if spec == "" {
return
}
for chain, rules := range state.Rules {
var out []ManagedRule
for _, r := range rules {
i := slices.IndexFunc(r.Exprs, func(e expr.Any) bool { _, ok := e.(*expr.Log); return ok })
if i < 0 {
out = append(out, r)
continue
}
logRule := r
logRule.Exprs = slices.Concat(r.Exprs[:i], parseRateLimit(spec), r.Exprs[i:i+1])
out = append(out, logRule)
if slices.ContainsFunc(r.Exprs[i+1:], func(e expr.Any) bool { _, ok := e.(*expr.Limit); return !ok }) {
r.Exprs = slices.Concat(r.Exprs[:i], r.Exprs[i+1:])
out = append(out, r)
}
}
state.Rules[chain] = out
}
}
func (c *Compiler) compileConntrackFastPath(state *FirewallState) error { func (c *Compiler) compileConntrackFastPath(state *FirewallState) error {
invalid := c.cfg.Settings.InvalidDisposition
if invalid == "" {
invalid = config.PolicyDrop
}
for _, chain := range []string{"input", "forward", "output"} { for _, chain := range []string{"input", "forward", "output"} {
state.Rules[chain] = append(state.Rules[chain], state.Rules[chain] = append(state.Rules[chain], ManagedRule{
ManagedRule{
Chain: chain, Chain: chain,
Exprs: append(matchCtState(ctStateEstablished|ctStateRelated), Exprs: append(matchCtState(ctStateEstablished|ctStateRelated),
&expr.Verdict{Kind: expr.VerdictAccept}), &expr.Verdict{Kind: expr.VerdictAccept}),
Tag: "ct:fastpath:" + chain, Tag: "ct:fastpath:" + chain,
}, })
ManagedRule{ for _, d := range []struct {
name string
state uint32
action config.PolicyAction
}{
{"invalid", ctStateInvalid, invalid},
{"untracked", ctStateUntracked, c.cfg.Settings.UntrackedDisposition},
} {
if d.action == "" || d.action == config.PolicyContinue {
continue
}
state.Rules[chain] = append(state.Rules[chain], ManagedRule{
Chain: chain, Chain: chain,
Exprs: append(matchCtState(ctStateInvalid), Exprs: append(matchCtState(d.state), policyVerdict(d.action, c.cfg.Settings.AddressFamily)...),
&expr.Verdict{Kind: expr.VerdictDrop}), Tag: "ct:" + d.name + ":" + chain,
Tag: "ct:invalid:" + chain, })
}, }
)
} }
return nil return nil
} }
@@ -303,12 +346,12 @@ func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string,
(dstAddr == "" || strings.HasPrefix(dstAddr, "!")) { (dstAddr == "" || strings.HasPrefix(dstAddr, "!")) {
return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone) return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone)
} }
srcIfaces, dstIfaces := c.resolveZoneInterfaces(srcZone, srcAddr), []string{""} srcIfaces, dstIfaces := c.resolveZone(srcZone, srcAddr), []zoneMatch{{}}
if chain == "raw_prerouting" && c.resolveZoneInterfaces(dstZone, dstAddr) == nil { if chain == "raw_prerouting" && c.resolveZone(dstZone, dstAddr) == nil {
return nil return nil
} }
if chain == "raw_output" { if chain == "raw_output" {
srcIfaces, dstIfaces = []string{""}, c.resolveZoneInterfaces(dstZone, dstAddr) srcIfaces, dstIfaces = []zoneMatch{{}}, c.resolveZone(dstZone, dstAddr)
} }
out := chain out := chain
if ct.Action == config.ConntrackHelper { if ct.Action == config.ConntrackHelper {
@@ -401,24 +444,24 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
break break
} }
var extra []expr.Any var match, extra []expr.Any
var replaceVerdict []expr.Any var replaceVerdict []expr.Any
if rule.User != "" { if rule.User != "" {
extra = append(extra, matchUID(rule.User)...) match = append(match, matchUID(rule.User)...)
} }
if rule.Mark != "" { if rule.Mark != "" {
extra = append(extra, matchMark(rule.Mark)...) match = append(match, matchMark(rule.Mark)...)
}
if rule.ConnLimit != "" {
match = append(match, matchConnLimit(rule.ConnLimit)...)
}
if rule.Time != nil {
match = append(match, matchTime(rule.Time)...)
} }
if rule.RateLimit != "" { if rule.RateLimit != "" {
extra = append(extra, parseRateLimit(rule.RateLimit)...) extra = append(extra, parseRateLimit(rule.RateLimit)...)
} }
if rule.ConnLimit != "" {
extra = append(extra, matchConnLimit(rule.ConnLimit)...)
}
if rule.Time != nil {
extra = append(extra, matchTime(rule.Time)...)
}
if rule.SetMark != "" { if rule.SetMark != "" {
extra = append(extra, setMarkExprs(rule.SetMark)...) extra = append(extra, setMarkExprs(rule.SetMark)...)
} }
@@ -426,21 +469,27 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul
replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}} replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}}
} }
if len(extra) > 0 || len(replaceVerdict) > 0 { if len(match) > 0 || len(extra) > 0 || len(replaceVerdict) > 0 {
existingExprs := rules[idx].Exprs var pre, post, verdict []expr.Any
var verdict []expr.Any for _, e := range rules[idx].Exprs {
var nonVerdict []expr.Any switch e.(type) {
for _, e := range existingExprs { case *expr.Verdict:
if _, ok := e.(*expr.Verdict); ok {
verdict = append(verdict, e) verdict = append(verdict, e)
case *expr.Log:
post = append(post, e)
default:
if len(post) > 0 {
post = append(post, e)
} else { } else {
nonVerdict = append(nonVerdict, e) pre = append(pre, e)
}
} }
} }
if len(replaceVerdict) > 0 { if len(replaceVerdict) > 0 {
verdict = replaceVerdict verdict = replaceVerdict
} }
rules[idx].Exprs = append(append(nonVerdict, extra...), verdict...) // matches go before the log so it only fires for packets the rule matches
rules[idx].Exprs = slices.Concat(pre, match, post, extra, verdict)
} }
} }
state.Rules[chain] = rules state.Rules[chain] = rules
@@ -626,8 +675,8 @@ func splitAddrs(addr string) []string {
func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, origDest, proto string, func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, origDest, proto string,
dports, sports config.PortSpec, action config.RuleAction, logLevel string, dports, sports config.PortSpec, action config.RuleAction, logLevel string,
fwZone string, section config.RuleSection) error { fwZone string, section config.RuleSection) error {
srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr) srcIfaces := c.resolveZone(srcZone, srcAddr)
dstIfaces := c.resolveZoneInterfaces(dstZone, dstAddr) dstIfaces := c.resolveZone(dstZone, dstAddr)
chain := c.selectChain(srcZone, dstZone, fwZone) chain := c.selectChain(srcZone, dstZone, fwZone)
// ponytail: forward daddr is post-DNAT; lift with `ct original daddr` (expr.Ct Direction, google/nftables v0.3.0). // ponytail: forward daddr is post-DNAT; lift with `ct original daddr` (expr.Ct Direction, google/nftables v0.3.0).
if origDest != "" && chain == "forward" { if origDest != "" && chain == "forward" {
@@ -696,7 +745,7 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
dnatPort = uint16(p) dnatPort = uint16(p)
} }
srcIfaces := c.resolveZoneInterfaces(srcZone, srcAddr) srcIfaces := c.resolveZone(srcZone, srcAddr)
var odExprs []expr.Any var odExprs []expr.Any
if origDest != "" { if origDest != "" {
@@ -717,12 +766,12 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
} }
for _, srcIface := range srcIfaces { for _, srcIface := range srcIfaces {
for _, m := range matches { zm, err := zoneMatchExprs(srcIface, true)
var exprs []expr.Any if err != nil {
return err
if srcIface != "" {
exprs = append(exprs, matchIfaceName(true, srcIface)...)
} }
for _, m := range matches {
exprs := slices.Clone(zm)
if srcAddr != "" { if srcAddr != "" {
src, err := matchSourceCIDR(srcAddr) src, err := matchSourceCIDR(srcAddr)
@@ -809,31 +858,35 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
func (c *Compiler) compilePolicies(state *FirewallState) error { func (c *Compiler) compilePolicies(state *FirewallState) error {
fwZone := c.cfg.FirewallZone() fwZone := c.cfg.FirewallZone()
overridden := map[string]bool{}
for i, pol := range c.cfg.Policy { for i, pol := range c.cfg.Policy {
tag := fmt.Sprintf("policy:%d", i) tag := fmt.Sprintf("policy:%d", i)
explicitIntra := pol.Source == pol.Dest && !isGlobalZone(pol.Source)
srcZones := c.expandZoneRef(pol.Source) srcZones := c.expandZoneRef(pol.Source)
dstZones := c.expandZoneRef(pol.Dest) dstZones := c.expandZoneRef(pol.Dest)
for _, sz := range srcZones { for _, sz := range srcZones {
for _, dz := range dstZones { for _, dz := range dstZones {
if sz == dz && !strings.HasSuffix(pol.Source, "+") { if sz == dz {
if sz == fwZone || (!explicitIntra && !strings.HasSuffix(pol.Source, "+") && !strings.HasSuffix(pol.Dest, "+")) {
continue continue
} }
overridden[sz] = true
}
chain := c.selectChain(sz, dz, fwZone) chain := c.selectChain(sz, dz, fwZone)
srcIfaces := c.resolveZoneInterfaces(sz, "") srcIfaces := c.resolveZone(sz, "")
dstIfaces := c.resolveZoneInterfaces(dz, "") dstIfaces := c.resolveZone(dz, "")
for _, si := range srcIfaces { for _, si := range srcIfaces {
for _, di := range dstIfaces { for _, di := range dstIfaces {
var exprs []expr.Any if sz == dz && intraZoneSkip(si, di) {
continue
if si != "" {
exprs = append(exprs, matchIfaceName(true, si)...)
} }
if di != "" && chain != "input" { exprs, err := zonePairExprs(si, di, chain)
exprs = append(exprs, matchIfaceName(false, di)...) if err != nil {
return fmt.Errorf("policy[%d]: %w", i, err)
} }
if pol.RateLimit != "" { if pol.RateLimit != "" {
@@ -861,9 +914,66 @@ func (c *Compiler) compilePolicies(state *FirewallState) error {
} }
} }
return c.compileImplicitIntraZone(state, overridden)
}
// compileImplicitIntraZone accepts traffic between different interfaces of one zone, shorewall's implicit intra-zone ACCEPT policy.
func (c *Compiler) compileImplicitIntraZone(state *FirewallState, overridden map[string]bool) error {
fwZone := c.cfg.FirewallZone()
zones := make([]string, 0, len(c.cfg.Zones))
for z := range c.cfg.Zones {
zones = append(zones, z)
}
sort.Strings(zones)
for _, z := range zones {
if z == fwZone || overridden[z] {
continue
}
if len(c.cfg.ZoneInterfaces(z)) == 0 && !slices.ContainsFunc(c.cfg.Hosts, func(h config.Host) bool { return h.Zone == z }) {
continue
}
matches := c.resolveZone(z, "")
for _, si := range matches {
for _, di := range matches {
if intraZoneSkip(si, di) {
continue
}
exprs, err := zonePairExprs(si, di, "forward")
if err != nil {
return fmt.Errorf("zone %s: %w", z, err)
}
state.Rules["forward"] = append(state.Rules["forward"], ManagedRule{
Chain: "forward",
Exprs: append(exprs, &expr.Verdict{Kind: expr.VerdictAccept}),
Tag: "intra:" + z,
})
}
}
}
return nil return nil
} }
// intraZoneSkip drops intra-zone pairs on one interface unless both are routeback hosts entries
// (interface routeback is compileIntraZone's job), and pairs whose address families can never both match.
func intraZoneSkip(si, di zoneMatch) bool {
if si.iface != "" && si.iface == di.iface && !(si.routeback && di.routeback) {
return true
}
a, b := matchFamily(si), matchFamily(di)
return a != 0 && b != 0 && a != b
}
// matchFamily is the NFPROTO a zoneMatch is guarded by: its host address's family, else fam.
func matchFamily(m zoneMatch) byte {
if p, err := parsePrefix(m.addr); err == nil {
if p.Addr().Is4() {
return unix.NFPROTO_IPV4
}
return unix.NFPROTO_IPV6
}
return m.fam
}
func (c *Compiler) compileSNAT(state *FirewallState) error { func (c *Compiler) compileSNAT(state *FirewallState) error {
for i, snat := range c.cfg.SNAT { for i, snat := range c.cfg.SNAT {
tag := fmt.Sprintf("snat:%d", i) tag := fmt.Sprintf("snat:%d", i)
@@ -1165,11 +1275,25 @@ func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string {
return "forward" return "forward"
} }
// resolveZoneInterfaces returns nil (fail closed) for an unknown zone, or one with no interfaces unless a non-negated address match narrows the rule. // zoneMatch classifies a packet into a zone: an interface (empty: any) and, for a hosts entry, one host
func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string { // address; excl carves out hosts exclusions and the hosts of sub-zones, which shorewall matches first.
// fam (an NFPROTO; addr's family when set, else 0: any) guards addr and excl so IPv4 offsets are
// never compared against IPv6 bytes. routeback marks a hosts entry with the routeback option. notIface
// carves narrower sub-zone host interfaces out of a wildcard iface; they get entries of their own.
type zoneMatch struct {
iface, addr string
notIface []string
excl []string
fam byte
routeback bool
}
// resolveZone returns nil (fail closed) for an unknown zone, or one with neither interfaces nor hosts
// unless a non-negated address match narrows the rule.
func (c *Compiler) resolveZone(zone, addr string) []zoneMatch {
switch zone { switch zone {
case "", "all", "all+", "any", "any+": case "", "all", "all+", "any", "any+":
return []string{""} return []zoneMatch{{}}
} }
z, ok := c.cfg.Zones[zone] z, ok := c.cfg.Zones[zone]
if !ok { if !ok {
@@ -1177,24 +1301,188 @@ func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string {
return nil return nil
} }
if z.Type == config.ZoneFirewall { if z.Type == config.ZoneFirewall {
return []string{""} return []zoneMatch{{}}
} }
if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 { var out []zoneMatch
return ifaces hasHosts := false
for _, iface := range c.cfg.ZoneInterfaces(zone) {
for _, sp := range c.subZoneSplit(zone, iface) {
if len(sp.sub) == 0 {
out = append(out, zoneMatch{iface: sp.iface, notIface: sp.not})
} else {
v4, v6 := splitFamily(sp.sub)
out = append(out, zoneMatch{iface: sp.iface, notIface: sp.not, excl: v4, fam: unix.NFPROTO_IPV4},
zoneMatch{iface: sp.iface, notIface: sp.not, excl: v6, fam: unix.NFPROTO_IPV6})
}
for _, b := range sp.back {
if addrsOverlap(b, addr) {
out = append(out, zoneMatch{iface: sp.iface, notIface: sp.not, addr: b})
}
}
}
}
for _, h := range c.cfg.Hosts {
if h.Zone != zone {
continue
}
hasHosts = true
for _, sp := range c.subZoneSplit(zone, h.Interface) {
v4, v6 := splitFamily(slices.Concat(h.Exclusions, sp.sub))
for _, a := range h.Addresses {
if addrsOverlap(a, addr) {
m := zoneMatch{iface: sp.iface, notIface: sp.not, addr: a, excl: v6, routeback: h.Options.RouteBack}
if p, err := parsePrefix(a); err == nil && p.Addr().Is4() {
m.excl = v4
}
out = append(out, m)
}
}
}
}
if len(out) > 0 || hasHosts {
return out
} }
if addr != "" && !strings.HasPrefix(addr, "!") { if addr != "" && !strings.HasPrefix(addr, "!") {
return []string{""} return []zoneMatch{{}}
} }
if !c.warned[zone] { if !c.warned[zone] {
if c.warned == nil { if c.warned == nil {
c.warned = map[string]bool{} c.warned = map[string]bool{}
} }
c.warned[zone] = true c.warned[zone] = true
slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) slog.Warn("compiler: zone has no interfaces or hosts, skipping its rules", "zone", zone)
} }
return nil return nil
} }
// subZoneSplit partitions iface for zone's sub-zone hosts, as shorewall matches them: the hosts on an
// interface covering iface exclude their addresses (sub) on all of it, while a host interface strictly
// inside a wildcard iface is carved out (not) into its own entry, so its addresses stay in zone on every
// other interface the wildcard matches. back lists the sub-zone hosts' exclusions, which fall back to zone.
type subZoneSplit struct {
iface string
not []string
sub, back []string
}
func (c *Compiler) subZoneSplit(zone, iface string) []subZoneSplit {
top := subZoneSplit{iface: iface}
var inner []string
for _, h := range c.cfg.Hosts {
switch {
case !c.cfg.IsSubZone(h.Zone, zone):
case ifaceCovers(h.Interface, iface):
top.sub = append(top.sub, h.Addresses...)
top.back = append(top.back, h.Exclusions...)
case ifaceCovers(iface, h.Interface) && !slices.Contains(inner, h.Interface):
inner = append(inner, h.Interface)
}
}
var rest []subZoneSplit
for _, h := range inner {
if slices.ContainsFunc(inner, func(o string) bool { return o != h && ifaceCovers(o, h) }) {
continue
}
top.not = append(top.not, h)
rest = append(rest, c.subZoneSplit(zone, h)...)
}
return append([]subZoneSplit{top}, rest...)
}
// ifaceCovers reports whether every interface name b matches is also matched by a ("+" suffix: prefix wildcard).
func ifaceCovers(a, b string) bool {
if pa, ok := strings.CutSuffix(a, "+"); ok {
return strings.HasPrefix(strings.TrimSuffix(b, "+"), pa)
}
return a == b
}
// addrsOverlap reports whether a host address can match a rule address; unparsable or negated rule addresses keep the host.
func addrsOverlap(host, rule string) bool {
if rule == "" || strings.HasPrefix(rule, "!") {
return true
}
h, err := parsePrefix(host)
if err != nil {
return true
}
for _, r := range strings.Split(rule, ",") {
if p, err := parsePrefix(r); err != nil || p.Overlaps(h) {
return true
}
}
return false
}
// splitFamily partitions addresses by family; unparsable ones go to v6 so zoneMatchExprs still rejects them.
func splitFamily(addrs []string) (v4, v6 []string) {
for _, a := range addrs {
if p, err := parsePrefix(a); err == nil && p.Addr().Is4() {
v4 = append(v4, a)
} else {
v6 = append(v6, a)
}
}
return v4, v6
}
func parsePrefix(s string) (netip.Prefix, error) {
if a, err := netip.ParseAddr(s); err == nil {
return netip.PrefixFrom(a, a.BitLen()), nil
}
return netip.ParsePrefix(s)
}
// zoneMatchExprs matches a zone on the in (src) or out interface plus its host address and exclusions,
// all guarded by m.fam so an IPv4 address never matches IPv6 bytes in an inet table.
func zoneMatchExprs(m zoneMatch, src bool) ([]expr.Any, error) {
var out []expr.Any
if m.iface != "" {
out = matchIfaceName(src, m.iface)
}
for _, n := range m.notIface {
e := matchIfaceName(src, n)
e[1].(*expr.Cmp).Op = expr.CmpOpNeq
out = append(out, e...)
}
var addr []expr.Any
if m.addr != "" {
p, err := parsePrefix(m.addr)
if err != nil {
return nil, fmt.Errorf("invalid host address %q", m.addr)
}
if addr, err = matchAddrCIDR(m.addr, src); err != nil {
return nil, err
}
m.fam = unix.NFPROTO_IPV6
if p.Addr().Is4() {
m.fam = unix.NFPROTO_IPV4
}
}
if m.fam != 0 {
out = append(out, matchNFProto(m.fam)...)
}
out = append(out, addr...)
if len(m.excl) > 0 {
e, err := matchAddrCIDR("!"+strings.Join(m.excl, ","), src)
if err != nil {
return nil, err
}
out = append(out, e...)
}
return out, nil
}
// zonePairExprs matches the source zone inbound and, outside input, the dest zone outbound.
func zonePairExprs(src, dst zoneMatch, chain string) ([]expr.Any, error) {
out, err := zoneMatchExprs(src, true)
if err != nil || chain == "input" {
return out, err
}
d, err := zoneMatchExprs(dst, false)
return append(out, d...), err
}
func (c *Compiler) expandZoneRef(ref string) []string { func (c *Compiler) expandZoneRef(ref string) []string {
base := ref base := ref
var excluded map[string]bool var excluded map[string]bool
@@ -1224,14 +1512,10 @@ func (c *Compiler) expandZoneRef(ref string) []string {
return []string{base} return []string{base}
} }
func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) { func (c *Compiler) buildMatchExprs(srcIface, dstIface zoneMatch, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) {
var exprs []expr.Any exprs, err := zonePairExprs(srcIface, dstIface, chain)
if err != nil {
if srcIface != "" { return nil, err
exprs = append(exprs, matchIfaceName(true, srcIface)...)
}
if dstIface != "" && chain != "input" {
exprs = append(exprs, matchIfaceName(false, dstIface)...)
} }
if srcAddr != "" { if srcAddr != "" {
@@ -2049,6 +2333,13 @@ func rejectExprs(proto byte, family config.AddressFamily) []expr.Any {
Code: 0, Code: 0,
}} }}
} }
// icmpx is inet-only; ip and ip6 tables silently drop on it.
switch family {
case config.FamilyIP:
return []expr.Any{&expr.Reject{Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 3}} // port-unreachable
case config.FamilyIP6:
return []expr.Any{&expr.Reject{Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 4}} // port-unreachable
}
return []expr.Any{&expr.Reject{ return []expr.Any{&expr.Reject{
Type: unix.NFT_REJECT_ICMPX_UNREACH, Type: unix.NFT_REJECT_ICMPX_UNREACH,
Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH, Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH,
+654 -21
View File
@@ -7,6 +7,7 @@ import (
"log/slog" "log/slog"
"net" "net"
"reflect" "reflect"
"slices"
"strings" "strings"
"testing" "testing"
@@ -144,19 +145,17 @@ func TestCompiler_ResolveZoneInterfaces(t *testing.T) {
} }
c := NewCompiler(cfg) c := NewCompiler(cfg)
ifaces := c.resolveZoneInterfaces("net", "") for _, tt := range []struct {
if len(ifaces) != 1 || ifaces[0] != "eth0" { zone string
t.Errorf("resolveZoneInterfaces(net) = %v, want [eth0]", ifaces) want []zoneMatch
}{
{"net", []zoneMatch{{iface: "eth0"}}},
{"all", []zoneMatch{{}}},
{"fw", []zoneMatch{{}}},
} {
if got := c.resolveZone(tt.zone, ""); !reflect.DeepEqual(got, tt.want) {
t.Errorf("resolveZone(%s) = %v, want %v", tt.zone, got, tt.want)
} }
ifaces = c.resolveZoneInterfaces("all", "")
if len(ifaces) != 1 || ifaces[0] != "" {
t.Errorf("resolveZoneInterfaces(all) = %v, want [\"\"]", ifaces)
}
ifaces = c.resolveZoneInterfaces("fw", "")
if len(ifaces) != 1 || ifaces[0] != "" {
t.Errorf("resolveZoneInterfaces(fw) = %v, want [\"\"]", ifaces)
} }
} }
@@ -1849,13 +1848,13 @@ func TestCompile_InterfacelessZonesFailClosed(t *testing.T) {
c.Zones["hst"] = config.Zone{Type: config.ZoneIP} c.Zones["hst"] = config.Zone{Type: config.ZoneIP}
c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}} c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}}
c.Policy = []config.Policy{{Source: "fw", Dest: "hst", Action: config.PolicyAccept}} c.Policy = []config.Policy{{Source: "fw", Dest: "hst", Action: config.PolicyAccept}}
}, "policy:0", 0, []string{"hst"}}, }, "policy:0", 1, nil},
{"fw all expansion keeps zones with interfaces", func(c *config.Config) { {"fw all expansion keeps zones with interfaces", func(c *config.Config) {
ipsec(c) ipsec(c)
c.Zones["hst"] = config.Zone{Type: config.ZoneIP} c.Zones["hst"] = config.Zone{Type: config.ZoneIP}
c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}} c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}}
c.Policy = []config.Policy{{Source: "fw", Dest: "all", Action: config.PolicyDrop}} c.Policy = []config.Policy{{Source: "fw", Dest: "all", Action: config.PolicyDrop}}
}, "policy:0", 1, []string{"hst", "ips"}}, }, "policy:0", 2, []string{"ips"}},
{"negated address does not scope", rule("ips:!192.0.2.1"), "rule:0", 0, []string{"ips"}}, {"negated address does not scope", rule("ips:!192.0.2.1"), "rule:0", 0, []string{"ips"}},
{"address scopes", rule("ips:192.0.2.1"), "rule:0", 1, []string{"ips"}}, {"address scopes", rule("ips:192.0.2.1"), "rule:0", 1, []string{"ips"}},
} }
@@ -1895,6 +1894,19 @@ func TestCompile_InterfacelessZonesFailClosed(t *testing.T) {
func describeRule(r ManagedRule) string { func describeRule(r ManagedRule) string {
var parts []string var parts []string
for i, e := range r.Exprs { for i, e := range r.Exprs {
if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseNetworkHeader && i+2 < len(r.Exprs) {
bw, okb := r.Exprs[i+1].(*expr.Bitwise)
cmp, okc := r.Exprs[i+2].(*expr.Cmp)
if okb && okc {
name := map[uint32]string{12: "saddr", 16: "daddr", 8: "saddr", 24: "daddr"}[p.Offset]
if cmp.Op == expr.CmpOpNeq {
name = "!" + name
}
ones, _ := net.IPMask(bw.Mask).Size()
parts = append(parts, fmt.Sprintf("%s=%s/%d", name, net.IP(cmp.Data), ones))
continue
}
}
cmp, ok := func() (*expr.Cmp, bool) { cmp, ok := func() (*expr.Cmp, bool) {
if i+1 >= len(r.Exprs) { if i+1 >= len(r.Exprs) {
return nil, false return nil, false
@@ -1909,9 +1921,19 @@ func describeRule(r ManagedRule) string {
case *expr.Meta: case *expr.Meta:
switch m.Key { switch m.Key {
case expr.MetaKeyIIFNAME: case expr.MetaKeyIIFNAME:
parts = append(parts, "iif="+strings.TrimRight(string(cmp.Data), "\x00")) op := "="
if cmp.Op == expr.CmpOpNeq {
op = "!="
}
parts = append(parts, "iif"+op+strings.TrimRight(string(cmp.Data), "\x00"))
case expr.MetaKeyOIFNAME: case expr.MetaKeyOIFNAME:
parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00")) op := "="
if cmp.Op == expr.CmpOpNeq {
op = "!="
}
parts = append(parts, "oif"+op+strings.TrimRight(string(cmp.Data), "\x00"))
case expr.MetaKeyNFPROTO:
parts = append(parts, map[byte]string{unix.NFPROTO_IPV4: "ip4", unix.NFPROTO_IPV6: "ip6"}[cmp.Data[0]])
} }
case *expr.Payload: case *expr.Payload:
if m.Base == expr.PayloadBaseNetworkHeader && (m.Len == 4 || m.Len == 16) { if m.Base == expr.PayloadBaseNetworkHeader && (m.Len == 4 || m.Len == 16) {
@@ -2050,25 +2072,25 @@ func TestCompile_CommaZoneLists(t *testing.T) {
{ {
name: "dnat origdest", name: "dnat origdest",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"}, rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"},
want: map[string][]string{"prerouting": {"iif=eth0 daddr=203.0.113.5"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 daddr=203.0.113.5"},
"forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}},
}, },
{ {
name: "dnat origdest list", name: "dnat origdest list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5,203.0.113.6"}, rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5,203.0.113.6"},
want: map[string][]string{"prerouting": {"iif=eth0 daddr=203.0.113.5", "iif=eth0 daddr=203.0.113.6"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 daddr=203.0.113.5", "iif=eth0 ip4 daddr=203.0.113.6"},
"forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}},
}, },
{ {
name: "dnat negated origdest list", name: "dnat negated origdest list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,203.0.113.6"}, rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,203.0.113.6"},
want: map[string][]string{"prerouting": {"iif=eth0 !daddr=203.0.113.5 !daddr=203.0.113.6"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 !daddr=203.0.113.5 !daddr=203.0.113.6"},
"forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}},
}, },
{ {
name: "accept origdest", name: "accept origdest",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"}, rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"},
want: map[string][]string{"input": {"iif=eth0 daddr=203.0.113.5"}}, want: map[string][]string{"input": {"iif=eth0 ip4 daddr=203.0.113.5"}},
}, },
{ {
name: "origdest does not scope interface-less zone", name: "origdest does not scope interface-less zone",
@@ -2078,7 +2100,7 @@ func TestCompile_CommaZoneLists(t *testing.T) {
{ {
name: "accept ipv6 origdest", name: "accept ipv6 origdest",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "2001:db8::5"}, rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "2001:db8::5"},
want: map[string][]string{"input": {"iif=eth0 daddr=2001:db8::5"}}, want: map[string][]string{"input": {"iif=eth0 ip6 daddr=2001:db8::5"}},
}, },
{ {
name: "blrule zone list", name: "blrule zone list",
@@ -3105,3 +3127,614 @@ func TestSpecCount_CommaAllMatchesExpansion(t *testing.T) {
} }
} }
} }
func TestCompile_Dispositions(t *testing.T) {
cases := []struct {
invalid, untracked config.PolicyAction
want map[string]expr.Any // tag prefix -> verdict expr, nil = absent
}{
{"", "", map[string]expr.Any{"ct:invalid:": &expr.Verdict{Kind: expr.VerdictDrop}, "ct:untracked:": nil}},
{config.PolicyContinue, config.PolicyContinue, map[string]expr.Any{"ct:invalid:": nil, "ct:untracked:": nil}},
{config.PolicyReject, config.PolicyAccept, map[string]expr.Any{
"ct:invalid:": rejectExprs(0, config.FamilyINET)[0],
"ct:untracked:": &expr.Verdict{Kind: expr.VerdictAccept},
}},
}
for _, tc := range cases {
cfg := &config.Config{
Settings: config.Settings{
TableName: "test",
AddressFamily: config.FamilyINET,
InvalidDisposition: tc.invalid,
UntrackedDisposition: tc.untracked,
},
Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}},
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}},
Policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}},
PortGroups: make(map[string]config.PortGroup),
}
state, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatalf("Compile: %v", err)
}
for _, chain := range []string{"input", "forward", "output"} {
for prefix, want := range tc.want {
var got *ManagedRule
for i, r := range state.Rules[chain] {
if r.Tag == prefix+chain {
got = &state.Rules[chain][i]
}
}
switch {
case want == nil && got != nil:
t.Errorf("%q/%q: unexpected %s%s rule", tc.invalid, tc.untracked, prefix, chain)
case want != nil && got == nil:
t.Errorf("%q/%q: missing %s%s rule", tc.invalid, tc.untracked, prefix, chain)
case want != nil && !reflect.DeepEqual(got.Exprs[len(got.Exprs)-1], want):
t.Errorf("%q/%q: %s%s verdict = %#v, want %#v", tc.invalid, tc.untracked, prefix, chain, got.Exprs[len(got.Exprs)-1], want)
}
}
}
}
}
func TestCompile_LogLimit(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{
TableName: "test",
AddressFamily: config.FamilyINET,
LogLimit: "1/sec:10",
},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}},
Policy: []config.Policy{
{Source: "net", Dest: "all", Action: config.PolicyDrop, Log: "info"},
},
Rules: []config.Rule{
{Action: config.RuleAccept, Source: "net:192.0.2.1", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, Log: "info"},
{Action: config.RuleLog, Source: "net:198.51.100.0/24", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"23"}, Log: "info"},
},
PortGroups: make(map[string]config.PortGroup),
}
state, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
want := &expr.Limit{Type: expr.LimitTypePkts, Rate: 1, Unit: expr.LimitTimeSecond, Burst: 10}
byTag := map[string][]ManagedRule{}
for _, r := range state.Rules["input"] {
byTag[r.Tag] = append(byTag[r.Tag], r)
}
for tag, n := range map[string]int{"policy:0": 2, "rule:0": 2, "rule:1": 1} {
rs := byTag[tag]
if len(rs) != n {
t.Fatalf("%s: %d rules, want %d", tag, len(rs), n)
}
logRule := rs[0].Exprs
l := len(logRule)
if l < 2 || !reflect.DeepEqual(logRule[l-2], want) {
t.Errorf("%s: want limit before log, got %#v", tag, logRule)
}
if _, ok := logRule[l-1].(*expr.Log); !ok {
t.Errorf("%s: log rule must end in log, got %T", tag, logRule[l-1])
}
for _, e := range rs[n-1].Exprs[:len(rs[n-1].Exprs)-1] {
if _, ok := e.(*expr.Log); n == 2 && ok {
t.Errorf("%s: verdict rule still logs", tag)
}
}
}
if logs := byTag["policy:0"]; len(logs) == 2 {
if _, ok := logs[1].Exprs[len(logs[1].Exprs)-1].(*expr.Verdict); !ok {
t.Errorf("policy verdict rule must end in a verdict")
}
}
}
func TestCompile_NoLogLimitKeepsInlineLog(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}},
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}},
Policy: []config.Policy{{Source: "net", Dest: "all", Action: config.PolicyDrop, Log: "info"}},
PortGroups: make(map[string]config.PortGroup),
}
state, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatal(err)
}
for _, r := range state.Rules["input"] {
for _, e := range r.Exprs {
if _, ok := e.(*expr.Limit); ok {
t.Errorf("%s: unexpected limit without log_limit", r.Tag)
}
}
}
}
func TestCompile_LogLimitSplitsAroundExtrasAndNAT(t *testing.T) {
cfg := listCfg(func(c *config.Config) {
c.Settings.LogLimit = "1/sec:10"
c.Zones["loc"] = config.Zone{Type: config.ZoneIP}
c.Interfaces = append(c.Interfaces, config.Interface{Zone: "loc", Interface: "eth1"})
c.Rules = []config.Rule{
{Action: config.RuleAccept, Source: "net:192.0.2.1", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, Log: "info",
Mark: "0x1", User: "0", Time: &config.TimeSpec{Start: "08:00", Stop: "17:00"}, RateLimit: "5/sec:20"},
{Action: config.RuleLog, Source: "net:192.0.2.2", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"23"}, Log: "info",
Mark: "0x2", RateLimit: "5/sec:20"},
{Action: config.RuleDNAT, Source: "net", Dest: "loc:198.51.100.10:80", Proto: "tcp", DPort: config.PortSpec{"8000"}, Log: "info"},
{Action: config.RuleAccept, Source: "net:192.0.2.3", Dest: "loc", Proto: "tcp", DPort: config.PortSpec{"443"}, Log: "info"},
}
c.SNAT = []config.SNATRule{{Action: config.SNATAddress, Source: "198.51.100.0/24", Dest: "eth0", Address: "203.0.113.7", Mark: "0x3", Log: "info"}}
})
state, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
logLimit := &expr.Limit{Type: expr.LimitTypePkts, Rate: 1, Unit: expr.LimitTimeSecond, Burst: 10}
kinds := func(exprs []expr.Any) (logs, limits, marks, uids int) {
for _, e := range exprs {
switch v := e.(type) {
case *expr.Log:
logs++
case *expr.Limit:
limits++
case *expr.Meta:
if v.Key == expr.MetaKeyMARK && !v.SourceRegister {
marks++
}
if v.Key == expr.MetaKeySKUID {
uids++
}
}
}
return
}
check := func(name string, rs []ManagedRule, n, wantMarks, wantUIDs int) []expr.Any {
t.Helper()
if len(rs) != n {
t.Fatalf("%s: %d rules, want %d", name, len(rs), n)
}
l := rs[0].Exprs
if len(l) < 2 || !reflect.DeepEqual(l[len(l)-2], logLimit) {
t.Errorf("%s: log rule must end limit+log, got %#v", name, l)
}
if logs, limits, marks, uids := kinds(l); logs != 1 || limits != 1 || marks != wantMarks || uids != wantUIDs {
t.Errorf("%s: log rule logs=%d limits=%d marks=%d uids=%d, want 1/1/%d/%d", name, logs, limits, marks, uids, wantMarks, wantUIDs)
}
if n == 1 {
return nil
}
v := rs[1].Exprs
if logs, _, marks, uids := kinds(v); logs != 0 || marks != wantMarks || uids != wantUIDs {
t.Errorf("%s: action rule logs=%d marks=%d uids=%d, want 0/%d/%d", name, logs, marks, uids, wantMarks, wantUIDs)
}
return v
}
v := check("rule:0", taggedRules(state, "input", "rule:0"), 2, 1, 1)
if _, limits, _, _ := kinds(v); limits != 1 {
t.Errorf("rule:0: action rule must keep its ratelimit, got %d limits", limits)
}
if _, ok := v[len(v)-1].(*expr.Verdict); !ok {
t.Errorf("rule:0: action rule must end in a verdict, got %T", v[len(v)-1])
}
check("rule:1", taggedRules(state, "input", "rule:1"), 1, 1, 0)
if v := check("rule:2", taggedRules(state, "prerouting", "rule:2"), 2, 0, 0); v != nil {
if _, ok := v[len(v)-1].(*expr.NAT); !ok {
t.Errorf("rule:2: DNAT rule must end in nat, got %T", v[len(v)-1])
}
}
check("rule:3", taggedRules(state, "forward", "rule:3"), 2, 0, 0)
var snat []ManagedRule
for _, r := range state.Rules["postrouting"] {
if logs, _, _, _ := kinds(r.Exprs); logs > 0 || strings.HasPrefix(r.Tag, "snat:0") {
snat = append(snat, r)
}
}
if v := check("snat:0", snat, 2, 1, 0); v != nil {
if _, ok := v[len(v)-1].(*expr.NAT); !ok {
t.Errorf("snat:0: rule must end in nat, got %T", v[len(v)-1])
}
}
}
// hostsCfg models a shorewall setup where lan:net is defined by hosts on net's interfaces.
func hostsCfg(mod func(*config.Config)) *config.Config {
cfg := &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP},
"lan": {Type: config.ZoneIP, Parents: []string{"net"}}, "vpn": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "wlo1"}, {Zone: "net", Interface: "enp2s0"}, {Zone: "vpn", Interface: "tun0"},
},
Hosts: []config.Host{
{Zone: "lan", Interface: "wlo1", Addresses: []string{"192.0.2.0/24"}},
{Zone: "lan", Interface: "enp2s0", Addresses: []string{"198.51.100.0/24"}},
},
Policy: []config.Policy{
{Source: "net", Dest: "all", Action: config.PolicyDrop},
{Source: "lan", Dest: "fw", Action: config.PolicyReject},
{Source: "all", Dest: "all", Action: config.PolicyReject},
},
Rules: []config.Rule{
{Action: config.RuleAccept, Source: "lan,vpn", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"6768"}},
},
PortGroups: make(map[string]config.PortGroup),
}
if mod != nil {
mod(cfg)
}
return cfg
}
func describeTagged(state *FirewallState, chain, tag string) []string {
var out []string
for _, r := range taggedRules(state, chain, tag) {
out = append(out, describeRule(r))
}
return out
}
func TestCompile_HostsZoneRules(t *testing.T) {
var logs bytes.Buffer
prev := slog.Default()
slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil)))
defer slog.SetDefault(prev)
state := mustCompile(t, hostsCfg(nil))
want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24", "iif=tun0"}
if got := describeTagged(state, "input", "rule:0"); !reflect.DeepEqual(got, want) {
t.Errorf("input rule:0 = %q, want %q", got, want)
}
if strings.Contains(logs.String(), "skipping") {
t.Errorf("unexpected warning:\n%s", logs.String())
}
}
func TestCompile_HostsSubZoneBeforeParent(t *testing.T) {
state := mustCompile(t, hostsCfg(nil))
want := []string{
"iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6",
"iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6",
}
if got := describeTagged(state, "input", "policy:0"); !reflect.DeepEqual(got, want) {
t.Errorf("net->fw policy = %q, want %q", got, want)
}
want = []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}
if got := describeTagged(state, "input", "policy:1"); !reflect.DeepEqual(got, want) {
t.Errorf("lan->fw policy = %q, want %q", got, want)
}
for _, r := range taggedRules(state, "input", "policy:1") {
if !slices.ContainsFunc(r.Exprs, func(e expr.Any) bool {
m, ok := e.(*expr.Meta)
return ok && m.Key == expr.MetaKeyNFPROTO
}) {
t.Errorf("host match lacks an nfproto guard: %v", describeRule(r))
}
}
}
func TestCompile_HostsZoneMatches(t *testing.T) {
tests := []struct {
name string
mod func(*config.Config)
chain string
tag string
want []string
}{
{"forward to hosts zone", func(c *config.Config) {
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "vpn", Dest: "lan", Proto: "tcp"}}
}, "forward", "rule:0", []string{"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.0/24", "iif=tun0 oif=enp2s0 ip4 daddr=198.51.100.0/24"}},
{"rule address prunes non-overlapping hosts", func(c *config.Config) {
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "lan:192.0.2.5", Dest: "fw", Proto: "tcp"}}
}, "input", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 saddr=192.0.2.5"}},
{"host exclusions and sub-zone exclusions fall back to the parent", func(c *config.Config) {
c.Hosts[1].Exclusions = []string{"198.51.100.7"}
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp"}}
}, "input", "rule:0", []string{
"iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6",
"iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6",
"iif=enp2s0 ip4 saddr=198.51.100.7",
}},
{"host exclusion", func(c *config.Config) {
c.Hosts[1].Exclusions = []string{"198.51.100.7"}
}, "input", "policy:1", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24 !saddr=198.51.100.7"}},
{"DNAT from hosts zone", func(c *config.Config) {
c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "lan", Dest: "vpn:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"80"}}}
}, "prerouting", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}},
{"conntrack from hosts zone", func(c *config.Config) {
c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Source: "lan", Proto: "udp"}}
}, "raw_prerouting", "conntrack:0:raw_prerouting", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}},
{"blrule from hosts zone", func(c *config.Config) {
c.Blrules = []config.BlruleRule{{Action: config.BlruleDrop, Source: "lan", Dest: "all"}}
}, "input", "blrule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}},
{"all expansion includes hosts zone", func(c *config.Config) {
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "all!net,vpn", Dest: "fw", Proto: "tcp"}}
}, "input", "rule:0", []string{"iif=wlo1 ip4 saddr=192.0.2.0/24", "iif=enp2s0 ip4 saddr=198.51.100.0/24"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
state := mustCompile(t, hostsCfg(tt.mod))
if got := describeTagged(state, tt.chain, tt.tag); !reflect.DeepEqual(got, tt.want) {
t.Errorf("%s %s = %q, want %q", tt.chain, tt.tag, got, tt.want)
}
})
}
}
func TestCompile_IntraZoneMultiInterface(t *testing.T) {
tests := []struct {
name string
policy []config.Policy
tag string
want []string
}{
{
name: "implicit accept between distinct interfaces",
policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}},
tag: "intra:lxd",
want: []string{"iif=lxdbr0 oif=docker0", "iif=lxdbr0 oif=br-", "iif=docker0 oif=lxdbr0", "iif=docker0 oif=br-", "iif=br- oif=lxdbr0", "iif=br- oif=docker0"},
},
{
name: "explicit zone policy overrides",
policy: []config.Policy{{Source: "lxd", Dest: "lxd", Action: config.PolicyDrop, Log: "info"}, {Source: "all", Dest: "all", Action: config.PolicyDrop}},
tag: "policy:0",
want: []string{"iif=lxdbr0 oif=docker0", "iif=lxdbr0 oif=br-", "iif=docker0 oif=lxdbr0", "iif=docker0 oif=br-", "iif=br- oif=lxdbr0", "iif=br- oif=docker0"},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
"lxd": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "lxd", Interface: "lxdbr0"},
{Zone: "lxd", Interface: "docker0"},
{Zone: "lxd", Interface: "br-+"},
},
Policy: tc.policy,
PortGroups: map[string]config.PortGroup{},
}
state, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
var got, all []string
for _, r := range state.Rules["forward"] {
if r.Tag == tc.tag {
got = append(got, describeRule(r))
}
if strings.HasPrefix(r.Tag, "intra:") {
all = append(all, r.Tag)
}
}
if !reflect.DeepEqual(got, tc.want) {
t.Errorf("%s rules = %q, want %q", tc.tag, got, tc.want)
}
if tc.tag != "intra:lxd" && len(all) != 0 {
t.Errorf("explicit policy must replace implicit accept, got %q", all)
}
last := taggedRules(state, "forward", tc.tag)
if len(last) == 0 {
return
}
if v, ok := last[0].Exprs[len(last[0].Exprs)-1].(*expr.Verdict); !ok || (tc.tag == "intra:lxd") != (v.Kind == expr.VerdictAccept) {
t.Errorf("%s verdict = %#v", tc.tag, last[0].Exprs[len(last[0].Exprs)-1])
}
})
}
}
func TestCompile_HostsAddressMatchesFamilyGuarded(t *testing.T) {
state := mustCompile(t, hostsCfg(func(c *config.Config) {
c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64")
c.Hosts[0].Exclusions = []string{"192.0.2.9"}
}))
want := []string{
"iif=wlo1 ip4 !saddr=192.0.2.0/24", "iif=wlo1 ip6 !saddr=2001:db8::/64",
"iif=wlo1 ip4 saddr=192.0.2.9",
"iif=enp2s0 ip4 !saddr=198.51.100.0/24", "iif=enp2s0 ip6",
}
if got := describeTagged(state, "input", "policy:0"); !reflect.DeepEqual(got, want) {
t.Errorf("net->fw DROP policy = %q, want %q", got, want)
}
want = []string{
"iif=wlo1 ip4 saddr=192.0.2.0/24 !saddr=192.0.2.9", "iif=wlo1 ip6 saddr=2001:db8::/64",
"iif=enp2s0 ip4 saddr=198.51.100.0/24",
}
if got := describeTagged(state, "input", "policy:1"); !reflect.DeepEqual(got, want) {
t.Errorf("lan->fw policy = %q, want %q", got, want)
}
for chain, rules := range state.Rules {
for _, r := range rules {
var fam byte
for i, e := range r.Exprs {
if m, ok := e.(*expr.Meta); ok && m.Key == expr.MetaKeyNFPROTO {
fam = r.Exprs[i+1].(*expr.Cmp).Data[0]
}
p, ok := e.(*expr.Payload)
if !ok || p.Base != expr.PayloadBaseNetworkHeader {
continue
}
if p.Len == 4 && fam != unix.NFPROTO_IPV4 || p.Len == 16 && fam != unix.NFPROTO_IPV6 {
t.Errorf("%s %s: %d-byte address compare without its family guard: %s", chain, r.Tag, p.Len, describeRule(r))
}
}
}
}
}
func TestCompile_FirewallSelfPolicySkipped(t *testing.T) {
for _, action := range []config.PolicyAction{config.PolicyAccept, config.PolicyDrop} {
t.Run(string(action), func(t *testing.T) {
state := mustCompile(t, listCfg(func(c *config.Config) {
c.Policy = []config.Policy{
{Source: "fw", Dest: "fw", Action: action},
{Source: "net", Dest: "fw", Action: config.PolicyDrop, Log: "info"},
}
}))
for _, chain := range []string{"input", "output", "forward"} {
if got := taggedRules(state, chain, "policy:0"); len(got) != 0 {
t.Errorf("fw->fw emitted %d rules in %s", len(got), chain)
}
}
got := taggedRules(state, "input", "policy:1")
if len(got) != 1 || describeRule(got[0]) != "iif=eth0" {
t.Errorf("net->fw input rules = %d, want one scoped to eth0", len(got))
}
})
}
}
func TestCompile_DestPlusOverridesIntraZone(t *testing.T) {
state := mustCompile(t, listCfg(func(c *config.Config) {
c.Zones["lxd"] = config.Zone{Type: config.ZoneIP}
c.Interfaces = append(c.Interfaces, config.Interface{Zone: "lxd", Interface: "lxdbr0"}, config.Interface{Zone: "lxd", Interface: "docker0"})
c.Policy = []config.Policy{{Source: "lxd", Dest: "all+", Action: config.PolicyDrop}}
}))
if got := taggedRules(state, "forward", "intra:lxd"); len(got) != 0 {
t.Errorf("lxd all+ must override implicit intra-zone accept, got %d rules", len(got))
}
if got := taggedRules(state, "forward", "policy:0"); len(got) == 0 {
t.Error("lxd all+ emitted no forward rules")
}
}
func TestCompile_HostsIntraZone(t *testing.T) {
state := mustCompile(t, hostsCfg(nil))
want := []string{
"iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 ip4 daddr=198.51.100.0/24",
"iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24",
}
if got := describeTagged(state, "forward", "intra:lan"); !reflect.DeepEqual(got, want) {
t.Errorf("lan intra = %q, want %q", got, want)
}
want = []string{
"iif=wlo1 ip4 !saddr=192.0.2.0/24 oif=enp2s0 ip4 !daddr=198.51.100.0/24",
"iif=wlo1 ip6 oif=enp2s0 ip6",
"iif=enp2s0 ip4 !saddr=198.51.100.0/24 oif=wlo1 ip4 !daddr=192.0.2.0/24",
"iif=enp2s0 ip6 oif=wlo1 ip6",
}
if got := describeTagged(state, "forward", "intra:net"); !reflect.DeepEqual(got, want) {
t.Errorf("net intra = %q, want %q", got, want)
}
state = mustCompile(t, hostsCfg(func(c *config.Config) {
c.Policy = append([]config.Policy{{Source: "lan", Dest: "lan", Action: config.PolicyDrop}}, c.Policy...)
}))
if got := describeTagged(state, "forward", "intra:lan"); len(got) != 0 {
t.Errorf("explicit lan lan policy must replace implicit accept, got %q", got)
}
want = []string{
"iif=wlo1 ip4 saddr=192.0.2.0/24 oif=enp2s0 ip4 daddr=198.51.100.0/24",
"iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24",
}
if got := describeTagged(state, "forward", "policy:0"); !reflect.DeepEqual(got, want) {
t.Errorf("lan lan policy = %q, want %q", got, want)
}
}
func TestCompile_HostsRouteBack(t *testing.T) {
sameIface := func(routeback bool) func(*config.Config) {
return func(c *config.Config) {
c.Hosts = []config.Host{
{Zone: "lan", Interface: "wlo1", Addresses: []string{"192.0.2.0/24"}, Options: config.HostOptions{RouteBack: routeback}},
{Zone: "lan", Interface: "wlo1", Addresses: []string{"198.51.100.0/24"}},
}
}
}
if got := describeTagged(mustCompile(t, hostsCfg(sameIface(false))), "forward", "intra:lan"); len(got) != 0 {
t.Errorf("same-interface hosts without routeback = %q, want none", got)
}
want := []string{"iif=wlo1 ip4 saddr=192.0.2.0/24 oif=wlo1 ip4 daddr=192.0.2.0/24"}
if got := describeTagged(mustCompile(t, hostsCfg(sameIface(true))), "forward", "intra:lan"); !reflect.DeepEqual(got, want) {
t.Errorf("routeback hosts intra = %q, want %q", got, want)
}
state := mustCompile(t, hostsCfg(func(c *config.Config) {
sameIface(true)(c)
c.Policy = append([]config.Policy{{Source: "lan", Dest: "lan", Action: config.PolicyDrop}}, c.Policy...)
}))
if got := describeTagged(state, "forward", "policy:0"); !reflect.DeepEqual(got, want) {
t.Errorf("routeback hosts lan lan policy = %q, want %q", got, want)
}
}
func TestCompile_WildcardParentExcludesSubZoneHosts(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP},
"lan": {Type: config.ZoneIP, Parents: []string{"net"}}, "lxd": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "enp+"}, {Zone: "net", Interface: "wlo1"}, {Zone: "lxd", Interface: "lxdbr0"},
},
Hosts: []config.Host{
{Zone: "lan", Interface: "enp2s0", Addresses: []string{"192.0.2.0/24"}},
{Zone: "lan", Interface: "enp4+", Addresses: []string{"203.0.113.0/24"}},
{Zone: "lan", Interface: "wlo1", Addresses: []string{"198.51.100.0/24"}},
},
Policy: []config.Policy{
{Source: "lxd", Dest: "net", Action: config.PolicyAccept},
{Source: "fw", Dest: "net", Action: config.PolicyDrop},
{Source: "net", Dest: "all", Action: config.PolicyDrop},
{Source: "all", Dest: "all", Action: config.PolicyReject},
},
PortGroups: make(map[string]config.PortGroup),
}
state := mustCompile(t, cfg)
for _, tt := range []struct{ chain, tag, dir string }{
{"forward", "policy:0", "iif=lxdbr0 oif"}, {"output", "policy:1", "oif"}, {"input", "policy:2", "iif"},
} {
d, a := tt.dir, "daddr"
if tt.chain == "input" {
a = "saddr"
}
want := []string{
d + "=enp " + d[len(d)-3:] + "!=enp2s0 " + d[len(d)-3:] + "!=enp4",
d + "=enp2s0 ip4 !" + a + "=192.0.2.0/24", d + "=enp2s0 ip6",
d + "=enp4 ip4 !" + a + "=203.0.113.0/24", d + "=enp4 ip6",
d + "=wlo1 ip4 !" + a + "=198.51.100.0/24", d + "=wlo1 ip6",
}
if got := describeTagged(state, tt.chain, tt.tag); !reflect.DeepEqual(got, want) {
t.Errorf("%s %s = %q, want %q", tt.chain, tt.tag, got, want)
}
}
for chain, want := range map[string]string{
"forward": "iif=lxdbr0 oif=enp2s0 ip4 daddr=192.0.2.0/24",
"output": "oif=enp2s0 ip4 daddr=192.0.2.0/24",
"input": "iif=enp2s0 ip4 saddr=192.0.2.0/24",
} {
if got := describeTagged(state, chain, "policy:3"); !slices.Contains(got, want) {
t.Errorf("%s policy:3 (lan reject) = %q, want it to contain %q", chain, got, want)
}
}
}
func TestIfaceCovers(t *testing.T) {
for _, tt := range []struct {
a, b string
want bool
}{
{"enp2s0", "enp2s0", true}, {"enp2s0", "enp3s0", false},
{"enp+", "enp2s0", true}, {"enp2s0", "enp+", false}, {"enp+", "wlo1", false},
{"en+", "enp+", true}, {"enp+", "en+", false}, {"enp+", "eno+", false}, {"enp+", "enp", true},
} {
if got := ifaceCovers(tt.a, tt.b); got != tt.want {
t.Errorf("ifaceCovers(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.want)
}
}
}
+137 -18
View File
@@ -5,6 +5,8 @@ import (
"github.com/google/nftables" "github.com/google/nftables"
"github.com/google/nftables/expr" "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"
) )
@@ -15,16 +17,101 @@ type Engine struct {
} }
func NewEngine(cfg *config.Config) (*Engine, error) { func NewEngine(cfg *config.Config) (*Engine, error) {
conn, err := nftables.New() conn, err := nftables.New(nftables.WithSockOptions(largeBuffers))
if err != nil { if err != nil {
return nil, fmt.Errorf("connecting to nftables: %w", err) return nil, fmt.Errorf("connecting to nftables: %w", err)
} }
return &Engine{cfg: cfg, conn: conn}, nil return &Engine{cfg: cfg, conn: conn}, nil
} }
var tableFamilies = map[config.AddressFamily]nftables.TableFamily{
config.FamilyINET: nftables.TableFamilyINet,
config.FamilyIP: nftables.TableFamilyIPv4,
config.FamilyIP6: nftables.TableFamilyIPv6,
}
func (e *Engine) family() nftables.TableFamily {
if f, ok := tableFamilies[e.cfg.Settings.AddressFamily]; ok {
return f
}
return nftables.TableFamilyINet
}
func addressFamily(tf nftables.TableFamily) config.AddressFamily {
for f, t := range tableFamilies {
if t == tf {
return f
}
}
return config.FamilyINET
}
// withFamily is the engine for the same table name in another address family.
func (e *Engine) withFamily(f config.AddressFamily) *Engine {
cfg := *e.cfg
cfg.Settings.AddressFamily = f
return &Engine{cfg: &cfg, conn: e.conn}
}
// overlaps reports whether tables of families a and b filter the same traffic:
// inet covers both ip and ip6, which do not overlap each other.
func overlaps(a, b nftables.TableFamily) bool {
if a == b {
return false
}
return (a == nftables.TableFamilyINet && (b == nftables.TableFamilyIPv4 || b == nftables.TableFamilyIPv6)) ||
(b == nftables.TableFamilyINet && (a == nftables.TableFamilyIPv4 || a == nftables.TableFamilyIPv6))
}
// staleTables are our-named tables left by a different address_family that
// would still filter the traffic this family now owns.
func (e *Engine) staleTables() ([]*nftables.Table, error) {
tables, err := e.conn.ListTables()
if err != nil {
return nil, fmt.Errorf("listing tables: %w", err)
}
var stale []*nftables.Table
for _, t := range tables {
if t.Name == e.cfg.Settings.TableName && overlaps(e.family(), t.Family) {
stale = append(stale, t)
}
}
return stale, nil
}
// batchBufSize bounds one batch: the kernel rejects a batch larger than the
// send buffer (EMSGSIZE) and drops ACKs beyond the receive buffer (ENOBUFS)
// after committing it.
// ponytail: fixed cap of tens of thousands of rules; size per batch if exceeded.
const batchBufSize = 64 << 20
// largeBuffers raises both socket buffers, ignoring rmem_max/wmem_max when
// CAP_NET_ADMIN allows it and falling back to the capped sizes otherwise.
func largeBuffers(c *netlink.Conn) error {
rc, err := c.SyscallConn()
if err != nil {
return err
}
var serr error
err = rc.Control(func(fd uintptr) {
for _, o := range [][2]int{{unix.SO_SNDBUFFORCE, unix.SO_SNDBUF}, {unix.SO_RCVBUFFORCE, unix.SO_RCVBUF}} {
if unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[0], batchBufSize) == nil {
continue
}
if serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[1], batchBufSize); serr != nil {
return
}
}
})
if err != nil {
return err
}
return serr
}
func (e *Engine) ensureTable() *nftables.Table { func (e *Engine) ensureTable() *nftables.Table {
return e.conn.AddTable(&nftables.Table{ return e.conn.AddTable(&nftables.Table{
Family: nftables.TableFamilyINet, Family: e.family(),
Name: e.cfg.Settings.TableName, Name: e.cfg.Settings.TableName,
}) })
} }
@@ -132,6 +219,13 @@ func (e *Engine) Apply(changes *ChangeSet) error {
} }
func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error { func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error {
stale, err := e.staleTables()
if err != nil {
return err
}
for _, t := range stale {
e.conn.DelTable(t)
}
table := e.ensureTable() table := e.ensureTable()
chains := e.ensureChains(table, policies) chains := e.ensureChains(table, policies)
@@ -173,18 +267,24 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol
} }
func (e *Engine) Flush() error { func (e *Engine) Flush() error {
tables, err := e.conn.ListTables() tables, err := e.staleTables()
if err != nil { if err != nil {
return fmt.Errorf("listing tables: %w", err) return err
} }
own, err := e.findTable()
for _, t := range tables { if err != nil {
if t.Name == e.cfg.Settings.TableName { return err
e.conn.DelTable(t)
return e.conn.Flush()
} }
if own != nil {
tables = append(tables, own)
} }
if len(tables) == 0 {
return nil return nil
}
for _, t := range tables {
e.conn.DelTable(t)
}
return e.conn.Flush()
} }
func (e *Engine) findTable() (*nftables.Table, error) { func (e *Engine) findTable() (*nftables.Table, error) {
@@ -193,7 +293,7 @@ func (e *Engine) findTable() (*nftables.Table, error) {
return nil, fmt.Errorf("listing tables: %w", err) return nil, fmt.Errorf("listing tables: %w", err)
} }
for _, t := range tables { for _, t := range tables {
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { if t.Name == e.cfg.Settings.TableName && t.Family == e.family() {
return t, nil return t, nil
} }
} }
@@ -222,7 +322,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
} }
} }
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)
} }
@@ -252,6 +352,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
// survives the process that took it. // survives the process that took it.
type Snapshot struct { type Snapshot struct {
Table string `json:"table"` Table string `json:"table"`
Family config.AddressFamily `json:"family,omitempty"`
Present bool `json:"present"` Present bool `json:"present"`
Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"` Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"`
Rules map[string][]SnapshotRule `json:"rules,omitempty"` Rules map[string][]SnapshotRule `json:"rules,omitempty"`
@@ -264,16 +365,30 @@ type SnapshotRule struct {
Exprs [][]byte `json:"exprs"` Exprs [][]byte `json:"exprs"`
} }
// Snapshot captures the live tomswall table so Restore can roll back to it. // Snapshot captures the live tomswall table so Restore can roll back to it,
// falling back to the overlapping table of another family that apply replaces.
func (e *Engine) Snapshot() (*Snapshot, error) { func (e *Engine) Snapshot() (*Snapshot, error) {
snap := &Snapshot{Table: e.cfg.Settings.TableName}
t, err := e.findTable() t, err := e.findTable()
if err != nil || t == nil { if err != nil {
return snap, err return nil, err
}
if t == nil {
stale, err := e.staleTables()
if err != nil {
return nil, err
}
// ponytail: captures one stale table; an inet config replacing both ip and ip6 restores only the first.
if len(stale) > 0 {
return e.withFamily(addressFamily(stale[0].Family)).Snapshot()
}
}
snap := &Snapshot{Table: e.cfg.Settings.TableName, Family: addressFamily(e.family())}
if t == nil {
return snap, nil
} }
snap.Present = true snap.Present = true
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) chains, err := e.conn.ListChainsOfTableFamily(e.family())
if err != nil { if err != nil {
return nil, fmt.Errorf("listing chains: %w", err) return nil, fmt.Errorf("listing chains: %w", err)
} }
@@ -288,7 +403,7 @@ func (e *Engine) Snapshot() (*Snapshot, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
snap.Rules, err = encodeState(state) snap.Rules, err = encodeState(state, byte(e.family()))
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -302,10 +417,14 @@ func (e *Engine) Restore(s *Snapshot) error {
if s.Table != e.cfg.Settings.TableName { if s.Table != e.cfg.Settings.TableName {
return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName) return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName)
} }
// Snapshots predating the family field are of the inet table.
if f := addressFamily(tableFamilies[s.Family]); f != addressFamily(e.family()) {
return e.withFamily(f).Restore(s)
}
if !s.Present { if !s.Present {
return e.Flush() return e.Flush()
} }
want, err := decodeState(s.Rules) want, err := decodeState(s.Rules, byte(e.family()))
if err != nil { if err != nil {
return err return err
} }
+59
View File
@@ -0,0 +1,59 @@
package nftables
import (
"os"
"strconv"
"strings"
"testing"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
)
func TestLargeBuffersRaisesSocketBuffers(t *testing.T) {
c, err := netlink.Dial(unix.NETLINK_NETFILTER, nil)
if err != nil {
t.Skipf("netlink unavailable: %v", err)
}
defer c.Close()
if err := largeBuffers(c); err != nil {
t.Fatal(err)
}
// Without CAP_NET_ADMIN the kernel caps at the sysctl max; it doubles either way.
for opt, sysctl := range map[int]string{unix.SO_RCVBUF: "rmem_max", unix.SO_SNDBUF: "wmem_max"} {
want := 2 * min(batchBufSize, procInt(t, "/proc/sys/net/core/"+sysctl))
if got := sockBuf(t, c, opt); got < want {
t.Errorf("%s-bounded buffer = %d, want >= %d", sysctl, got, want)
}
}
}
func procInt(t *testing.T, path string) int {
t.Helper()
b, err := os.ReadFile(path)
if err != nil {
t.Skipf("reading %s: %v", path, err)
}
v, err := strconv.Atoi(strings.TrimSpace(string(b)))
if err != nil {
t.Fatal(err)
}
return v
}
func sockBuf(t *testing.T, c *netlink.Conn, opt int) int {
t.Helper()
rc, err := c.SyscallConn()
if err != nil {
t.Fatal(err)
}
var v int
var serr error
if err := rc.Control(func(fd uintptr) { v, serr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt) }); err != nil {
t.Fatal(err)
}
if serr != nil {
t.Fatal(serr)
}
return v
}
+113
View File
@@ -0,0 +1,113 @@
package nftables
import (
"testing"
"github.com/google/nftables"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
type sentTable struct {
msg int
family nftables.TableFamily
}
// familyEngine fakes a kernel holding tomswall tables of the given families
// and records table creations/deletions.
func familyEngine(t *testing.T, af config.AddressFamily, live ...nftables.TableFamily) (*Engine, *[]sentTable) {
var sent []sentTable
e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) {
var out []netlink.Message
for _, m := range req {
switch m.Header.Type {
case nftType(unix.NFT_MSG_GETTABLE):
for _, f := range live {
attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}})
out = append(out, netlink.Message{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append([]byte{byte(f), 0, 0, 0}, attrs...)})
}
case nftType(unix.NFT_MSG_NEWTABLE):
sent = append(sent, sentTable{unix.NFT_MSG_NEWTABLE, nftables.TableFamily(m.Data[0])})
case nftType(unix.NFT_MSG_DELTABLE):
sent = append(sent, sentTable{unix.NFT_MSG_DELTABLE, nftables.TableFamily(m.Data[0])})
}
}
return out, nil
})
e.cfg.Settings.AddressFamily = af
return e, &sent
}
func TestApplyIPFamilyReplacesInetTable(t *testing.T) {
e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyINet, nftables.TableFamilyIPv6)
if err := e.Apply(&ChangeSet{}); err != nil {
t.Fatal(err)
}
want := []sentTable{{unix.NFT_MSG_DELTABLE, nftables.TableFamilyINet}, {unix.NFT_MSG_NEWTABLE, nftables.TableFamilyIPv4}}
if len(*sent) != len(want) || (*sent)[0] != want[0] || (*sent)[1] != want[1] {
t.Errorf("got %+v, want %+v (the ip6 table must survive)", *sent, want)
}
}
func TestApplyInetFamilyReplacesIPTables(t *testing.T) {
e, sent := familyEngine(t, config.FamilyINET, nftables.TableFamilyIPv4, nftables.TableFamilyIPv6)
if err := e.Apply(&ChangeSet{}); err != nil {
t.Fatal(err)
}
var dels int
for _, s := range *sent {
if s.msg == unix.NFT_MSG_DELTABLE {
dels++
}
}
if dels != 2 {
t.Errorf("want ip and ip6 tables deleted, got %+v", *sent)
}
}
func TestFlushIPFamilyKeepsIP6Table(t *testing.T) {
e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyIPv6, nftables.TableFamilyIPv4)
if err := e.Flush(); err != nil {
t.Fatal(err)
}
if len(*sent) != 1 || (*sent)[0] != (sentTable{unix.NFT_MSG_DELTABLE, nftables.TableFamilyIPv4}) {
t.Errorf("got %+v, want only the ip table deleted", *sent)
}
}
func TestSnapshotFallsBackToReplacedTable(t *testing.T) {
e, _ := familyEngine(t, config.FamilyIP, nftables.TableFamilyINet)
snap, err := e.Snapshot()
if err != nil {
t.Fatal(err)
}
if !snap.Present || snap.Family != config.FamilyINET {
t.Fatalf("want present inet snapshot, got %+v", snap)
}
// Reverting to the inet snapshot drops the tried ip table.
e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyIPv4)
if err := e.Restore(&Snapshot{Table: "tomswall", Family: config.FamilyINET, Present: true}); err != nil {
t.Fatal(err)
}
want := []sentTable{{unix.NFT_MSG_DELTABLE, nftables.TableFamilyIPv4}, {unix.NFT_MSG_NEWTABLE, nftables.TableFamilyINet}}
if len(*sent) != 2 || (*sent)[0] != want[0] || (*sent)[1] != want[1] {
t.Errorf("got %+v, want %+v", *sent, want)
}
}
func TestRejectExprsFamily(t *testing.T) {
for af, want := range map[config.AddressFamily]expr.Reject{
config.FamilyINET: {Type: unix.NFT_REJECT_ICMPX_UNREACH, Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH},
config.FamilyIP: {Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 3},
config.FamilyIP6: {Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 4},
} {
got := rejectExprs(unix.IPPROTO_UDP, af)[0].(*expr.Reject)
if *got != want {
t.Errorf("%s: got %+v, want %+v", af, *got, want)
}
}
}
+7 -10
View File
@@ -4,14 +4,11 @@ import (
"encoding/binary" "encoding/binary"
"fmt" "fmt"
"github.com/google/nftables"
"github.com/google/nftables/expr" "github.com/google/nftables/expr"
"github.com/mdlayher/netlink" "github.com/mdlayher/netlink"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
const inet = byte(nftables.TableFamilyINet)
// exprByName mirrors the expression types google/nftables can parse back from the kernel. // exprByName mirrors the expression types google/nftables can parse back from the kernel.
var exprByName = map[string]func() expr.Any{ var exprByName = map[string]func() expr.Any{
"ct": func() expr.Any { return &expr.Ct{} }, "ct": func() expr.Any { return &expr.Ct{} },
@@ -40,13 +37,13 @@ var exprByName = map[string]func() expr.Any{
"notrack": func() expr.Any { return &expr.Notrack{} }, "notrack": func() expr.Any { return &expr.Notrack{} },
} }
func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) { func encodeState(state *FirewallState, fam byte) (map[string][]SnapshotRule, error) {
out := make(map[string][]SnapshotRule, len(state.Rules)) out := make(map[string][]SnapshotRule, len(state.Rules))
for chain, rules := range state.Rules { for chain, rules := range state.Rules {
for _, r := range rules { for _, r := range rules {
sr := SnapshotRule{Tag: r.Tag} sr := SnapshotRule{Tag: r.Tag}
for _, e := range r.Exprs { for _, e := range r.Exprs {
b, err := expr.Marshal(inet, e) b, err := expr.Marshal(fam, e)
if err != nil { if err != nil {
return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err) return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err)
} }
@@ -58,13 +55,13 @@ func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) {
return out, nil return out, nil
} }
func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) { func decodeState(rules map[string][]SnapshotRule, fam byte) (*FirewallState, error) {
state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))} state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))}
for chain, rs := range rules { for chain, rs := range rules {
for _, sr := range rs { for _, sr := range rs {
r := ManagedRule{Chain: chain, Tag: sr.Tag} r := ManagedRule{Chain: chain, Tag: sr.Tag}
for _, b := range sr.Exprs { for _, b := range sr.Exprs {
e, err := decodeExpr(b) e, err := decodeExpr(b, fam)
if err != nil { if err != nil {
return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err) return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err)
} }
@@ -77,7 +74,7 @@ func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) {
} }
// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules. // decodeExpr reverses expr.Marshal, as google/nftables does when reading rules.
func decodeExpr(b []byte) (expr.Any, error) { func decodeExpr(b []byte, fam byte) (expr.Any, error) {
ad, err := netlink.NewAttributeDecoder(b) ad, err := netlink.NewAttributeDecoder(b)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -104,13 +101,13 @@ func decodeExpr(b []byte) (expr.Any, error) {
if name == "notrack" { if name == "notrack" {
return e, nil return e, nil
} }
if err := expr.Unmarshal(inet, data, e); err != nil { if err := expr.Unmarshal(fam, data, e); err != nil {
return nil, err return nil, err
} }
// A verdict is an immediate into the verdict register with no data. // A verdict is an immediate into the verdict register with no data.
if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 { if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 {
v := &expr.Verdict{} v := &expr.Verdict{}
if err := expr.Unmarshal(inet, data, v); err != nil { if err := expr.Unmarshal(fam, data, v); err != nil {
return nil, err return nil, err
} }
return v, nil return v, nil
+5 -5
View File
@@ -29,7 +29,7 @@ func TestSnapshotRulesRoundTrip(t *testing.T) {
"input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}}, "input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}},
}} }}
rules, err := encodeState(state) rules, err := encodeState(state, byte(nftables.TableFamilyINet))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -45,7 +45,7 @@ func TestSnapshotRulesRoundTrip(t *testing.T) {
if snap.Policies["input"] != nftables.ChainPolicyAccept { if snap.Policies["input"] != nftables.ChainPolicyAccept {
t.Errorf("policy lost: %v", snap.Policies) t.Errorf("policy lost: %v", snap.Policies)
} }
got, err := decodeState(snap.Rules) got, err := decodeState(snap.Rules, byte(nftables.TableFamilyINet))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -79,7 +79,7 @@ func TestSnapshotAndRestoreAbsentTable(t *testing.T) {
for _, m := range req { for _, m := range req {
sent = append(sent, m.Header.Type) sent = append(sent, m.Header.Type)
if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent { if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent {
data := []byte{inet, 0, 0, 0} data := []byte{byte(nftables.TableFamilyINet), 0, 0, 0}
attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}}) attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}})
return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil
} }
@@ -116,7 +116,7 @@ func TestRestorePresentTable(t *testing.T) {
{Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}}, {Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}},
{Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}, {Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}},
} { } {
enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}) enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}, byte(nftables.TableFamilyINet))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -132,7 +132,7 @@ func TestRestorePresentTable(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
return append([]byte{inet, 0, 0, 0}, b...) return append([]byte{byte(nftables.TableFamilyINet), 0, 0, 0}, b...)
} }
handle := make([]byte, 8) handle := make([]byte, 8)
binary.BigEndian.PutUint64(handle, 7) binary.BigEndian.PutUint64(handle, 7)
+36 -6
View File
@@ -2,6 +2,7 @@ package shorewall
import ( import (
"fmt" "fmt"
"log/slog"
"strconv" "strconv"
"strings" "strings"
@@ -99,6 +100,15 @@ func convertDir(dir string, ipv6 bool) (*config.Config, error) {
return cfg, nil return cfg, nil
} }
// disposition maps a shorewall *_DISPOSITION value; unset means CONTINUE and A_ (audit) variants map to their base action.
func disposition(v string) config.PolicyAction {
v = strings.TrimPrefix(strings.ToLower(v), "a_")
if v == "" {
return config.PolicyContinue
}
return config.PolicyAction(v)
}
func subst(s string, params map[string]string) string { func subst(s string, params map[string]string) string {
if !strings.Contains(s, "$") { if !strings.Contains(s, "$") {
return s return s
@@ -115,6 +125,8 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6
if err != nil { if err != nil {
return err return err
} }
cfg.Settings.InvalidDisposition = disposition(conf["INVALID_DISPOSITION"])
cfg.Settings.UntrackedDisposition = disposition(conf["UNTRACKED_DISPOSITION"])
if conf == nil { if conf == nil {
return nil return nil
} }
@@ -130,13 +142,23 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6
} else { } else {
cfg.Settings.LogLevel = "info" cfg.Settings.LogLevel = "info"
} }
if v := conf["LOGLIMIT"]; v != "" {
if strings.HasPrefix(v, "s:") || strings.HasPrefix(v, "d:") {
slog.Warn("shorewall: per-address LOGLIMIT is not supported, limiting each log site globally", "loglimit", v)
v = v[2:]
}
if name, rest, ok := strings.Cut(v, ":"); ok && !strings.Contains(name, "/") {
slog.Warn("shorewall: named LOGLIMIT is not supported, dropping the name", "loglimit", conf["LOGLIMIT"])
v = rest
}
cfg.Settings.LogLimit = v
}
if v, ok := conf["IP_FORWARDING"]; ok { if v, ok := conf["IP_FORWARDING"]; ok {
cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep" cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep"
} }
if v, ok := conf["IMPLICIT_CONTINUE"]; ok { if v, ok := conf["IMPLICIT_CONTINUE"]; ok {
cfg.Settings.ImplicitContinue = v == "Yes" cfg.Settings.ImplicitContinue = v == "Yes"
} }
return nil return nil
} }
@@ -226,6 +248,9 @@ func convertInterfaces(dir string, cfg *config.Config, params map[string]string)
for _, row := range rows { for _, row := range rows {
zone := subst(field(row, 0), params) zone := subst(field(row, 0), params)
iface := subst(field(row, 1), params) iface := subst(field(row, 1), params)
if isDash(zone) {
zone = ""
}
intf := config.Interface{ intf := config.Interface{
Zone: zone, Zone: zone,
@@ -370,12 +395,14 @@ func convertHosts(dir string, cfg *config.Config, params map[string]string) erro
zone := subst(field(row, 0), params) zone := subst(field(row, 0), params)
hostDef := subst(field(row, 1), params) hostDef := subst(field(row, 1), params)
hostDef, excl, _ := strings.Cut(hostDef, "!")
iface, addrs := splitHostDef(hostDef) iface, addrs := splitHostDef(hostDef)
host := config.Host{ host := config.Host{
Zone: zone, Zone: zone,
Interface: iface, Interface: iface,
Addresses: addrs, Addresses: addrs,
Exclusions: splitAddrList(excl),
} }
optsStr := subst(field(row, 2), params) optsStr := subst(field(row, 2), params)
@@ -393,16 +420,19 @@ func splitHostDef(s string) (string, []string) {
if idx < 0 { if idx < 0 {
return s, nil return s, nil
} }
iface := s[:idx] return s[:idx], splitAddrList(s[idx+1:])
addrPart := s[idx+1:] }
// splitAddrList splits a comma address list, unwrapping shorewall6 [addr]/len brackets.
func splitAddrList(s string) []string {
var addrs []string var addrs []string
for _, a := range strings.Split(addrPart, ",") { for _, a := range strings.Split(s, ",") {
a = strings.TrimSpace(a) a = strings.NewReplacer("[", "", "]", "").Replace(strings.TrimSpace(a))
if a != "" { if a != "" {
addrs = append(addrs, a) addrs = append(addrs, a)
} }
} }
return iface, addrs return addrs
} }
func parseHostOptions(s string) config.HostOptions { func parseHostOptions(s string) config.HostOptions {
+79
View File
@@ -3,6 +3,7 @@ package shorewall
import ( import (
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"testing" "testing"
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
@@ -725,3 +726,81 @@ func TestIsIPv6Dir(t *testing.T) {
} }
}) })
} }
func TestConvert_Dispositions(t *testing.T) {
cases := []struct {
conf string
invalid, untracked config.PolicyAction
}{
{"", config.PolicyContinue, config.PolicyContinue},
{"IP_FORWARDING=Yes", config.PolicyContinue, config.PolicyContinue},
{"INVALID_DISPOSITION=CONTINUE\nUNTRACKED_DISPOSITION=ACCEPT", config.PolicyContinue, config.PolicyAccept},
{"INVALID_DISPOSITION=DROP\nUNTRACKED_DISPOSITION=A_DROP", config.PolicyDrop, config.PolicyDrop},
{"INVALID_DISPOSITION=A_REJECT", config.PolicyReject, config.PolicyContinue},
}
for _, tc := range cases {
dir := t.TempDir()
writeFile(t, dir, "shorewall.conf", tc.conf)
writeFile(t, dir, "zones", "fw firewall\nnet ipv4\n")
writeFile(t, dir, "interfaces", "net eth0 -\n")
writeFile(t, dir, "policy", "all all DROP\n")
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("Convert(%q): %v", tc.conf, err)
}
if cfg.Settings.InvalidDisposition != tc.invalid || cfg.Settings.UntrackedDisposition != tc.untracked {
t.Errorf("%q: got invalid=%q untracked=%q, want %q/%q", tc.conf,
cfg.Settings.InvalidDisposition, cfg.Settings.UntrackedDisposition, tc.invalid, tc.untracked)
}
if err := cfg.Validate(); err != nil {
t.Errorf("%q: Validate: %v", tc.conf, err)
}
}
}
func TestConvert_LogLimit(t *testing.T) {
for in, want := range map[string]string{
`LOGLIMIT="s:1/sec:10"`: "1/sec:10",
`LOGLIMIT=2/min`: "2/min",
`LOGLIMIT=name:1/sec:5`: "1/sec:5",
`LOGLIMIT=s:name:1/sec:5`: "1/sec:5",
`LOGLIMIT=`: "",
} {
dir := minimalShorewallDir(t)
writeFile(t, dir, "shorewall.conf", "LOG_LEVEL=info\n"+in+"\n")
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("%s: %v", in, err)
}
if cfg.Settings.LogLimit != want {
t.Errorf("%s: log_limit = %q, want %q", in, cfg.Settings.LogLimit, want)
}
}
}
func TestConvert_HostsExclusions(t *testing.T) {
dir := minimalShorewallDir(t)
writeFile(t, dir, "interfaces", `
net eth0
- eth1
`)
writeFile(t, dir, "hosts", `
loc eth0:192.0.2.0/24,198.51.100.0/24!192.0.2.7,192.0.2.8 routeback
loc eth1:[2001:db8::]/64
`)
cfg, err := Convert(dir)
if err != nil {
t.Fatalf("Convert: %v", err)
}
want := []config.Host{
{Zone: "loc", Interface: "eth0", Addresses: []string{"192.0.2.0/24", "198.51.100.0/24"},
Exclusions: []string{"192.0.2.7", "192.0.2.8"}, Options: config.HostOptions{RouteBack: true}},
{Zone: "loc", Interface: "eth1", Addresses: []string{"2001:db8::/64"}},
}
if !reflect.DeepEqual(cfg.Hosts, want) {
t.Errorf("hosts = %+v, want %+v", cfg.Hosts, want)
}
if err := cfg.Validate(); err != nil {
t.Errorf("Validate: %v", err)
}
}
+61 -26
View File
@@ -25,15 +25,15 @@ const Unit = "tomswall-try-revert"
var ( var (
// Dir holds the lock and the pending snapshot. // Dir holds the lock and the pending snapshot.
Dir = "/var/lib/tomswall" Dir = "/var/lib/tomswall"
// run executes a systemd command; replaced in tests. // Run executes a systemd command. Test hook; production code must not reassign.
run = func(name string, args ...string) error { Run = func(name string, args ...string) error {
if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { if out, err := exec.Command(name, args...).CombinedOutput(); err != nil {
return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out)))
} }
return nil return nil
} }
// restore rolls the live table back to a snapshot; replaced in tests. // Restore rolls the live table back to a snapshot. Test hook; production code must not reassign.
restore = func(s *nftables.Snapshot) error { Restore = func(s *nftables.Snapshot) error {
engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}})
if err != nil { if err != nil {
return err return err
@@ -94,23 +94,7 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error)
if err != nil { if err != nil {
return "", err return "", err
} }
f, err := os.CreateTemp(Dir, ".try-snapshot-*") if err := WriteFile(snapshotPath(), b); err != nil {
if err != nil {
return "", err
}
defer os.Remove(f.Name())
if _, err := f.Write(b); err != nil {
f.Close()
return "", err
}
if err := f.Sync(); err != nil {
f.Close()
return "", err
}
if err := f.Close(); err != nil {
return "", err
}
if err := os.Rename(f.Name(), snapshotPath()); err != nil {
return "", err return "", err
} }
@@ -119,13 +103,46 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error)
return "", discardWith(err) return "", discardWith(err)
} }
_ = disarm() // a leftover timer from an earlier try would block the unit name _ = disarm() // a leftover timer from an earlier try would block the unit name
if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, if err := Run("systemd-run", "--quiet", "--collect", "--unit", Unit,
fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil { fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil {
return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) return "", discardWith(fmt.Errorf("arming revert timer: %w", err))
} }
return id, nil return id, nil
} }
// WriteFile durably replaces path with b: temp file, fsync, rename, fsync the directory.
func WriteFile(path string, b []byte) error {
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
f, err := os.CreateTemp(dir, "."+filepath.Base(path)+"-*")
if err != nil {
return err
}
defer os.Remove(f.Name())
if _, err := f.Write(b); err != nil {
f.Close()
return err
}
if err := f.Sync(); err != nil {
f.Close()
return err
}
if err := f.Close(); err != nil {
return err
}
if err := os.Rename(f.Name(), path); err != nil {
return err
}
d, err := os.Open(dir)
if err != nil {
return err
}
defer d.Close()
return d.Sync()
}
// Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. // Discard drops the pending snapshot and timer without restoring. The caller must hold the lock.
func Discard() error { func Discard() error {
_ = disarm() _ = disarm()
@@ -143,7 +160,7 @@ func discardWith(err error) error {
} }
func disarm() error { func disarm() error {
return run("systemctl", "stop", Unit+".timer") return Run("systemctl", "stop", Unit+".timer")
} }
// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e. // Confirm keeps the tried ruleset. ok is false when no try was pending, i.e.
@@ -175,10 +192,28 @@ func Revert(id string) (reverted bool, err error) {
if err != nil || p == nil || (id != "" && p.ID != id) { if err != nil || p == nil || (id != "" && p.ID != id) {
return false, err return false, err
} }
if err := restore(p.Snapshot); err != nil { return true, restorePending(p)
return false, fmt.Errorf("restoring snapshot: %w", err) }
// Abort restores the pending snapshot after a failed apply, which may have
// committed partially. A failed restore keeps the snapshot and timer so the
// timer still reverts. The caller must hold the lock.
func Abort() error {
p, err := load()
if err != nil {
return err
} }
return true, Discard() if p == nil {
return errors.New("no pending try to abort")
}
return restorePending(p)
}
func restorePending(p *pending) error {
if err := Restore(p.Snapshot); err != nil {
return fmt.Errorf("restoring snapshot: %w", err)
}
return Discard()
} }
func load() (*pending, error) { func load() (*pending, error) {
+45 -7
View File
@@ -15,12 +15,12 @@ func setup(t *testing.T) *[]string {
t.Helper() t.Helper()
Dir = t.TempDir() Dir = t.TempDir()
var cmds []string var cmds []string
orig := run orig := Run
run = func(name string, args ...string) error { Run = func(name string, args ...string) error {
cmds = append(cmds, name+" "+strings.Join(args, " ")) cmds = append(cmds, name+" "+strings.Join(args, " "))
return nil return nil
} }
t.Cleanup(func() { run = orig }) t.Cleanup(func() { Run = orig })
return &cmds return &cmds
} }
@@ -91,7 +91,7 @@ func TestAcquireRefusesWhilePending(t *testing.T) {
func TestArmFailureDiscardsSnapshot(t *testing.T) { func TestArmFailureDiscardsSnapshot(t *testing.T) {
setup(t) setup(t)
run = func(name string, args ...string) error { Run = func(name string, args ...string) error {
if name == "systemd-run" { if name == "systemd-run" {
return errors.New("no systemd") return errors.New("no systemd")
} }
@@ -145,12 +145,12 @@ func TestConfirmAfterRevertFails(t *testing.T) {
func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot { func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot {
t.Helper() t.Helper()
var got []*nftables.Snapshot var got []*nftables.Snapshot
orig := restore orig := Restore
restore = func(s *nftables.Snapshot) error { Restore = func(s *nftables.Snapshot) error {
got = append(got, s) got = append(got, s)
return err return err
} }
t.Cleanup(func() { restore = orig }) t.Cleanup(func() { Restore = orig })
return &got return &got
} }
@@ -209,3 +209,41 @@ func TestRevertStaleIDIgnored(t *testing.T) {
t.Errorf("newer try's snapshot removed: %v", err) t.Errorf("newer try's snapshot removed: %v", err)
} }
} }
func TestAbortRestoresAndDisarms(t *testing.T) {
cmds := setup(t)
restored := stubRestore(t, nil)
snap := &nftables.Snapshot{Table: "tomswall", Present: true}
arm(t, snap)
if err := Abort(); err != nil {
t.Fatal(err)
}
if len(*restored) != 1 || !reflect.DeepEqual((*restored)[0], snap) {
t.Errorf("restored %+v, want the armed snapshot", *restored)
}
if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) {
t.Error("snapshot not removed")
}
if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" {
t.Errorf("timer not stopped, last command %q", last)
}
}
func TestAbortFailureKeepsSnapshotAndTimer(t *testing.T) {
cmds := setup(t)
boom := errors.New("netlink down")
stubRestore(t, boom)
arm(t, &nftables.Snapshot{Table: "tomswall"})
armed := len(*cmds)
if err := Abort(); !errors.Is(err, boom) {
t.Fatalf("Abort error = %v, want %v", err, boom)
}
if _, err := os.Stat(snapshotPath()); err != nil {
t.Fatalf("snapshot gone after failed abort: %v", err)
}
if len(*cmds) != armed {
t.Errorf("revert timer touched after failed abort: %v", (*cmds)[armed:])
}
}
+11
View File
@@ -47,6 +47,17 @@ contents:
file_info: file_info:
mode: 0640 mode: 0640
# systemd unit + environment file for applying a local config at boot.
- src: packaging/tomswall.service
dst: /usr/lib/systemd/system/tomswall.service
file_info:
mode: 0644
- src: packaging/tomswall.env
dst: /etc/tomswall/tomswall.env
type: config|noreplace
file_info:
mode: 0644
# Shell completions (generated by scripts/build-rpm.sh before packaging). # Shell completions (generated by scripts/build-rpm.sh before packaging).
- src: dist/completions/tomswall.bash - src: dist/completions/tomswall.bash
dst: /usr/share/bash-completion/completions/tomswall dst: /usr/share/bash-completion/completions/tomswall
+1
View File
@@ -3,6 +3,7 @@ Description=tomswall control-plane agent (pull and apply firewall config)
Documentation=https://git.unkin.net/unkin/tomswall Documentation=https://git.unkin.net/unkin/tomswall
After=network-online.target After=network-online.target
Wants=network-online.target Wants=network-online.target
Conflicts=tomswall.service
[Service] [Service]
Type=simple Type=simple
+3
View File
@@ -0,0 +1,3 @@
# Config applied by tomswall.service: a tomswall YAML file or a shorewall directory.
TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
#TOMSWALL_CONFIG=/etc/shorewall
+25
View File
@@ -0,0 +1,25 @@
[Unit]
Description=tomswall firewall (apply local config at boot)
Documentation=https://git.unkin.net/unkin/tomswall
DefaultDependencies=no
Wants=network-pre.target
Before=network-pre.target shutdown.target
After=local-fs.target systemd-sysctl.service
Conflicts=shutdown.target tomswall-agent.service
StartLimitIntervalSec=60
StartLimitBurst=5
[Service]
Type=oneshot
RemainAfterExit=yes
Environment=TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
EnvironmentFile=-/etc/tomswall/tomswall.env
ExecStart=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
ExecReload=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
# Fails open: after StartLimitBurst failures within StartLimitIntervalSec, boot continues without the ruleset.
Restart=on-failure
RestartSec=5
# No ExecStop: stopping the unit leaves the ruleset in place (flush would open the firewall).
[Install]
WantedBy=sysinit.target
+6
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: