30 Commits

Author SHA1 Message Date
unkin-agent 78dbb6ad18 Merge remote-tracking branch 'origin/main' into benvin/nft-decodable
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
# Conflicts:
#	internal/nftables/compiler_test.go
2026-10-10 01:10:37 +11:00
benvin 9aee3ad7eb Merge pull request 'Exclude sub-zone hosts from wildcard parent interfaces' (#39) from benvin/wildcard-subzone-exclusion into main
Reviewed-on: #39
2026-10-10 01:08:33 +11:00
unkin-agent 9bfdf292ad Count limiter expansion after table-family filtering
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 23:52:02 +11:00
unkin-agent bc65312647 Cover per-family expansion across rules, NAT, tunnels and blrules
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 23:48:53 +11:00
unkin-agent 9a65973d7f Expand address matches per family, fail compile on conflicting guards 2026-10-09 23:48:53 +11:00
unkin-agent 460eb20db5 Carve sub-zone host interfaces out of wildcard parent matches
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:41:34 +11:00
unkin-agent 42a4dab6a3 Emit one family guard per rule, none in ip/ip6 tables
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:39:23 +11:00
unkin-agent 190ff72643 Exclude sub-zone hosts from wildcard parent interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 23:35:32 +11:00
benvin b8ad59b053 Merge pull request 'Match zones defined by hosts entries' (#38) from benvin/hosts-zones into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #38
2026-10-09 23:28:54 +11:00
unkin-agent 3174eabd94 Honour hosts routeback for same-interface intra-zone pairs
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 23:15:50 +11:00
unkin-agent 0b68110220 Merge remote-tracking branch 'origin/main' into benvin/hosts-zones
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
# Conflicts:
#	internal/nftables/compiler.go
#	internal/nftables/compiler_test.go
2026-10-09 23:12:36 +11:00
benvin 70df237121 Merge pull request 'Accept intra-zone traffic between different interfaces' (#37) from benvin/intrazone-multi-iface into main
Reviewed-on: #37
2026-10-09 23:10:48 +11:00
unkin-agent 799c7f3524 Guard zone host exclusions with the address family
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:54:07 +11:00
unkin-agent ecc349cb6f Skip fw->fw policies and treat dest-side + as intra-zone override
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:48:33 +11:00
unkin-agent 695869c80b Match zones defined by hosts entries
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-09 22:48:12 +11:00
unkin-agent 96a1ba8351 Accept intra-zone traffic between different interfaces
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-09 22:45:53 +11:00
benvin afa056b454 Merge pull request 'Add boot unit applying a local config' (#34) from benvin/boot-unit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #34
2026-10-05 21:46:58 +11:00
benvin f593f7625d Merge pull request 'Revert agent generations that cut off the control plane' (#35) from benvin/agent-safe-apply into main
Reviewed-on: #35
2026-10-05 21:46:08 +11:00
benvin 7c8bd87ec0 Merge pull request 'ci: use container-rpmbuilder image' (#36) from benvin/rpmbuilder-image into main
Reviewed-on: #36
2026-10-05 21:34:18 +11:00
unkin-agent 9854b0e7b6 ci: use container-rpmbuilder image
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 14:39:45 +11:00
unkin-agent c7e02c089a Apply the cached config without safe-apply and keep reverted generations in memory
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:56:11 +11:00
unkin-agent 9092b463a0 Apply agent generations as a pending try with a revert timer
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:52:32 +11:00
unkin-agent 502d06bdda Cap boot unit restarts so it fails open
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:50 +11:00
unkin-agent 4ad55fc65e Revert agent generations that cut off the control plane
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:46:20 +11:00
unkin-agent 7fcbb5fad8 Order boot unit before sysinit and retry on failure
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:44:58 +11:00
unkin-agent dc406c4f56 Add boot unit applying a local config
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-05 13:41:57 +11:00
benvin 2832dd0e7d Merge pull request 'Rate-limit log sites with LOGLIMIT' (#33) from benvin/loglimit into main
ci/woodpecker/tag/release Pipeline was successful
Reviewed-on: #33
2026-10-04 18:19:06 +11:00
unkin-agent 15ab32431d Drop the name from named shorewall LOGLIMIT values
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
2026-10-04 15:58:12 +11:00
unkin-agent be391ed385 Keep rule match extras on the limited log rule
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:56:23 +11:00
unkin-agent 2e8d51759d Rate-limit log sites with shorewall LOGLIMIT
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
2026-10-04 15:52:37 +11:00
23 changed files with 2346 additions and 265 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.`,
+231 -29
View File
@@ -2,8 +2,13 @@ package agent
import ( import (
"context" "context"
"encoding/json"
"errors"
"fmt" "fmt"
"log/slog" "log/slog"
"net/url"
"os"
"path/filepath"
"time" "time"
"git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/config"
@@ -12,11 +17,17 @@ import (
) )
// Applier applies a translated config to the firewall. Abstracted so the run // Applier applies a translated config to the firewall. Abstracted so the run
// loop is testable without touching the kernel. // loop is testable without touching the kernel. With safe, a change is applied
// as a pending try: revert restores the previous ruleset (a failed revert leaves
// the revert timer armed) and keep drops the snapshot. Both are nil when nothing
// changed or safe is false.
type Applier interface { type Applier interface {
Apply(ctx context.Context, cfg *config.Config) error Apply(ctx context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error)
} }
// revertDelay is when the revert timer fires if the agent dies mid-apply.
const revertDelay = time.Minute
// Agent runs the pull-apply-report loop for one device. // Agent runs the pull-apply-report loop for one device.
type Agent struct { type Agent struct {
Client *Client Client *Client
@@ -25,6 +36,9 @@ type Agent struct {
Applier Applier Applier Applier
// Resolver overrides the DNS resolver (tests); nil derives it per-config. // Resolver overrides the DNS resolver (tests); nil derives it per-config.
Resolver *Resolver Resolver *Resolver
// lastReverted covers a reverted generation whose persistence failed.
lastReverted *reverted
} }
// Run loops until ctx is cancelled, applying one cycle per Interval (and once // Run loops until ctx is cancelled, applying one cycle per Interval (and once
@@ -62,16 +76,31 @@ func (a *Agent) RunOnce(ctx context.Context) error {
return fmt.Errorf("control plane unreachable and no cached config: %w", err) return fmt.Errorf("control plane unreachable and no cached config: %w", err)
} }
// Re-apply last known-good; do not report a generation we didn't fetch. // Re-apply last known-good; do not report a generation we didn't fetch.
return a.applyConfig(ctx, cached, false) return a.applyConfig(ctx, cached, nil)
} }
if err := a.Cache.Write(raw); err != nil { rv, err := a.readReverted()
slog.Warn("agent: caching config failed", "err", err) if err != nil {
return err
} }
return a.applyConfig(ctx, rc, true) if a.lastReverted != nil && (rv == nil || a.lastReverted.Generation > rv.Generation) {
rv = a.lastReverted
}
if rv != nil {
a.reportReverted(ctx, rv)
if rc.Generation <= rv.Generation {
slog.Info("agent: generation was reverted, waiting for a newer one", "generation", rc.Generation)
return nil
}
}
return a.applyConfig(ctx, rc, raw)
} }
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error { // applyConfig applies rc. A fetched config (raw != nil) is applied as a pending
// try, verified by reaching the API through the new ruleset and reverted if that
// fails; only then is it cached. The cached config is the last verified-good one,
// so it is applied plainly: there is nothing to verify it against.
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error {
resolver := a.Resolver resolver := a.Resolver
if resolver == nil { if resolver == nil {
resolver = NewResolver(rc.Resolver) resolver = NewResolver(rc.Resolver)
@@ -82,46 +111,219 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
if err != nil { if err != nil {
return fmt.Errorf("translate: %w", err) return fmt.Errorf("translate: %w", err)
} }
if err := a.Applier.Apply(ctx, cfg); err != nil {
return fmt.Errorf("apply: %w", err) unlock, err := tryapply.Acquire()
if errors.Is(err, tryapply.ErrPending) {
slog.Warn("agent: a 'tomswall try' is pending, skipping cycle")
return nil
}
if err != nil {
return err
}
defer unlock()
revert, keep, err := a.Applier.Apply(ctx, cfg, raw != nil)
if err != nil {
err = fmt.Errorf("apply: %w", err)
if raw != nil && revert != nil {
return a.revertGeneration(ctx, rc.Generation, StatusFailed, err, revert)
}
if revert != nil {
if rerr := revert(); rerr != nil {
err = fmt.Errorf("%w; restore: %v; revert timer pending", err, rerr)
}
}
if raw != nil {
if rerr := a.Client.ReportStatus(ctx, Status{Status: StatusFailed, Generation: rc.Generation, Error: err.Error()}); rerr != nil {
slog.Warn("agent: reporting status failed", "err", rerr)
}
}
return err
}
if raw == nil {
slog.Info("agent: applied cached config", "generation", rc.Generation, "rules", len(cfg.Rules))
return nil
}
if err := a.confirm(ctx, rc.Generation); err != nil {
if keep != nil && ctx.Err() == nil {
return a.revertGeneration(ctx, rc.Generation, StatusReverted, err, revert)
}
// Shutdown is not a verdict on the generation: keep it.
if keep != nil {
if kerr := keep(); kerr != nil {
slog.Warn("agent: dropping snapshot failed", "err", kerr)
}
}
return err
}
if keep != nil {
if err := keep(); err != nil {
return fmt.Errorf("dropping snapshot: %w", err)
}
} }
slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules))
if report { if err := a.Cache.Write(raw); err != nil {
if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { slog.Warn("agent: caching config failed", "err", err)
slog.Warn("agent: reporting status failed", "err", err) }
} a.lastReverted = nil
// Report the FIB so the control plane can scope router enforcement. if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) {
if fib := CollectFIB(ctx); len(fib) > 0 { slog.Warn("agent: clearing reverted generation failed", "err", err)
if err := a.Client.ReportRoutes(ctx, fib); err != nil { }
slog.Warn("agent: reporting routes failed", "err", err) // Report the FIB so the control plane can scope router enforcement.
} if fib := CollectFIB(ctx); len(fib) > 0 {
if err := a.Client.ReportRoutes(ctx, fib); err != nil {
slog.Warn("agent: reporting routes failed", "err", err)
} }
} }
return nil return nil
} }
// EngineApplier applies via the real nftables differential engine. // revertGeneration marks generation as reverted before restoring, so a failed
type EngineApplier struct{} // restore can never lead to re-applying it, then reports status. It is also kept
// in memory in case persisting fails. A failed
// restore leaves the snapshot and timer armed and is reported as failed.
func (a *Agent) revertGeneration(ctx context.Context, generation int64, status string, cause error, revert func() error) error {
rv := &reverted{Generation: generation, Status: status, Error: cause.Error()}
a.lastReverted = rv
if err := a.writeReverted(rv); err != nil {
slog.Error("agent: persisting reverted generation failed", "err", err)
}
suffix := ""
if rerr := revert(); rerr != nil {
suffix = fmt.Sprintf("; restore: %v; revert timer pending", rerr)
rv.Status = StatusFailed
rv.Error += suffix
if werr := a.writeReverted(rv); werr != nil {
slog.Error("agent: persisting reverted generation failed", "err", werr)
}
}
a.reportReverted(ctx, rv)
return fmt.Errorf("generation %d %s: %w%s", generation, rv.Status, cause, suffix)
}
// Apply computes and applies the differential change set for cfg. It refuses var (
// while a 'tomswall try' awaits confirmation. verifyAttempts = 3
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { verifyDelay = 2 * time.Second
unlock, err := tryapply.Acquire() verifyTimeout = 5 * time.Second
)
// errUnreachable means the API could not be reached through the new ruleset.
var errUnreachable = errors.New("control plane unreachable after apply")
// confirm reports generation as applied over a fresh connection, which proves
// the API is reachable through the new ruleset. Any HTTP response counts as
// reachable; only repeated transport failures return errUnreachable.
func (a *Agent) confirm(ctx context.Context, generation int64) error {
var err error
for i := 0; i < verifyAttempts; i++ {
if i > 0 {
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(verifyDelay):
}
}
actx, cancel := context.WithTimeout(ctx, verifyTimeout)
err = a.Client.ReportStatus(actx, Status{Status: StatusApplied, Generation: generation})
cancel()
if ctx.Err() != nil {
return ctx.Err()
}
var uerr *url.Error
if !errors.As(err, &uerr) {
if err != nil {
slog.Warn("agent: reporting status failed", "err", err)
}
return nil
}
}
return fmt.Errorf("%w: %v", errUnreachable, err)
}
// reverted is a generation rolled back after a failed apply or for severing the
// API; persisted so it is not re-applied until a newer generation arrives.
type reverted struct {
Generation int64 `json:"generation"`
Status string `json:"status,omitempty"`
Error string `json:"error,omitempty"`
Reported bool `json:"reported"`
}
func (a *Agent) revertedPath() string {
return filepath.Join(filepath.Dir(a.Cache.Path), "reverted.json")
}
func (a *Agent) readReverted() (*reverted, error) {
b, err := os.ReadFile(a.revertedPath())
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
var rv reverted
if err := json.Unmarshal(b, &rv); err != nil {
return nil, fmt.Errorf("parsing %s: %w", a.revertedPath(), err)
}
return &rv, nil
}
func (a *Agent) writeReverted(rv *reverted) error {
b, err := json.Marshal(rv)
if err != nil { if err != nil {
return err return err
} }
defer unlock() return tryapply.WriteFile(a.revertedPath(), b)
}
// reportReverted reports rv until the control plane accepts it.
func (a *Agent) reportReverted(ctx context.Context, rv *reverted) {
if rv.Reported {
return
}
status := rv.Status
if status == "" {
status = StatusReverted
}
if err := a.Client.ReportStatus(ctx, Status{Status: status, Generation: rv.Generation, Error: rv.Error}); err != nil {
slog.Warn("agent: reporting reverted generation failed, retrying next cycle", "generation", rv.Generation, "err", err)
return
}
rv.Reported = true
if err := a.writeReverted(rv); err != nil {
slog.Warn("agent: persisting reverted generation failed", "err", err)
}
}
// EngineApplier applies via the real nftables differential engine.
type EngineApplier struct{}
// Apply computes and applies the differential change set for cfg, with safe
// under a pending try as 'tomswall try' does. The caller holds the try lock.
func (EngineApplier) Apply(_ context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error) {
engine, err := nftables.NewEngine(cfg) engine, err := nftables.NewEngine(cfg)
if err != nil { if err != nil {
return fmt.Errorf("initializing nftables: %w", err) return nil, nil, fmt.Errorf("initializing nftables: %w", err)
} }
changes, err := engine.Plan() changes, err := engine.Plan()
if err != nil { if err != nil {
return fmt.Errorf("computing changes: %w", err) return nil, nil, fmt.Errorf("computing changes: %w", err)
} }
if changes.Empty() { if changes.Empty() {
return nil return nil, nil, nil
} }
return engine.Apply(changes) if !safe {
return nil, nil, engine.Apply(changes)
}
snap, err := engine.Snapshot()
if err != nil {
return nil, nil, fmt.Errorf("snapshotting ruleset: %w", err)
}
// PID 0: 'tomswall confirm' must not signal the agent.
if _, err := tryapply.Arm(snap, 0, revertDelay); err != nil {
return nil, nil, err
}
return tryapply.Abort, tryapply.Discard, engine.Apply(changes)
} }
+2 -2
View File
@@ -132,10 +132,10 @@ type fakeApplier struct {
lastGen int lastGen int
} }
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error { func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config, _ bool) (func() error, func() error, error) {
atomic.AddInt32(&f.count, 1) atomic.AddInt32(&f.count, 1)
f.lastGen = len(cfg.Rules) f.lastGen = len(cfg.Rules)
return nil return nil, nil, nil
} }
const renderedYAML = `generation: 7 const renderedYAML = `generation: 7
+4 -10
View File
@@ -2,7 +2,8 @@ package agent
import ( import (
"os" "os"
"path/filepath"
"git.unkin.net/unkin/tomswall/internal/tryapply"
) )
// Cache persists the last known-good rendered config to disk so the agent can // Cache persists the last known-good rendered config to disk so the agent can
@@ -11,16 +12,9 @@ type Cache struct {
Path string Path string
} }
// Write atomically stores the raw config bytes. // Write durably stores the raw config bytes.
func (c Cache) Write(raw []byte) error { func (c Cache) Write(raw []byte) error {
if err := os.MkdirAll(filepath.Dir(c.Path), 0o755); err != nil { return tryapply.WriteFile(c.Path, raw)
return err
}
tmp := c.Path + ".tmp"
if err := os.WriteFile(tmp, raw, 0o600); err != nil {
return err
}
return os.Rename(tmp, c.Path)
} }
// Read returns the cached config, or (nil, nil) when no cache exists yet. // Read returns the cached config, or (nil, nil) when no cache exists yet.
+26 -4
View File
@@ -26,7 +26,9 @@ func NewClient(baseURL, device, token string) *Client {
BaseURL: baseURL, BaseURL: baseURL,
Device: device, Device: device,
Token: token, Token: token,
HTTP: &http.Client{Timeout: 30 * time.Second}, // No keep-alives: every request, the post-apply check included, opens a
// fresh connection that must pass the current ruleset.
HTTP: &http.Client{Timeout: 30 * time.Second, Transport: noKeepAlive()},
} }
} }
@@ -94,10 +96,24 @@ func (c *Client) ReportRoutes(ctx context.Context, prefixes []string) error {
return nil return nil
} }
// ReportStatus tells the control plane which generation this device has applied. // Status values reported to POST /api/v1/devices/{name}/status.
func (c *Client) ReportStatus(ctx context.Context, generation int64) error { const (
StatusApplied = "applied"
StatusReverted = "reverted"
StatusFailed = "failed"
)
// Status is the outcome of applying one generation.
type Status struct {
Status string `json:"status"`
Generation int64 `json:"generation"`
Error string `json:"error,omitempty"`
}
// ReportStatus tells the control plane the outcome of applying a generation.
func (c *Client) ReportStatus(ctx context.Context, st Status) error {
url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device) url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device)
payload, _ := json.Marshal(map[string]int64{"generation": generation}) payload, _ := json.Marshal(st)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload)) req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
if err != nil { if err != nil {
return err return err
@@ -116,3 +132,9 @@ func (c *Client) ReportStatus(ctx context.Context, generation int64) error {
} }
return nil return nil
} }
func noKeepAlive() http.RoundTripper {
t := http.DefaultTransport.(*http.Transport).Clone()
t.DisableKeepAlives = true
return t
}
+396
View File
@@ -0,0 +1,396 @@
package agent
import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"git.unkin.net/unkin/tomswall/internal/config"
"git.unkin.net/unkin/tomswall/internal/nftables"
"git.unkin.net/unkin/tomswall/internal/tryapply"
)
func TestMain(m *testing.M) {
dir, err := os.MkdirTemp("", "tomswall-agent-test")
if err != nil {
panic(err)
}
tryapply.Dir = dir
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
return nil
}
verifyDelay = time.Millisecond
verifyTimeout = time.Second
code := m.Run()
os.RemoveAll(dir)
os.Exit(code)
}
// fakeAPI serves a config generation and records status reports; while cut it
// drops connections to the status endpoint, as a severing ruleset would.
type fakeAPI struct {
*httptest.Server
gen atomic.Int64
cut atomic.Bool
code atomic.Int32
mu sync.Mutex
reports []Status
}
func newFakeAPI(t *testing.T, gen int64) *fakeAPI {
f := &fakeAPI{}
f.gen.Store(gen)
f.code.Store(http.StatusNoContent)
f.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v1/devices/fw-a/config":
_, _ = w.Write([]byte(strings.Replace(renderedYAML, "generation: 7", "generation: "+itoa(f.gen.Load()), 1)))
case "/api/v1/devices/fw-a/status":
if f.cut.Load() {
conn, _, _ := w.(http.Hijacker).Hijack()
conn.Close()
return
}
var st Status
_ = json.NewDecoder(r.Body).Decode(&st)
f.mu.Lock()
f.reports = append(f.reports, st)
f.mu.Unlock()
w.WriteHeader(int(f.code.Load()))
default:
w.WriteHeader(http.StatusNotFound)
}
}))
t.Cleanup(f.Close)
return f
}
// timerCmds records the systemd commands tryapply runs.
var timerCmds []string
func itoa(n int64) string { b, _ := json.Marshal(n); return string(b) }
func (f *fakeAPI) last() Status {
f.mu.Lock()
defer f.mu.Unlock()
if len(f.reports) == 0 {
return Status{}
}
return f.reports[len(f.reports)-1]
}
// fakeEngine always changes the ruleset, when safe under a real tryapply pending
// try; onApply simulates its effect and restoreErr fails the restore.
type fakeEngine struct {
applies, plain, restores int
err, restoreErr error
onApply func()
onRestore func()
}
func (f *fakeEngine) Apply(_ context.Context, _ *config.Config, safe bool) (func() error, func() error, error) {
if !safe {
f.plain++
return nil, nil, f.err
}
if _, err := tryapply.Arm(&nftables.Snapshot{Table: "tomswall"}, 0, time.Minute); err != nil {
return nil, nil, err
}
tryapply.Restore = func(*nftables.Snapshot) error {
f.restores++
if f.onRestore != nil {
f.onRestore()
}
return f.restoreErr
}
f.applies++
if f.onApply != nil {
f.onApply()
}
return tryapply.Abort, tryapply.Discard, f.err
}
// pending reports whether a snapshot is still armed and its timer not stopped since.
func pending(t *testing.T) bool {
t.Helper()
_, err := os.Stat(filepath.Join(tryapply.Dir, "try-snapshot.json"))
armed := len(timerCmds) > 0 && timerCmds[len(timerCmds)-1] == "systemd-run"
if (err == nil) != armed {
t.Fatalf("snapshot present=%v but timer armed=%v", err == nil, armed)
}
return armed
}
func newAgent(t *testing.T, api *fakeAPI, eng *fakeEngine) *Agent {
return &Agent{
Client: NewClient(api.URL, "fw-a", "tok"),
Cache: Cache{Path: filepath.Join(t.TempDir(), "rendered.yaml")},
Applier: eng,
}
}
func cachedGen(t *testing.T, a *Agent) int64 {
rc, err := a.Cache.Read()
if err != nil {
t.Fatal(err)
}
if rc == nil {
return 0
}
return rc.Generation
}
func TestSafeApplyReachableApplies(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.restores != 0 || api.last() != (Status{Status: StatusApplied, Generation: 7}) || cachedGen(t, a) != 7 || pending(t) {
t.Fatalf("restores=%d last=%+v cache=%d", eng.restores, api.last(), cachedGen(t, a))
}
}
func TestSafeApplyUnreachableRevertsAndReports(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, onRestore: func() { api.cut.Store(false) }}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
t.Fatalf("want errUnreachable, got %v", err)
}
if eng.restores != 1 || cachedGen(t, a) != 0 || pending(t) {
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
}
if st := api.last(); st.Status != StatusReverted || st.Generation != 7 || st.Error == "" {
t.Fatalf("last report %+v", st)
}
rv, _ := a.readReverted()
if rv == nil || rv.Generation != 7 || !rv.Reported {
t.Fatalf("persisted %+v", rv)
}
}
func TestSafeApplyRevertReportedOnceReachable(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }}
a := newAgent(t, api, eng)
_ = a.RunOnce(context.Background())
if eng.restores != 1 || api.last().Status != "" {
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
}
api.cut.Store(false)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 1 || api.last() != (Status{Status: StatusReverted, Generation: 7, Error: api.last().Error}) {
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
}
}
func TestSafeApplyShutdownDoesNotRevert(t *testing.T) {
api := newFakeAPI(t, 7)
ctx, cancel := context.WithCancel(context.Background())
eng := &fakeEngine{onApply: func() { api.cut.Store(true); cancel() }}
a := newAgent(t, api, eng)
if err := a.RunOnce(ctx); !errors.Is(err, context.Canceled) {
t.Fatalf("want context.Canceled, got %v", err)
}
if rv, _ := a.readReverted(); eng.restores != 0 || rv != nil || cachedGen(t, a) != 0 {
t.Fatalf("restores=%d reverted=%+v cache=%d", eng.restores, rv, cachedGen(t, a))
}
}
func TestSafeApplyHTTPErrorDoesNotRevert(t *testing.T) {
api := newFakeAPI(t, 7)
api.code.Store(http.StatusInternalServerError)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.restores != 0 || cachedGen(t, a) != 7 {
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
}
}
func TestSafeApplyApplyErrorRestoresAndReportsFailed(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{err: errors.New("netlink: boom")}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil {
t.Fatal("want error")
}
if st := api.last(); eng.restores != 1 || st.Status != StatusFailed || !strings.Contains(st.Error, "boom") || pending(t) {
t.Fatalf("restores=%d last=%+v", eng.restores, st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported {
t.Fatalf("persisted %+v", rv)
}
}
func TestSafeApplyApplyErrorRestoreFailsKeepsTimer(t *testing.T) {
t.Cleanup(func() { _ = tryapply.Discard() })
api := newFakeAPI(t, 7)
eng := &fakeEngine{err: errors.New("netlink: boom"), restoreErr: errors.New("netlink: stuck")}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "restore: restoring snapshot: netlink: stuck") {
t.Fatalf("got %v", err)
}
want := "apply: netlink: boom; restore: restoring snapshot: netlink: stuck; revert timer pending"
if st := api.last(); st != (Status{Status: StatusFailed, Generation: 7, Error: want}) || !pending(t) {
t.Fatalf("last=%+v", st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || rv.Status != StatusFailed {
t.Fatalf("persisted %+v", rv)
}
// The next cycle waits for the timer instead of re-applying.
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
t.Fatalf("err=%v applies=%d", err, eng.applies)
}
}
func TestSafeApplyUnreachableRestoreFailsKeepsTimer(t *testing.T) {
t.Cleanup(func() { _ = tryapply.Discard() })
api := newFakeAPI(t, 7)
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, restoreErr: errors.New("netlink: stuck")}
eng.onRestore = func() { api.cut.Store(false) }
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) || !strings.Contains(err.Error(), "revert timer pending") {
t.Fatalf("got %v", err)
}
st := api.last()
if st.Status != StatusFailed || st.Generation != 7 || !strings.HasPrefix(st.Error, errUnreachable.Error()) ||
!strings.HasSuffix(st.Error, "; restore: restoring snapshot: netlink: stuck; revert timer pending") || !pending(t) {
t.Fatalf("last=%+v", st)
}
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported || cachedGen(t, a) != 0 {
t.Fatalf("persisted %+v cache=%d", rv, cachedGen(t, a))
}
}
func TestSafeApplyRevertedGenerationSkippedAfterRestart(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.writeReverted(&reverted{Generation: 7, Reported: true}); err != nil {
t.Fatal(err)
}
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 0 {
t.Fatalf("reverted generation re-applied")
}
api.gen.Store(8)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if rv, _ := a.readReverted(); eng.applies != 1 || api.last().Generation != 8 || rv != nil {
t.Fatalf("applies=%d last=%+v reverted=%+v", eng.applies, api.last(), rv)
}
}
func TestSafeApplySkipsWhileTryPending(t *testing.T) {
marker := filepath.Join(tryapply.Dir, "try-snapshot.json")
if err := os.WriteFile(marker, []byte("{}"), 0o600); err != nil {
t.Fatal(err)
}
defer os.Remove(marker)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.applies != 0 || api.last().Status != "" {
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
}
}
// failArm makes arming the revert timer fail, as without systemd.
func failArm(t *testing.T) {
orig := tryapply.Run
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
if name == "systemd-run" {
return errors.New("no systemd")
}
return nil
}
t.Cleanup(func() { tryapply.Run = orig })
}
func TestSafeApplyArmFailureReportsFailedAndRetries(t *testing.T) {
failArm(t)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "no systemd") {
t.Fatalf("got %v", err)
}
if st := api.last(); eng.applies != 0 || st.Status != StatusFailed || st.Generation != 7 || pending(t) || cachedGen(t, a) != 0 {
t.Fatalf("applies=%d last=%+v", eng.applies, st)
}
if rv, _ := a.readReverted(); rv != nil || a.lastReverted != nil {
t.Fatalf("arm failure marked generation reverted: %+v", rv)
}
tryapply.Run = func(name string, args ...string) error {
timerCmds = append(timerCmds, name)
return nil
}
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 || cachedGen(t, a) != 7 {
t.Fatalf("retry err=%v applies=%d", err, eng.applies)
}
}
func TestCachedConfigAppliesWithoutArm(t *testing.T) {
failArm(t)
api := newFakeAPI(t, 7)
eng := &fakeEngine{}
a := newAgent(t, api, eng)
if err := a.Cache.Write([]byte(renderedYAML)); err != nil {
t.Fatal(err)
}
api.Close()
if err := a.RunOnce(context.Background()); err != nil {
t.Fatal(err)
}
if eng.plain != 1 || eng.applies != 0 || pending(t) {
t.Fatalf("plain=%d safe=%d", eng.plain, eng.applies)
}
}
func TestSafeApplyRevertedKeptInMemoryWhenPersistFails(t *testing.T) {
api := newFakeAPI(t, 7)
eng := &fakeEngine{onRestore: func() { api.cut.Store(false) }}
a := newAgent(t, api, eng)
// A non-empty directory in its place makes persisting reverted.json fail.
eng.onApply = func() {
api.cut.Store(true)
_ = os.MkdirAll(filepath.Join(a.revertedPath(), "x"), 0o755)
}
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
t.Fatalf("want errUnreachable, got %v", err)
}
if err := os.RemoveAll(a.revertedPath()); err != nil {
t.Fatal(err)
}
if eng.restores != 1 || api.last().Status != StatusReverted || pending(t) {
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
}
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
t.Fatalf("reverted generation re-applied: err=%v applies=%d", err, eng.applies)
}
}
+9 -1
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"regexp"
"strings" "strings"
"gopkg.in/yaml.v3" "gopkg.in/yaml.v3"
@@ -55,7 +56,9 @@ type Settings struct {
AddressFamily AddressFamily `yaml:"address_family,omitempty"` AddressFamily AddressFamily `yaml:"address_family,omitempty"`
IPForwarding bool `yaml:"ip_forwarding"` IPForwarding bool `yaml:"ip_forwarding"`
LogLevel string `yaml:"log_level"` LogLevel string `yaml:"log_level"`
TableName string `yaml:"table_name"` // LogLimit rate-limits every log site, shorewall LOGLIMIT syntax rate/unit[:burst]; unset logs every hit.
LogLimit string `yaml:"log_limit,omitempty"`
TableName string `yaml:"table_name"`
// When true, auto-generate CONTINUE policies for sub-zones to their parent zones. // When true, auto-generate CONTINUE policies for sub-zones to their parent zones.
ImplicitContinue bool `yaml:"implicit_continue,omitempty"` ImplicitContinue bool `yaml:"implicit_continue,omitempty"`
@@ -112,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,
} }
@@ -120,6 +125,9 @@ 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{ for name, d := range map[string]PolicyAction{
"invalid_disposition": c.Settings.InvalidDisposition, "invalid_disposition": c.Settings.InvalidDisposition,
"untracked_disposition": c.Settings.UntrackedDisposition, "untracked_disposition": c.Settings.UntrackedDisposition,
+33
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 {
@@ -1069,3 +1090,15 @@ func TestValidateDispositions(t *testing.T) {
checkErr(t, c.Validate(), tc.wantErr) 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
} }
+533 -103
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"net" "net"
"net/netip"
"slices" "slices"
"sort" "sort"
"strconv" "strconv"
@@ -65,10 +66,136 @@ 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)
if err := familyGuards(state, c.cfg.Settings.AddressFamily); err != nil {
return nil, err
}
return state, nil return state, nil
} }
// familyGuards leaves each rule at most one meta nfproto guard, ahead of its first network-header
// payload as nft emits it; nft list cannot decode repeated or conflicting guards. Rules are built per
// family, so a conflict is a compiler bug and fails the compile. An ip/ip6 table is its own guard:
// guards go and other-family rules are dropped.
func familyGuards(state *FirewallState, family config.AddressFamily) error {
table := tableFamily(family)
for chain, rules := range state.Rules {
var out []ManagedRule
for _, r := range rules {
fam, ok := guardFamily(r.Exprs)
if !ok {
return fmt.Errorf("%s %s: conflicting address families", chain, r.Tag)
}
if table != 0 && fam != 0 && fam != table {
continue
}
r.Exprs = normalizeGuards(r.Exprs, fam, table != 0)
out = append(out, r)
}
state.Rules[chain] = out
}
return nil
}
// tableFamily is the NFPROTO an ip/ip6 table is restricted to, 0 for inet.
func tableFamily(family config.AddressFamily) byte {
return map[config.AddressFamily]byte{config.FamilyIP: unix.NFPROTO_IPV4, config.FamilyIP6: unix.NFPROTO_IPV6}[family]
}
// guardFamily is the family exprs' nfproto guards require (0: none); ok is false when they conflict.
func guardFamily(exprs []expr.Any) (fam byte, ok bool) {
for i := range exprs {
if p := nfprotoGuard(exprs, i); p != 0 {
if fam != 0 && p != fam {
return 0, false
}
fam = p
}
}
return fam, true
}
// famsAgree reports whether families (0: any) can all hold for one packet.
func famsAgree(fams ...byte) bool {
var f byte
for _, p := range fams {
if p != 0 && f != 0 && p != f {
return false
}
if p != 0 {
f = p
}
}
return true
}
func normalizeGuards(in []expr.Any, fam byte, strip bool) []expr.Any {
at := -1
out := make([]expr.Any, 0, len(in))
for i := 0; i < len(in); i++ {
if nfprotoGuard(in, i) != 0 {
if at < 0 {
at = len(out)
}
i++
continue
}
out = append(out, in[i])
}
if fam == 0 || strip {
return out
}
if l3 := slices.IndexFunc(out, func(e expr.Any) bool {
p, ok := e.(*expr.Payload)
return ok && p.Base == expr.PayloadBaseNetworkHeader
}); l3 >= 0 && l3 < at {
at = l3
}
return slices.Insert(out, at, matchNFProto(fam)...)
}
// nfprotoGuard is the family a meta nfproto == match at in[i] guards for, else 0.
func nfprotoGuard(in []expr.Any, i int) byte {
m, ok := in[i].(*expr.Meta)
if !ok || m.Key != expr.MetaKeyNFPROTO || i+1 >= len(in) {
return 0
}
c, ok := in[i+1].(*expr.Cmp)
if !ok || c.Op != expr.CmpOpEq || c.Register != m.Register || len(c.Data) != 1 {
return 0
}
return c.Data[0]
}
// 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 invalid := c.cfg.Settings.InvalidDisposition
if invalid == "" { if invalid == "" {
@@ -313,15 +440,15 @@ func (c *Compiler) compileConntrack(state *FirewallState) error {
func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule, func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule,
srcZone, srcAddr, dstZone, dstAddr string) error { srcZone, srcAddr, dstZone, dstAddr string) error {
if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" && if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" &&
(dstAddr == "" || strings.HasPrefix(dstAddr, "!")) { !narrows(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 {
@@ -335,6 +462,9 @@ func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string,
} }
for _, m := range matches { for _, m := range matches {
exprs := m.exprs exprs := m.exprs
if _, ok := guardFamily(exprs); !ok {
continue
}
switch ct.Action { switch ct.Action {
case config.ConntrackNoTrack: case config.ConntrackNoTrack:
exprs = append(exprs, &expr.Notrack{}) exprs = append(exprs, &expr.Notrack{})
@@ -378,7 +508,7 @@ func (c *Compiler) compileRules(state *FirewallState) error {
return fmt.Errorf("rule[%d]: %w", i, err) return fmt.Errorf("rule[%d]: %w", i, err)
} }
if len(matches)*c.specCount(rule.Source, rule.Dest, rule.OrigDest, fwZone, rule.Action) > 1 { if len(matches)*c.specCount(rule.Source, rule.Dest, rule.OrigDest, fwZone, rule.Action) > 1 {
return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists (each expanded rule would get its own limiter)", i) return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists, or with a negated address in an inet table (it expands to one rule per family); each expanded rule would get its own limiter", i)
} }
} }
@@ -414,24 +544,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)...)
} }
@@ -439,21 +569,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)
} else { case *expr.Log:
nonVerdict = append(nonVerdict, e) post = append(post, e)
default:
if len(post) > 0 {
post = append(post, e)
} else {
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
@@ -550,19 +686,39 @@ func (c *Compiler) compileDNATAccept(state *FirewallState, tag, srcZone, srcAddr
return nil return nil
} }
// specCount is how many zone/address combinations compileOneRule expands src and dst into. // specCount is how many rules compileOneRule emits for src and dst once familyGuards has dropped
// cross-family address combinations and those outside an ip/ip6 table's family.
func (c *Compiler) specCount(srcSpec, dstSpec, origDest, fwZone string, action config.RuleAction) int { func (c *Compiler) specCount(srcSpec, dstSpec, origDest, fwZone string, action config.RuleAction) int {
table := tableFamily(c.cfg.Settings.AddressFamily)
count := func(addrs ...string) int {
n := 0
var walk func(i int, fams []byte)
walk = func(i int, fams []byte) {
if !famsAgree(fams...) {
return
}
if i == len(addrs) {
n++
return
}
for _, a := range splitAddrs(addrs[i]) {
walk(i+1, append(fams, addrFamily(a)))
}
}
walk(0, []byte{table})
return n
}
n := 0 n := 0
if action == config.RuleDNAT || action == config.RuleRedirect { if action == config.RuleDNAT || action == config.RuleRedirect {
for _, src := range c.dnatSourceSpecs(srcSpec, fwZone) { for _, src := range c.dnatSourceSpecs(srcSpec, fwZone) {
n += len(splitAddrs(src.Addr)) n += count(src.Addr, origDest)
} }
return n * len(splitAddrs(origDest)) return n
} }
for _, p := range c.zonePairs(srcSpec, dstSpec, fwZone) { for _, p := range c.zonePairs(srcSpec, dstSpec, fwZone) {
n += len(splitAddrs(p[0].Addr)) * len(splitAddrs(p[1].Addr)) n += count(p[0].Addr, p[1].Addr, origDest)
} }
return n * len(splitAddrs(origDest)) return n
} }
// zonePairs is the src/dst zone expansion of a non-DNAT rule, with fw added beside all/any. // zonePairs is the src/dst zone expansion of a non-DNAT rule, with fw added beside all/any.
@@ -628,19 +784,40 @@ func isZoneExclusion(spec string) bool {
return ok && (base == "all" || base == "any") return ok && (base == "all" || base == "any")
} }
// splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. // splitAddrs yields one alternative per listed address. A negated list ("everything except") becomes
// one AND-ed match per family: the family's negations, or the bare family (a /0) when it has none.
func splitAddrs(addr string) []string { func splitAddrs(addr string) []string {
if addr == "" || strings.HasPrefix(addr, "!") { if addr == "" {
return []string{addr} return []string{""}
} }
return strings.Split(addr, ",") if !strings.HasPrefix(addr, "!") {
return strings.Split(addr, ",")
}
v4, v6 := splitFamily(strings.Split(addr[1:], ","))
out := []string{"0.0.0.0/0", "::/0"}
if len(v4) > 0 {
out[0] = "!" + strings.Join(v4, ",")
}
if len(v6) > 0 {
out[1] = "!" + strings.Join(v6, ",")
}
return out
}
// narrows reports whether a splitAddrs alternative restricts addresses within its family.
func narrows(addr string) bool {
if addr == "" || strings.HasPrefix(addr, "!") {
return false
}
p, err := parsePrefix(addr)
return err != nil || p.Bits() > 0
} }
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" {
@@ -654,7 +831,7 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr,
return err return err
} }
if origDest != "" { if origDest != "" {
od, err := matchOrigDest(origDest) od, err := matchDestCIDR(origDest)
if err != nil { if err != nil {
return fmt.Errorf("origdest: %w", err) return fmt.Errorf("origdest: %w", err)
} }
@@ -665,6 +842,9 @@ func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr,
for _, m := range matches { for _, m := range matches {
exprs := m.exprs exprs := m.exprs
if _, ok := guardFamily(exprs); !ok {
continue
}
if section != "" && section != config.SectionAll { if section != "" && section != config.SectionAll {
exprs = append(exprs, matchSection(section)...) exprs = append(exprs, matchSection(section)...)
@@ -709,12 +889,12 @@ 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 != "" {
var err error var err error
if odExprs, err = matchOrigDest(origDest); err != nil { if odExprs, err = matchDestCIDR(origDest); err != nil {
return fmt.Errorf("origdest: %w", err) return fmt.Errorf("origdest: %w", err)
} }
} }
@@ -728,14 +908,18 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
if ip == nil { if ip == nil {
return fmt.Errorf("invalid DNAT address %q", dnatAddr) return fmt.Errorf("invalid DNAT address %q", dnatAddr)
} }
natFam := byte(0)
if action != config.RuleRedirect {
natFam = addrFamily(dnatAddr)
}
for _, srcIface := range srcIfaces { for _, srcIface := range srcIfaces {
zm, err := zoneMatchExprs(srcIface, true)
if err != nil {
return err
}
for _, m := range matches { for _, m := range matches {
var exprs []expr.Any exprs := slices.Clone(zm)
if srcIface != "" {
exprs = append(exprs, matchIfaceName(true, srcIface)...)
}
if srcAddr != "" { if srcAddr != "" {
src, err := matchSourceCIDR(srcAddr) src, err := matchSourceCIDR(srcAddr)
@@ -745,6 +929,9 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr,
exprs = append(exprs, src...) exprs = append(exprs, src...)
} }
exprs = append(exprs, odExprs...) exprs = append(exprs, odExprs...)
if f, ok := guardFamily(exprs); !ok || !famsAgree(f, natFam) {
continue
}
exprs = append(exprs, m.exprs...) exprs = append(exprs, m.exprs...)
@@ -822,31 +1009,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 {
continue if sz == fwZone || (!explicitIntra && !strings.HasSuffix(pol.Source, "+") && !strings.HasSuffix(pol.Dest, "+")) {
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) || chain != "input" && !famsAgree(matchFamily(si), matchFamily(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 != "" {
@@ -874,32 +1065,91 @@ 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)
var exprs []expr.Any
destIface, _ := splitZoneSpec(snat.Dest) destIface, _ := splitZoneSpec(snat.Dest)
exprs = append(exprs, matchIfaceName(false, destIface)...) var heads [][]expr.Any
for _, src := range splitAddrs(snat.Source) {
if snat.Source != "" { head := matchIfaceName(false, destIface)
srcExprs, err := matchSourceCIDR(snat.Source) if src != "" {
if err != nil { srcExprs, err := matchSourceCIDR(src)
return fmt.Errorf("snat[%d]: %w", i, err) if err != nil {
return fmt.Errorf("snat[%d]: %w", i, err)
}
head = append(head, srcExprs...)
}
if snat.Action != config.SNATAddress || famsAgree(addrFamily(src), addrFamily(snat.Address)) {
heads = append(heads, head)
} }
exprs = append(exprs, srcExprs...)
} }
matches, err := l4Matches(snat.Proto, snat.DPort, snat.SPort) matches, err := l4Matches(snat.Proto, snat.DPort, snat.SPort)
if err != nil { if err != nil {
return fmt.Errorf("snat[%d]: %w", i, err) return fmt.Errorf("snat[%d]: %w", i, err)
} }
head := exprs var exprs []expr.Any
exprs = nil
if snat.Mark != "" { if snat.Mark != "" {
exprs = append(exprs, matchMark(snat.Mark)...) exprs = append(exprs, matchMark(snat.Mark)...)
@@ -945,12 +1195,14 @@ func (c *Compiler) compileSNAT(state *FirewallState) error {
} }
} }
for _, m := range matches { for _, head := range heads {
state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ for _, m := range matches {
Chain: "postrouting", state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{
Exprs: append(append(append([]expr.Any{}, head...), m.exprs...), exprs...), Chain: "postrouting",
Tag: tag, Exprs: slices.Concat(head, m.exprs, exprs),
}) Tag: tag,
})
}
} }
} }
@@ -1178,11 +1430,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 {
@@ -1190,24 +1456,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})
}
}
}
} }
if addr != "" && !strings.HasPrefix(addr, "!") { for _, h := range c.cfg.Hosts {
return []string{""} 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 narrows(addr) {
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
@@ -1237,14 +1667,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 != "" {
@@ -1465,6 +1891,7 @@ func matchTCPFlags(flags, mask byte) []expr.Any {
func matchSmurfDrop(iface string) []expr.Any { func matchSmurfDrop(iface string) []expr.Any {
var exprs []expr.Any var exprs []expr.Any
exprs = append(exprs, matchIfaceName(true, iface)...) exprs = append(exprs, matchIfaceName(true, iface)...)
exprs = append(exprs, matchNFProto(unix.NFPROTO_IPV4)...)
exprs = append(exprs, exprs = append(exprs,
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Bitwise{ &expr.Bitwise{
@@ -1625,32 +2052,35 @@ func parseSPortOrRange(s string) ([]expr.Any, error) {
} }
func matchSourceCIDR(cidr string) ([]expr.Any, error) { func matchSourceCIDR(cidr string) ([]expr.Any, error) {
return matchAddrCIDR(cidr, true) return matchGuardedCIDR(cidr, true)
} }
func matchDestCIDR(cidr string) ([]expr.Any, error) { func matchDestCIDR(cidr string) ([]expr.Any, error) {
return matchAddrCIDR(cidr, false) return matchGuardedCIDR(cidr, false)
} }
// matchOrigDest guards the daddr match with the address's nfproto so it is family-correct in the inet table. // matchGuardedCIDR guards a single-family splitAddrs alternative with its nfproto; a /0 is the guard alone.
func matchOrigDest(addr string) ([]expr.Any, error) { func matchGuardedCIDR(cidr string, isSrc bool) ([]expr.Any, error) {
var proto byte if !narrows(cidr) && !strings.HasPrefix(cidr, "!") {
for i, a := range strings.Split(strings.TrimPrefix(addr, "!"), ",") { return matchNFProto(addrFamily(cidr)), nil
a, _, _ = strings.Cut(a, "/")
p := byte(unix.NFPROTO_IPV6)
if ip := net.ParseIP(a); ip != nil && ip.To4() != nil {
p = unix.NFPROTO_IPV4
}
if i > 0 && p != proto {
return nil, fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr)
}
proto = p
} }
dst, err := matchDestCIDR(addr) e, err := matchAddrCIDR(cidr, isSrc)
if err != nil { if err != nil {
return nil, err return nil, err
} }
return append(matchNFProto(proto), dst...), nil return append(matchNFProto(addrFamily(cidr)), e...), nil
}
// addrFamily is the NFPROTO of a single-family splitAddrs alternative, 0 for none; unparsable is IPv6.
func addrFamily(list string) byte {
if list == "" {
return 0
}
a, _, _ := strings.Cut(strings.TrimPrefix(list, "!"), ",")
if p, err := parsePrefix(a); err == nil && p.Addr().Is4() {
return unix.NFPROTO_IPV4
}
return unix.NFPROTO_IPV6
} }
func matchNFProto(proto byte) []expr.Any { func matchNFProto(proto byte) []expr.Any {
+682 -73
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"}}},
ifaces = c.resolveZoneInterfaces("all", "") {"all", []zoneMatch{{}}},
if len(ifaces) != 1 || ifaces[0] != "" { {"fw", []zoneMatch{{}}},
t.Errorf("resolveZoneInterfaces(all) = %v, want [\"\"]", ifaces) } {
} 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("fw", "") }
if len(ifaces) != 1 || ifaces[0] != "" {
t.Errorf("resolveZoneInterfaces(fw) = %v, want [\"\"]", ifaces)
} }
} }
@@ -211,7 +210,7 @@ func TestMatchSourceCIDR_IPv6(t *testing.T) {
} }
for _, tt := range tests { for _, tt := range tests {
exprs, err := matchSourceCIDR(tt.input) exprs, err := matchAddrCIDR(tt.input, true)
if tt.wantErr { if tt.wantErr {
if err == nil { if err == nil {
t.Errorf("matchSourceCIDR(%q) should fail", tt.input) t.Errorf("matchSourceCIDR(%q) should fail", tt.input)
@@ -242,7 +241,7 @@ func TestMatchDestCIDR_IPv6(t *testing.T) {
} }
for _, tt := range tests { for _, tt := range tests {
exprs, err := matchDestCIDR(tt.input) exprs, err := matchAddrCIDR(tt.input, false)
if tt.wantErr { if tt.wantErr {
if err == nil { if err == nil {
t.Errorf("matchDestCIDR(%q) should fail", tt.input) t.Errorf("matchDestCIDR(%q) should fail", tt.input)
@@ -962,7 +961,7 @@ func TestCompile_RateLimit(t *testing.T) {
} }
func TestNegatedAddress(t *testing.T) { func TestNegatedAddress(t *testing.T) {
exprs, err := matchSourceCIDR("!192.168.1.0/24") exprs, err := matchAddrCIDR("!192.168.1.0/24", true)
if err != nil { if err != nil {
t.Fatalf("matchSourceCIDR(!192.168.1.0/24) error: %v", err) t.Fatalf("matchSourceCIDR(!192.168.1.0/24) error: %v", err)
} }
@@ -974,7 +973,7 @@ func TestNegatedAddress(t *testing.T) {
t.Errorf("negated address should use CmpOpNeq, got %v", cmp.Op) t.Errorf("negated address should use CmpOpNeq, got %v", cmp.Op)
} }
exprs, err = matchDestCIDR("!10.0.0.1") exprs, err = matchAddrCIDR("!10.0.0.1", false)
if err != nil { if err != nil {
t.Fatalf("matchDestCIDR(!10.0.0.1) error: %v", err) t.Fatalf("matchDestCIDR(!10.0.0.1) error: %v", err)
} }
@@ -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) {
@@ -1954,57 +1976,57 @@ func TestCompile_CommaZoneLists(t *testing.T) {
{ {
name: "address list after colon belongs to one zone", name: "address list after colon belongs to one zone",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "net:192.0.2.1,198.51.100.1"}, rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "net:192.0.2.1,198.51.100.1"},
want: map[string][]string{"forward": {"iif=eth1 oif=eth0 daddr=192.0.2.1", "iif=eth1 oif=eth0 daddr=198.51.100.1"}}, want: map[string][]string{"forward": {"iif=eth1 oif=eth0 ip4 daddr=192.0.2.1", "iif=eth1 oif=eth0 ip4 daddr=198.51.100.1"}},
}, },
{ {
name: "zone:address inside a list", name: "zone:address inside a list",
rule: config.Rule{Action: config.RuleAccept, Source: "lan,svr:203.0.113.7", Dest: "fw"}, rule: config.Rule{Action: config.RuleAccept, Source: "lan,svr:203.0.113.7", Dest: "fw"},
want: map[string][]string{"input": {"iif=eth1", "iif=eth2 saddr=203.0.113.7"}}, want: map[string][]string{"input": {"iif=eth1", "iif=eth2 ip4 saddr=203.0.113.7"}},
}, },
{ {
name: "dnat source list", name: "dnat source list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"}, want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"},
"forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10"}}, "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth1 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat source list skips the target zone", name: "dnat source list skips the target zone",
rule: config.Rule{Action: config.RuleDNAT, Source: "net,svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "net,svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10"}}, want: map[string][]string{"prerouting": {"iif=eth0"}, "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat lone source zone may equal the target zone", name: "dnat lone source zone may equal the target zone",
rule: config.Rule{Action: config.RuleDNAT, Source: "svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth2"}, "forward": {"iif=eth2 oif=eth2 daddr=192.0.2.10"}}, want: map[string][]string{"prerouting": {"iif=eth2"}, "forward": {"iif=eth2 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat exclusion source skips the target zone", name: "dnat exclusion source skips the target zone",
rule: config.Rule{Action: config.RuleDNAT, Source: "all!fw,anycast", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "all!fw,anycast", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth1", "iif=eth0"}, want: map[string][]string{"prerouting": {"iif=eth1", "iif=eth0"},
"forward": {"iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10"}}, "forward": {"iif=eth1 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat intrazone exclusion source keeps the target zone", name: "dnat intrazone exclusion source keeps the target zone",
rule: config.Rule{Action: config.RuleDNAT, Source: "all+!fw,anycast,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "all+!fw,anycast,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth2"}, want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth2"},
"forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth2 oif=eth2 daddr=192.0.2.10"}}, "forward": {"iif=eth0 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth2 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat all source expands per zone and skips fw and the target zone", name: "dnat all source expands per zone and skips fw and the target zone",
rule: config.Rule{Action: config.RuleDNAT, Source: "all", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "all", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0"}, want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0"},
"forward": {"iif=eth3 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10"}}, "forward": {"iif=eth3 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth1 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat any+ source keeps the target zone", name: "dnat any+ source keeps the target zone",
rule: config.Rule{Action: config.RuleDNAT, Source: "any+", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "any+", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0", "iif=eth2"}, want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0", "iif=eth2"},
"forward": {"iif=eth3 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth2 oif=eth2 daddr=192.0.2.10"}}, "forward": {"iif=eth3 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth1 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 daddr=192.0.2.10", "iif=eth2 oif=eth2 ip4 daddr=192.0.2.10"}},
}, },
{ {
name: "dnat to fw accepts in input", name: "dnat to fw accepts in input",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.1", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.1", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0"}, "input": {"iif=eth0 daddr=192.0.2.1"}}, want: map[string][]string{"prerouting": {"iif=eth0"}, "input": {"iif=eth0 ip4 daddr=192.0.2.1"}},
}, },
{ {
name: "redirect accepts in input without daddr", name: "redirect accepts in input without daddr",
@@ -2014,13 +2036,13 @@ func TestCompile_CommaZoneLists(t *testing.T) {
{ {
name: "dnat source address list", name: "dnat source address list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net:192.0.2.5,198.51.100.5", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, rule: config.Rule{Action: config.RuleDNAT, Source: "net:192.0.2.5,198.51.100.5", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0 saddr=192.0.2.5", "iif=eth0 saddr=198.51.100.5"}, want: map[string][]string{"prerouting": {"iif=eth0 ip4 saddr=192.0.2.5", "iif=eth0 ip4 saddr=198.51.100.5"},
"forward": {"iif=eth0 oif=eth2 saddr=192.0.2.5 daddr=192.0.2.10", "iif=eth0 oif=eth2 saddr=198.51.100.5 daddr=192.0.2.10"}}, "forward": {"iif=eth0 oif=eth2 ip4 saddr=192.0.2.5 daddr=192.0.2.10", "iif=eth0 oif=eth2 ip4 saddr=198.51.100.5 daddr=192.0.2.10"}},
}, },
{ {
name: "negated address list stays one AND-ed rule", name: "negated address list is one AND-ed rule per family",
rule: config.Rule{Action: config.RuleAccept, Source: "net:!192.0.2.5,198.51.100.5", Dest: "fw"}, rule: config.Rule{Action: config.RuleAccept, Source: "net:!192.0.2.5,198.51.100.5", Dest: "fw"},
want: map[string][]string{"input": {"iif=eth0 !saddr=192.0.2.5 !saddr=198.51.100.5"}}, want: map[string][]string{"input": {"iif=eth0 ip4 !saddr=192.0.2.5 !saddr=198.51.100.5", "iif=eth0 ip6"}},
}, },
{ {
name: "zone named like all/any keyword is a plain zone", name: "zone named like all/any keyword is a plain zone",
@@ -2035,7 +2057,7 @@ func TestCompile_CommaZoneLists(t *testing.T) {
{ {
name: "interface-less zone kept when address narrows it", name: "interface-less zone kept when address narrows it",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn:192.0.2.1"}, rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn:192.0.2.1"},
want: map[string][]string{"forward": {"iif=eth1 daddr=192.0.2.1"}}, want: map[string][]string{"forward": {"iif=eth1 ip4 daddr=192.0.2.1"}},
}, },
{ {
name: "fw source matches dest zone oif", name: "fw source matches dest zone oif",
@@ -2045,30 +2067,30 @@ func TestCompile_CommaZoneLists(t *testing.T) {
{ {
name: "fw to all has no oif", name: "fw to all has no oif",
rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "all:192.0.2.1"}, rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "all:192.0.2.1"},
want: map[string][]string{"output": {"daddr=192.0.2.1"}}, want: map[string][]string{"output": {"ip4 daddr=192.0.2.1"}},
}, },
{ {
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 ip4 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 ip4 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 ip4 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",
@@ -2473,7 +2495,7 @@ func TestCompile_RejectPerProto(t *testing.T) {
} }
func TestNegatedAddressList(t *testing.T) { func TestNegatedAddressList(t *testing.T) {
exprs, err := matchDestCIDR("!192.0.2.1,198.51.100.1") exprs, err := matchAddrCIDR("!192.0.2.1,198.51.100.1", false)
if err != nil { if err != nil {
t.Fatalf("matchDestCIDR error: %v", err) t.Fatalf("matchDestCIDR error: %v", err)
} }
@@ -2591,25 +2613,6 @@ func TestCompile_ColonRanges(t *testing.T) {
} }
} }
func TestMatchOrigDest_FamilyGuard(t *testing.T) {
for addr, want := range map[string]byte{
"203.0.113.5": unix.NFPROTO_IPV4,
"!203.0.113.0/24,192.0.2.1": unix.NFPROTO_IPV4,
"2001:db8::5": unix.NFPROTO_IPV6,
} {
e, err := matchOrigDest(addr)
if err != nil {
t.Fatalf("%s: %v", addr, err)
}
if m, ok := e[0].(*expr.Meta); !ok || m.Key != expr.MetaKeyNFPROTO || e[1].(*expr.Cmp).Data[0] != want {
t.Errorf("%s: missing nfproto %d guard: %v", addr, want, e[:2])
}
}
if _, err := matchOrigDest("!203.0.113.5,2001:db8::5"); err == nil {
t.Error("mixed IPv4/IPv6 origdest: want error")
}
}
func TestCompile_OrigDestForwardRejected(t *testing.T) { func TestCompile_OrigDestForwardRejected(t *testing.T) {
cfg := &config.Config{ cfg := &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
@@ -2699,7 +2702,7 @@ func TestCompile_ConntrackZones(t *testing.T) {
{ {
name: "source and dest addresses", name: "source and dest addresses",
ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:192.0.2.1,198.51.100.1", Dest: "fw:203.0.113.1"}, ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:192.0.2.1,198.51.100.1", Dest: "fw:203.0.113.1"},
want: map[string][]string{"raw_prerouting": {"iif=eth0 saddr=192.0.2.1 daddr=203.0.113.1", "iif=eth0 saddr=198.51.100.1 daddr=203.0.113.1"}}, want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 saddr=192.0.2.1 daddr=203.0.113.1", "iif=eth0 ip4 saddr=198.51.100.1 daddr=203.0.113.1"}},
}, },
{ {
name: "fw source goes to raw_output with dest oif", name: "fw source goes to raw_output with dest oif",
@@ -2709,7 +2712,7 @@ func TestCompile_ConntrackZones(t *testing.T) {
{ {
name: "all matches no interface", name: "all matches no interface",
ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Dest: "fw:192.0.2.53"}, ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Dest: "fw:192.0.2.53"},
want: map[string][]string{"raw_prerouting": {"daddr=192.0.2.53"}}, want: map[string][]string{"raw_prerouting": {"ip4 daddr=192.0.2.53"}},
}, },
{ {
name: "interface-less zone fails closed", name: "interface-less zone fails closed",
@@ -2767,9 +2770,9 @@ func TestCompile_ConntrackZones(t *testing.T) {
want: map[string][]string{"raw_output": {"oif=eth1"}}, want: map[string][]string{"raw_output": {"oif=eth1"}},
}, },
{ {
name: "negated addresses stay one AND-ed match", name: "negated addresses are one AND-ed match per family",
ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:!192.0.2.1,198.51.100.1"}, ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:!192.0.2.1,198.51.100.1"},
want: map[string][]string{"raw_prerouting": {"iif=eth0 !saddr=192.0.2.1 !saddr=198.51.100.1"}}, want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 !saddr=192.0.2.1 !saddr=198.51.100.1", "iif=eth0 ip6"}},
}, },
{ {
name: "sport", name: "sport",
@@ -2789,7 +2792,7 @@ func TestCompile_ConntrackZones(t *testing.T) {
{ {
name: "fw dest zone with address matches daddr", name: "fw dest zone with address matches daddr",
ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw:192.0.2.1"}, ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw:192.0.2.1"},
want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=192.0.2.1"}}, want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 daddr=192.0.2.1"}},
}, },
{ {
name: "unknown zone fails closed", name: "unknown zone fails closed",
@@ -2844,7 +2847,7 @@ func TestCompile_ConntrackZones(t *testing.T) {
{ {
name: "dest zone with address matches daddr", name: "dest zone with address matches daddr",
ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan:203.0.113.10"}, ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan:203.0.113.10"},
want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, want: map[string][]string{"raw_prerouting": {"iif=eth0 ip4 daddr=203.0.113.10"}},
}, },
} }
for _, tt := range tests { for _, tt := range tests {
@@ -2943,7 +2946,7 @@ func TestCompile_ConntrackHelperZones(t *testing.T) {
{ {
name: "dest zone with address in prerouting", name: "dest zone with address in prerouting",
ct: config.ConntrackRule{Source: "net", Dest: "lan:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"21"}}, ct: config.ConntrackRule{Source: "net", Dest: "lan:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"21"}},
want: map[string][]string{"helper_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, want: map[string][]string{"helper_prerouting": {"iif=eth0 ip4 daddr=203.0.113.10"}},
}, },
{ {
name: "dest zone without address is rejected in prerouting", name: "dest zone without address is rejected in prerouting",
@@ -3009,7 +3012,7 @@ func TestCompile_AllIncludesFirewallMatches(t *testing.T) {
{ {
name: "all address kept on added fw rules", name: "all address kept on added fw rules",
rule: config.Rule{Action: config.RuleAccept, Source: "all:192.0.2.5", Dest: "all"}, rule: config.Rule{Action: config.RuleAccept, Source: "all:192.0.2.5", Dest: "all"},
want: map[string][]string{"input": {"saddr=192.0.2.5"}, "output": {"saddr=192.0.2.5"}, "forward": {"saddr=192.0.2.5"}}, want: map[string][]string{"input": {"ip4 saddr=192.0.2.5"}, "output": {"ip4 saddr=192.0.2.5"}, "forward": {"ip4 saddr=192.0.2.5"}},
}, },
{ {
name: "dnat with all source skips fw", name: "dnat with all source skips fw",
@@ -3155,3 +3158,609 @@ func TestCompile_Dispositions(t *testing.T) {
} }
} }
} }
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 daddr=198.51.100.0/24",
"iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 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 !daddr=198.51.100.0/24",
"iif=wlo1 ip6 oif=enp2s0",
"iif=enp2s0 ip4 !saddr=198.51.100.0/24 oif=wlo1 !daddr=192.0.2.0/24",
"iif=enp2s0 ip6 oif=wlo1",
}
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 daddr=198.51.100.0/24",
"iif=enp2s0 ip4 saddr=198.51.100.0/24 oif=wlo1 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 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_LimitNegatedAddrPerTableFamily(t *testing.T) {
for _, tt := range []struct {
name string
family config.AddressFamily
source string
wantErr bool
}{
{"ip v4 negation", config.FamilyIP, "net:!192.0.2.1", false},
{"ip v6 negation", config.FamilyIP, "net:!2001:db8::1", false},
{"ip mixed negation", config.FamilyIP, "net:!192.0.2.1,2001:db8::1", false},
{"ip6 v6 negation", config.FamilyIP6, "net:!2001:db8::1", false},
{"ip6 v4 negation", config.FamilyIP6, "net:!192.0.2.1", false},
{"inet v4 negation", config.FamilyINET, "net:!192.0.2.1", true},
} {
t.Run(tt.name, func(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: tt.family},
Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}},
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}},
Rules: []config.Rule{{Action: config.RuleAccept, Source: tt.source, Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, RateLimit: "10/sec"}},
PortGroups: make(map[string]config.PortGroup),
}
state, err := NewCompiler(cfg).Compile()
if tt.wantErr {
if err == nil || !strings.Contains(err.Error(), "negated address in an inet table") {
t.Fatalf("Compile() error = %v, want negated-address inet error", err)
}
return
}
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
n := 0
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
n++
}
}
if n != 1 {
t.Errorf("got %d rule:0 rules, want 1", n)
}
})
}
}
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)
}
}
}
+239
View File
@@ -0,0 +1,239 @@
package nftables
import (
"os"
"os/exec"
"reflect"
"strings"
"testing"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
func guardCfg(af config.AddressFamily) *config.Config {
return hostsCfg(func(c *config.Config) {
c.Settings.AddressFamily = af
c.Hosts[0].Addresses = append(c.Hosts[0].Addresses, "2001:db8::/64")
c.Hosts[0].Exclusions = []string{"192.0.2.9", "2001:db8::9"}
c.Interfaces[0].Options.NoSmurfs = true
c.Rules = append(c.Rules,
config.Rule{Action: config.RuleDrop, Source: "net:203.0.113.7", Dest: "lan:192.0.2.5"},
config.Rule{Action: config.RuleDrop, Source: "net:2001:db8:1::7", Dest: "lan:2001:db8::5"},
config.Rule{Action: config.RuleDrop, Source: "vpn", Dest: "net:!192.0.2.1"},
config.Rule{Action: config.RuleDrop, Source: "vpn", Dest: "net:!192.0.2.1,2001:db8::1"},
config.Rule{Action: config.RuleDrop, Source: "vpn:203.0.113.7", Dest: "net:!2001:db8::5"},
config.Rule{Action: config.RuleDNAT, Source: "vpn", Dest: "lan:192.0.2.10", Proto: "tcp",
DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,2001:db8::5"},
config.Rule{Action: config.RuleDrop, Source: "vpn:192.0.2.77,2001:db8:7::7", Dest: "fw"})
c.Blrules = []config.BlruleRule{{Action: config.BlruleDrop, Source: "vpn:!192.0.2.1,2001:db8::1", Dest: "fw"}}
c.SNAT = []config.SNATRule{
{Action: config.SNATMasquerade, Dest: "wlo1", Source: "!192.0.2.0/24,2001:db8::/48"},
{Action: config.SNATAddress, Address: "203.0.113.1", Dest: "wlo1", Source: "!192.0.2.9,2001:db8::9"},
{Action: config.SNATAddress, Address: "203.0.113.1", Dest: "wlo1", Source: "2001:db8::/48"},
}
c.Tunnels = []config.Tunnel{{Type: "gre", Zone: "vpn", Gateways: []string{"203.0.113.50", "2001:db8:5::1"}}}
c.StaticNAT = []config.StaticNAT{
{External: "203.0.113.60", Interface: "wlo1", Internal: "192.0.2.60"},
{External: "2001:db8:6::1", Interface: "wlo1", Internal: "2001:db8::60"},
}
})
}
func TestCompile_FamilyGuardsDecodable(t *testing.T) {
addrLen := map[config.AddressFamily]uint32{config.FamilyIP: 4, config.FamilyIP6: 16}
for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} {
t.Run(string(af), func(t *testing.T) {
for chain, rules := range mustCompile(t, guardCfg(af)).Rules {
for _, r := range rules {
guards, l3 := 0, false
for i, e := range r.Exprs {
if nfprotoGuard(r.Exprs, i) != 0 {
guards++
if l3 {
t.Errorf("%s %s: guard after a network payload: %s", chain, r.Tag, describeRule(r))
}
}
if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseNetworkHeader {
l3 = true
if n := addrLen[af]; n != 0 && (p.Len == 4 || p.Len == 16) && p.Len != n {
t.Errorf("%s %s: other-family address in %s table: %s", chain, r.Tag, af, describeRule(r))
}
}
}
if max := map[bool]int{true: 1, false: 0}[af == config.FamilyINET]; guards > max {
t.Errorf("%s %s: %d family guards in %s table: %s", chain, r.Tag, guards, af, describeRule(r))
}
}
}
})
}
}
func TestCompile_FamilyGuardsPerFamily(t *testing.T) {
state := mustCompile(t, guardCfg(config.FamilyINET))
for _, tt := range []struct {
chain, tag string
want []string
}{
{"forward", "rule:3", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 !daddr=192.0.2.1",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip6 !daddr=2001:db8::/64",
"iif=tun0 oif=wlo1 ip6 daddr=2001:db8::9",
"iif=tun0 oif=enp2s0 ip6",
}},
{"forward", "rule:4", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 !daddr=192.0.2.1",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 !daddr=192.0.2.1",
"iif=tun0 oif=wlo1 ip6 !daddr=2001:db8::/64 !daddr=2001:db8::1",
"iif=tun0 oif=wlo1 ip6 daddr=2001:db8::9 !daddr=2001:db8::1",
"iif=tun0 oif=enp2s0 ip6 !daddr=2001:db8::1",
}},
{"forward", "rule:5", []string{
"iif=tun0 oif=wlo1 ip4 !daddr=192.0.2.0/24 saddr=203.0.113.7",
"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.9 saddr=203.0.113.7",
"iif=tun0 oif=enp2s0 ip4 !daddr=198.51.100.0/24 saddr=203.0.113.7",
}},
{"prerouting", "rule:6", []string{"iif=tun0 ip4 !daddr=203.0.113.5"}},
{"forward", "rule:6:accept", []string{"iif=tun0 oif=wlo1 ip4 daddr=192.0.2.0/24 !daddr=192.0.2.9 daddr=192.0.2.10"}},
{"input", "rule:7", []string{"iif=tun0 ip4 saddr=192.0.2.77", "iif=tun0 ip6 saddr=2001:db8:7::7"}},
{"input", "blrule:0", []string{"iif=tun0 ip4 !saddr=192.0.2.1", "iif=tun0 ip6 !saddr=2001:db8::1"}},
{"postrouting", "snat:0", []string{"oif=wlo1 ip4 !saddr=192.0.2.0/24", "oif=wlo1 ip6 !saddr=2001:db8::/48"}},
{"postrouting", "snat:1", []string{"oif=wlo1 ip4 !saddr=192.0.2.9"}},
{"postrouting", "snat:2", nil},
{"input", "tunnel:0", []string{"ip4 saddr=203.0.113.50", "ip6 saddr=2001:db8:5::1"}},
{"prerouting", "staticnat:dnat:0", []string{"iif=wlo1 ip4 daddr=203.0.113.60"}},
{"postrouting", "staticnat:snat:1", []string{"oif=wlo1 ip6 saddr=2001:db8::60"}},
} {
if got := describeTagged(state, tt.chain, tt.tag); !reflect.DeepEqual(got, tt.want) {
t.Errorf("%s %s = %q, want %q", tt.chain, tt.tag, got, tt.want)
}
}
}
// TestCompile_NegatedV4DropKeepsV6 checks a single-family table keeps its family's half of a
// negated DROP: "everything except 192.0.2.1" still drops all IPv6.
func TestCompile_NegatedV4DropKeepsV6(t *testing.T) {
for af, want := range map[config.AddressFamily][]string{
config.FamilyIP: {"iif=tun0 oif=wlo1 !daddr=192.0.2.0/24 !daddr=192.0.2.1", "iif=tun0 oif=wlo1 daddr=192.0.2.9 !daddr=192.0.2.1", "iif=tun0 oif=enp2s0 !daddr=198.51.100.0/24 !daddr=192.0.2.1"},
config.FamilyIP6: {"iif=tun0 oif=wlo1 !daddr=2001:db8::/64", "iif=tun0 oif=wlo1 daddr=2001:db8::9", "iif=tun0 oif=enp2s0"},
} {
if got := describeTagged(mustCompile(t, guardCfg(af)), "forward", "rule:3"); !reflect.DeepEqual(got, want) {
t.Errorf("%s rule:3 = %q, want %q", af, got, want)
}
}
}
func TestSplitAddrs(t *testing.T) {
for in, want := range map[string][]string{
"": {""},
"192.0.2.1,2001:db8::1": {"192.0.2.1", "2001:db8::1"},
"!192.0.2.1": {"!192.0.2.1", "::/0"},
"!2001:db8::1": {"0.0.0.0/0", "!2001:db8::1"},
"!192.0.2.1,2001:db8::1,198.51.100.0/24": {"!192.0.2.1,198.51.100.0/24", "!2001:db8::1"},
} {
if got := splitAddrs(in); !reflect.DeepEqual(got, want) {
t.Errorf("splitAddrs(%q) = %q, want %q", in, got, want)
}
}
}
func TestMatchGuardedCIDR(t *testing.T) {
for in, want := range map[string]string{
"192.0.2.1": "ip4 saddr=192.0.2.1",
"2001:db8::/48": "ip6 saddr=2001:db8::/48",
"!192.0.2.1,198.51.100.0/24": "ip4 !saddr=192.0.2.1 !saddr=198.51.100.0/24",
"!2001:db8::1": "ip6 !saddr=2001:db8::1",
"0.0.0.0/0": "ip4",
"::/0": "ip6",
} {
e, err := matchSourceCIDR(in)
if err != nil {
t.Fatalf("%s: %v", in, err)
}
if got := describeRule(ManagedRule{Exprs: e}); got != want {
t.Errorf("matchSourceCIDR(%q) = %q, want %q", in, got, want)
}
}
if e, _ := matchDestCIDR("!2001:db8::5"); describeRule(ManagedRule{Exprs: e}) != "ip6 !daddr=2001:db8::5" {
t.Errorf("matchDestCIDR(!2001:db8::5) = %q", describeRule(ManagedRule{Exprs: e}))
}
if _, err := matchSourceCIDR("!nonsense"); err == nil {
t.Error("invalid negated address: want error")
}
}
func TestFamilyGuards(t *testing.T) {
v4, _ := matchSourceCIDR("192.0.2.1")
v6, _ := matchDestCIDR("2001:db8::1")
both, _ := matchDestCIDR("198.51.100.1")
state := func(e ...[]expr.Any) *FirewallState {
var r []expr.Any
for _, x := range e {
r = append(r, x...)
}
return &FirewallState{Rules: map[string][]ManagedRule{"input": {{Exprs: r, Tag: "t"}}}}
}
s := state(v4, both)
if err := familyGuards(s, config.FamilyINET); err != nil || describeRule(s.Rules["input"][0]) != "ip4 saddr=192.0.2.1 daddr=198.51.100.1" {
t.Errorf("same-family guards not merged: %v %q", err, describeRule(s.Rules["input"][0]))
}
if err := familyGuards(state(v4, v6), config.FamilyINET); err == nil || !strings.Contains(err.Error(), "conflicting") {
t.Errorf("conflicting guards: want error, got %v", err)
}
s = state(v6)
if err := familyGuards(s, config.FamilyIP); err != nil || len(s.Rules["input"]) != 0 {
t.Errorf("ip table must drop IPv6 rules: %v %v", err, s.Rules["input"])
}
s = state(v6)
if err := familyGuards(s, config.FamilyIP6); err != nil || describeRule(s.Rules["input"][0]) != "daddr=2001:db8::1" {
t.Errorf("ip6 table must strip the guard: %v %q", err, describeRule(s.Rules["input"][0]))
}
if !famsAgree(0, unix.NFPROTO_IPV4, 0, unix.NFPROTO_IPV4) || famsAgree(unix.NFPROTO_IPV4, 0, unix.NFPROTO_IPV6) {
t.Error("famsAgree")
}
}
// TestNetnsNftListDecodes applies each family in a fresh user+net namespace and requires
// nft(8) to list the ruleset and a second plan to be empty. Needs unshare and nft.
func TestNetnsNftListDecodes(t *testing.T) {
if af := os.Getenv("TOMSWALL_NETNS_CHILD"); af != "" {
netnsChild(t, config.AddressFamily(af))
return
}
if os.Getenv("TOMSWALL_NETNS_TEST") == "" {
t.Skip("set TOMSWALL_NETNS_TEST=1 to run (needs unshare and nft)")
}
for _, af := range []config.AddressFamily{config.FamilyINET, config.FamilyIP, config.FamilyIP6} {
cmd := exec.Command("unshare", "-rn", os.Args[0], "-test.run=^TestNetnsNftListDecodes$", "-test.v")
cmd.Env = append(os.Environ(), "TOMSWALL_NETNS_CHILD="+string(af))
if out, err := cmd.CombinedOutput(); err != nil {
t.Errorf("%s: %v\n%s", af, err, out)
}
}
}
func netnsChild(t *testing.T, af config.AddressFamily) {
e, err := NewEngine(guardCfg(af))
if err != nil {
t.Fatal(err)
}
cs, err := e.Plan()
if err != nil {
t.Fatal(err)
}
if err := e.Apply(cs); err != nil {
t.Fatal(err)
}
if out, err := exec.Command("nft", "list", "ruleset").CombinedOutput(); err != nil {
t.Fatalf("nft list ruleset: %v\n%s", err, out)
}
if cs, err = e.Plan(); err != nil || len(cs.Add)+len(cs.Remove) != 0 {
t.Fatalf("second plan not empty: %d add, %d remove, err %v", len(cs.Add), len(cs.Remove), err)
}
}
+28 -8
View File
@@ -2,6 +2,7 @@ package shorewall
import ( import (
"fmt" "fmt"
"log/slog"
"strconv" "strconv"
"strings" "strings"
@@ -141,6 +142,17 @@ 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"
} }
@@ -236,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,
@@ -380,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)
@@ -403,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 {
+48
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"
@@ -756,3 +757,50 @@ func TestConvert_Dispositions(t *testing.T) {
} }
} }
} }
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)
}
}
+41 -24
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.
@@ -193,7 +210,7 @@ func Abort() error {
} }
func restorePending(p *pending) error { func restorePending(p *pending) error {
if err := restore(p.Snapshot); err != nil { if err := Restore(p.Snapshot); err != nil {
return fmt.Errorf("restoring snapshot: %w", err) return fmt.Errorf("restoring snapshot: %w", err)
} }
return Discard() return Discard()
+7 -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
} }
+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
+2
View File
@@ -6,6 +6,8 @@ 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) # ct state invalid/untracked verdict: accept, drop, reject, continue (pass to rules)