Add try/confirm safe-apply and agent auto-revert
This commit is contained in:
@@ -33,6 +33,8 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`,
|
||||
|
||||
root.AddCommand(
|
||||
applyCmd(),
|
||||
tryCmd(),
|
||||
confirmCmd(),
|
||||
planCmd(),
|
||||
validateCmd(),
|
||||
statusCmd(),
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.unkin.net/unkin/tomswall/internal/agent"
|
||||
)
|
||||
|
||||
var tryPIDFile = "/run/tomswall/try.pid"
|
||||
|
||||
func tryCmd() *cobra.Command {
|
||||
var timeout time.Duration
|
||||
cmd := &cobra.Command{
|
||||
Use: "try",
|
||||
Short: "Apply configuration and revert unless confirmed within a timeout",
|
||||
Long: `Try snapshots the live tomswall table, applies the configuration, then waits
|
||||
for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is restored
|
||||
atomically. Confirm from a new session to prove new connections still work.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
cfg, err := loadConfig()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Registered before apply so a confirm or hangup cannot be missed.
|
||||
signal.Ignore(syscall.SIGPIPE)
|
||||
confirm := make(chan os.Signal, 1)
|
||||
signal.Notify(confirm, syscall.SIGUSR1)
|
||||
abort := make(chan os.Signal, 1)
|
||||
signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM)
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(tryPIDFile), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(tryPIDFile, []byte(strconv.Itoa(os.Getpid())), 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
defer os.Remove(tryPIDFile)
|
||||
|
||||
revert, err := agent.EngineApplier{}.Apply(context.Background(), cfg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("applying changes: %w", err)
|
||||
}
|
||||
fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout)
|
||||
if err := confirmOrRevert(confirm, abort, timeout, revert); err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Println("Confirmed.")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().DurationVar(&timeout, "timeout", 60*time.Second, "time to wait for confirmation before reverting")
|
||||
return cmd
|
||||
}
|
||||
|
||||
// confirmOrRevert waits for confirm; on abort or timeout it runs revert.
|
||||
func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() error) error {
|
||||
reason := "not confirmed within " + timeout.String()
|
||||
select {
|
||||
case <-confirm:
|
||||
return nil
|
||||
case s := <-abort:
|
||||
reason = "interrupted by " + s.String()
|
||||
case <-time.After(timeout):
|
||||
}
|
||||
if err := revert(); err != nil {
|
||||
return fmt.Errorf("%s; revert failed: %w", reason, err)
|
||||
}
|
||||
return fmt.Errorf("%s: previous ruleset restored", reason)
|
||||
}
|
||||
|
||||
func confirmCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "confirm",
|
||||
Short: "Keep the configuration applied by a pending 'tomswall try'",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
b, err := os.ReadFile(tryPIDFile)
|
||||
if os.IsNotExist(err) {
|
||||
return fmt.Errorf("no pending 'tomswall try'")
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pid, err := strconv.Atoi(strings.TrimSpace(string(b)))
|
||||
if err != nil {
|
||||
return fmt.Errorf("parsing %s: %w", tryPIDFile, err)
|
||||
}
|
||||
// A stale pidfile must not signal an unrelated process.
|
||||
comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid))
|
||||
if err != nil || strings.TrimSpace(string(comm)) != "tomswall" {
|
||||
return fmt.Errorf("no pending 'tomswall try' (stale %s)", tryPIDFile)
|
||||
}
|
||||
if err := syscall.Kill(pid, syscall.SIGUSR1); err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Println("Confirmed.")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestConfirmOrRevert(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
confirm bool
|
||||
abort bool
|
||||
revertErr error
|
||||
wantRevert bool
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "confirmed keeps ruleset", confirm: true},
|
||||
{name: "timeout reverts", wantRevert: true, wantErr: true},
|
||||
{name: "hangup reverts", abort: true, wantRevert: true, wantErr: true},
|
||||
{name: "revert failure surfaces", revertErr: errors.New("boom"), wantRevert: true, wantErr: true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
confirm := make(chan os.Signal, 1)
|
||||
abort := make(chan os.Signal, 1)
|
||||
if tt.confirm {
|
||||
confirm <- syscall.SIGUSR1
|
||||
}
|
||||
if tt.abort {
|
||||
abort <- syscall.SIGHUP
|
||||
}
|
||||
reverted := false
|
||||
err := confirmOrRevert(confirm, abort, 20*time.Millisecond, func() error {
|
||||
reverted = true
|
||||
return tt.revertErr
|
||||
})
|
||||
if reverted != tt.wantRevert {
|
||||
t.Errorf("reverted = %v, want %v", reverted, tt.wantRevert)
|
||||
}
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("err = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if tt.revertErr != nil && !errors.Is(err, tt.revertErr) {
|
||||
t.Errorf("revert error not wrapped: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+48
-16
@@ -2,18 +2,21 @@ package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/url"
|
||||
"time"
|
||||
|
||||
"git.unkin.net/unkin/tomswall/internal/config"
|
||||
"git.unkin.net/unkin/tomswall/internal/nftables"
|
||||
)
|
||||
|
||||
// Applier applies a translated config to the firewall. Abstracted so the run
|
||||
// loop is testable without touching the kernel.
|
||||
// Applier applies a translated config to the firewall and returns a func that
|
||||
// restores the previous ruleset. Abstracted so the run loop is testable without
|
||||
// touching the kernel.
|
||||
type Applier interface {
|
||||
Apply(ctx context.Context, cfg *config.Config) error
|
||||
Apply(ctx context.Context, cfg *config.Config) (revert func() error, err error)
|
||||
}
|
||||
|
||||
// Agent runs the pull-apply-report loop for one device.
|
||||
@@ -24,6 +27,10 @@ type Agent struct {
|
||||
Applier Applier
|
||||
// Resolver overrides the DNS resolver (tests); nil derives it per-config.
|
||||
Resolver *Resolver
|
||||
|
||||
// reverted is the last generation rolled back for cutting off the API; it
|
||||
// is not re-applied until a newer generation is published.
|
||||
reverted int64
|
||||
}
|
||||
|
||||
// Run loops until ctx is cancelled, applying one cycle per Interval (and once
|
||||
@@ -61,16 +68,20 @@ func (a *Agent) RunOnce(ctx context.Context) error {
|
||||
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.
|
||||
return a.applyConfig(ctx, cached, false)
|
||||
return a.applyConfig(ctx, cached, nil)
|
||||
}
|
||||
|
||||
if err := a.Cache.Write(raw); err != nil {
|
||||
slog.Warn("agent: caching config failed", "err", err)
|
||||
if a.reverted != 0 && rc.Generation == a.reverted {
|
||||
slog.Warn("agent: skipping reverted generation", "generation", rc.Generation)
|
||||
return nil
|
||||
}
|
||||
return a.applyConfig(ctx, rc, true)
|
||||
return a.applyConfig(ctx, rc, raw)
|
||||
}
|
||||
|
||||
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error {
|
||||
// applyConfig applies rc. A freshly fetched config (raw != nil) is verified by
|
||||
// reporting status over a new connection through the new ruleset; if the API is
|
||||
// unreachable the previous ruleset is restored. Only verified configs are cached.
|
||||
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error {
|
||||
resolver := a.Resolver
|
||||
if resolver == nil {
|
||||
resolver = NewResolver(rc.Resolver)
|
||||
@@ -81,15 +92,28 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
|
||||
if err != nil {
|
||||
return fmt.Errorf("translate: %w", err)
|
||||
}
|
||||
if err := a.Applier.Apply(ctx, cfg); err != nil {
|
||||
revert, err := a.Applier.Apply(ctx, cfg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("apply: %w", err)
|
||||
}
|
||||
slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules))
|
||||
|
||||
if report {
|
||||
if raw != nil {
|
||||
a.Client.HTTP.CloseIdleConnections()
|
||||
if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil {
|
||||
var uerr *url.Error
|
||||
if errors.As(err, &uerr) {
|
||||
if rerr := revert(); rerr != nil {
|
||||
return fmt.Errorf("API unreachable after applying generation %d (%v); revert failed: %w", rc.Generation, err, rerr)
|
||||
}
|
||||
a.reverted = rc.Generation
|
||||
return fmt.Errorf("reverted generation %d: API unreachable through new ruleset: %w", rc.Generation, err)
|
||||
}
|
||||
slog.Warn("agent: reporting status failed", "err", err)
|
||||
}
|
||||
if err := a.Cache.Write(raw); err != nil {
|
||||
slog.Warn("agent: caching config 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 {
|
||||
@@ -103,18 +127,26 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
|
||||
// EngineApplier applies via the real nftables differential engine.
|
||||
type EngineApplier struct{}
|
||||
|
||||
// Apply computes and applies the differential change set for cfg.
|
||||
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error {
|
||||
// Apply snapshots the live table, then computes and applies the differential
|
||||
// change set for cfg. The returned func atomically restores the snapshot.
|
||||
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) {
|
||||
engine, err := nftables.NewEngine(cfg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("initializing nftables: %w", err)
|
||||
return nil, fmt.Errorf("initializing nftables: %w", err)
|
||||
}
|
||||
changes, err := engine.Plan()
|
||||
if err != nil {
|
||||
return fmt.Errorf("computing changes: %w", err)
|
||||
return nil, fmt.Errorf("computing changes: %w", err)
|
||||
}
|
||||
if changes.Empty() {
|
||||
return nil
|
||||
return func() error { return nil }, nil
|
||||
}
|
||||
return engine.Apply(changes)
|
||||
snap, err := engine.Snapshot()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("snapshotting ruleset: %w", err)
|
||||
}
|
||||
if err := engine.Apply(changes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return func() error { return engine.Restore(snap) }, nil
|
||||
}
|
||||
|
||||
@@ -126,16 +126,17 @@ func TestTranslateRejectsUnknownAction(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// fakeApplier records applied configs.
|
||||
// fakeApplier records applied configs and reverts.
|
||||
type fakeApplier struct {
|
||||
count int32
|
||||
reverts int32
|
||||
lastGen int
|
||||
}
|
||||
|
||||
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error {
|
||||
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) {
|
||||
atomic.AddInt32(&f.count, 1)
|
||||
f.lastGen = len(cfg.Rules)
|
||||
return nil
|
||||
return func() error { atomic.AddInt32(&f.reverts, 1); return nil }, nil
|
||||
}
|
||||
|
||||
const renderedYAML = `generation: 7
|
||||
@@ -238,3 +239,71 @@ func TestRunOnceNoCacheReturnsError(t *testing.T) {
|
||||
t.Fatal("expected error when unreachable and no cache exists")
|
||||
}
|
||||
}
|
||||
|
||||
// statusServer serves renderedYAML and answers status reports with handler.
|
||||
func statusServer(t *testing.T, status http.HandlerFunc) *httptest.Server {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodGet {
|
||||
_, _ = w.Write([]byte(renderedYAML))
|
||||
return
|
||||
}
|
||||
status(w, r)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv
|
||||
}
|
||||
|
||||
func TestRunOnceRevertsWhenAPIUnreachableAfterApply(t *testing.T) {
|
||||
// Dropping the connection mimics a pushed rule that severs the API path.
|
||||
srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, _, _ := w.(http.Hijacker).Hijack()
|
||||
conn.Close()
|
||||
})
|
||||
|
||||
applier := &fakeApplier{}
|
||||
a := &Agent{
|
||||
Client: NewClient(srv.URL, "fw-a", "tok"),
|
||||
Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")},
|
||||
Applier: applier,
|
||||
}
|
||||
if err := a.RunOnce(context.Background()); err == nil {
|
||||
t.Fatal("expected revert error")
|
||||
}
|
||||
if applier.reverts != 1 {
|
||||
t.Errorf("expected 1 revert, got %d", applier.reverts)
|
||||
}
|
||||
if cached, _ := a.Cache.Read(); cached != nil {
|
||||
t.Error("reverted config must not be cached as known-good")
|
||||
}
|
||||
|
||||
// The same generation is not re-applied on the next cycle.
|
||||
if err := a.RunOnce(context.Background()); err != nil {
|
||||
t.Fatalf("second RunOnce: %v", err)
|
||||
}
|
||||
if applier.count != 1 {
|
||||
t.Errorf("reverted generation re-applied: %d applies", applier.count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunOnceKeepsConfigOnAPIErrorStatus(t *testing.T) {
|
||||
// An HTTP error still proves the API is reachable — no revert.
|
||||
srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
})
|
||||
|
||||
applier := &fakeApplier{}
|
||||
a := &Agent{
|
||||
Client: NewClient(srv.URL, "fw-a", "tok"),
|
||||
Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")},
|
||||
Applier: applier,
|
||||
}
|
||||
if err := a.RunOnce(context.Background()); err != nil {
|
||||
t.Fatalf("RunOnce: %v", err)
|
||||
}
|
||||
if applier.reverts != 0 {
|
||||
t.Errorf("unexpected revert")
|
||||
}
|
||||
if cached, _ := a.Cache.Read(); cached == nil {
|
||||
t.Error("expected config cached")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
@@ -84,6 +85,32 @@ func computeDiff(current, desired *FirewallState) *ChangeSet {
|
||||
return cs
|
||||
}
|
||||
|
||||
// restoreChangeSet replaces every managed rule in current with the snapshot's,
|
||||
// in snapshot order, so a restore cannot reorder rules.
|
||||
func restoreChangeSet(current, snap *FirewallState) *ChangeSet {
|
||||
cs := &ChangeSet{}
|
||||
for _, rules := range current.Rules {
|
||||
for _, r := range rules {
|
||||
if r.Tag != "" {
|
||||
cs.Remove = append(cs.Remove, r)
|
||||
}
|
||||
}
|
||||
}
|
||||
chains := make([]string, 0, len(snap.Rules))
|
||||
for c := range snap.Rules {
|
||||
chains = append(chains, c)
|
||||
}
|
||||
sort.Strings(chains)
|
||||
for _, c := range chains {
|
||||
for _, r := range snap.Rules[c] {
|
||||
if r.Tag != "" {
|
||||
cs.Add = append(cs.Add, r)
|
||||
}
|
||||
}
|
||||
}
|
||||
return cs
|
||||
}
|
||||
|
||||
func rulesMatch(a, b []ManagedRule) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
)
|
||||
|
||||
func TestRestoreChangeSet(t *testing.T) {
|
||||
accept := []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}}
|
||||
drop := []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
|
||||
|
||||
snap := &FirewallState{Rules: map[string][]ManagedRule{
|
||||
"input": {
|
||||
{Chain: "input", Tag: "ssh", Exprs: accept, Handle: 4},
|
||||
{Chain: "input", Tag: "web", Exprs: accept, Handle: 5},
|
||||
{Chain: "input", Tag: "", Exprs: drop, Handle: 6},
|
||||
},
|
||||
"forward": {{Chain: "forward", Tag: "fwd", Exprs: accept, Handle: 7}},
|
||||
}}
|
||||
current := &FirewallState{Rules: map[string][]ManagedRule{
|
||||
"input": {
|
||||
{Chain: "input", Tag: "web", Exprs: accept, Handle: 10},
|
||||
{Chain: "input", Tag: "ssh", Exprs: drop, Handle: 11},
|
||||
{Chain: "input", Tag: "", Exprs: drop, Handle: 12},
|
||||
},
|
||||
}}
|
||||
|
||||
cs := restoreChangeSet(current, snap)
|
||||
|
||||
var removed []uint64
|
||||
for _, r := range cs.Remove {
|
||||
removed = append(removed, r.Handle)
|
||||
}
|
||||
if len(removed) != 2 || removed[0] != 10 || removed[1] != 11 {
|
||||
t.Errorf("expected managed handles [10 11] removed, untagged kept; got %v", removed)
|
||||
}
|
||||
|
||||
var added []string
|
||||
for _, r := range cs.Add {
|
||||
added = append(added, r.Tag)
|
||||
}
|
||||
want := []string{"fwd", "ssh", "web"}
|
||||
if len(added) != len(want) {
|
||||
t.Fatalf("added %v, want %v", added, want)
|
||||
}
|
||||
for i := range want {
|
||||
if added[i] != want[i] {
|
||||
t.Fatalf("added %v, want %v (snapshot order per chain)", added, want)
|
||||
}
|
||||
}
|
||||
if !exprsEqual(cs.Add[1].Exprs, accept) {
|
||||
t.Error("ssh not restored to its snapshot exprs")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) {
|
||||
current := &FirewallState{Rules: map[string][]ManagedRule{
|
||||
"input": {{Chain: "input", Tag: "x", Handle: 1}},
|
||||
}}
|
||||
cs := restoreChangeSet(current, &FirewallState{Rules: map[string][]ManagedRule{}})
|
||||
if len(cs.Remove) != 1 || len(cs.Add) != 0 {
|
||||
t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add))
|
||||
}
|
||||
}
|
||||
+48
-17
@@ -135,31 +135,32 @@ func (e *Engine) Flush() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Engine) findTable() (*nftables.Table, error) {
|
||||
tables, err := e.conn.ListTables()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("listing tables: %w", err)
|
||||
}
|
||||
for _, t := range tables {
|
||||
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet {
|
||||
return t, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (e *Engine) readCurrentState() (*FirewallState, error) {
|
||||
state := &FirewallState{
|
||||
Rules: make(map[string][]ManagedRule),
|
||||
}
|
||||
|
||||
tables, err := e.conn.ListTables()
|
||||
if err != nil {
|
||||
return state, nil
|
||||
}
|
||||
|
||||
var ourTable *nftables.Table
|
||||
for _, t := range tables {
|
||||
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet {
|
||||
ourTable = t
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if ourTable == nil {
|
||||
return state, nil
|
||||
ourTable, err := e.findTable()
|
||||
if err != nil || ourTable == nil {
|
||||
return state, err
|
||||
}
|
||||
|
||||
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet)
|
||||
if err != nil {
|
||||
return state, nil
|
||||
return nil, fmt.Errorf("listing chains: %w", err)
|
||||
}
|
||||
|
||||
for _, chain := range chains {
|
||||
@@ -168,7 +169,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
|
||||
}
|
||||
rules, err := e.conn.GetRules(ourTable, chain)
|
||||
if err != nil {
|
||||
continue
|
||||
return nil, fmt.Errorf("listing rules of %s: %w", chain.Name, err)
|
||||
}
|
||||
for _, rule := range rules {
|
||||
state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{
|
||||
@@ -183,6 +184,36 @@ func (e *Engine) readCurrentState() (*FirewallState, error) {
|
||||
return state, nil
|
||||
}
|
||||
|
||||
// Snapshot is the tomswall table as captured live; a nil state means it was absent.
|
||||
type Snapshot struct {
|
||||
state *FirewallState
|
||||
}
|
||||
|
||||
// Snapshot captures the live tomswall table so Restore can roll back to it.
|
||||
func (e *Engine) Snapshot() (*Snapshot, error) {
|
||||
t, err := e.findTable()
|
||||
if err != nil || t == nil {
|
||||
return &Snapshot{}, err
|
||||
}
|
||||
state, err := e.readCurrentState()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Snapshot{state: state}, nil
|
||||
}
|
||||
|
||||
// Restore atomically returns the tomswall table to the snapshot, rule order included.
|
||||
func (e *Engine) Restore(s *Snapshot) error {
|
||||
if s.state == nil {
|
||||
return e.Flush()
|
||||
}
|
||||
current, err := e.readCurrentState()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return e.Apply(restoreChangeSet(current, s.state))
|
||||
}
|
||||
|
||||
func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy {
|
||||
return &p
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user