Compare commits
30 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 78dbb6ad18 | |||
| 9aee3ad7eb | |||
| 9bfdf292ad | |||
| bc65312647 | |||
| 9a65973d7f | |||
| 460eb20db5 | |||
| 42a4dab6a3 | |||
| 190ff72643 | |||
| b8ad59b053 | |||
| 3174eabd94 | |||
| 0b68110220 | |||
| 70df237121 | |||
| 799c7f3524 | |||
| ecc349cb6f | |||
| 695869c80b | |||
| 96a1ba8351 | |||
| afa056b454 | |||
| f593f7625d | |||
| 7c8bd87ec0 | |||
| 9854b0e7b6 | |||
| c7e02c089a | |||
| 9092b463a0 | |||
| 502d06bdda | |||
| 4ad55fc65e | |||
| 7fcbb5fad8 | |||
| dc406c4f56 | |||
| 2832dd0e7d | |||
| 15ab32431d | |||
| be391ed385 | |||
| 2e8d51759d |
@@ -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]
|
||||||
|
|||||||
@@ -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`.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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,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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
[Unit]
|
||||||
|
Description=tomswall firewall (apply local config at boot)
|
||||||
|
Documentation=https://git.unkin.net/unkin/tomswall
|
||||||
|
DefaultDependencies=no
|
||||||
|
Wants=network-pre.target
|
||||||
|
Before=network-pre.target shutdown.target
|
||||||
|
After=local-fs.target systemd-sysctl.service
|
||||||
|
Conflicts=shutdown.target tomswall-agent.service
|
||||||
|
StartLimitIntervalSec=60
|
||||||
|
StartLimitBurst=5
|
||||||
|
|
||||||
|
[Service]
|
||||||
|
Type=oneshot
|
||||||
|
RemainAfterExit=yes
|
||||||
|
Environment=TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
|
||||||
|
EnvironmentFile=-/etc/tomswall/tomswall.env
|
||||||
|
ExecStart=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
|
||||||
|
ExecReload=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
|
||||||
|
# Fails open: after StartLimitBurst failures within StartLimitIntervalSec, boot continues without the ruleset.
|
||||||
|
Restart=on-failure
|
||||||
|
RestartSec=5
|
||||||
|
# No ExecStop: stopping the unit leaves the ruleset in place (flush would open the firewall).
|
||||||
|
|
||||||
|
[Install]
|
||||||
|
WantedBy=sysinit.target
|
||||||
@@ -6,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)
|
||||||
|
|||||||
Reference in New Issue
Block a user