From 4ad55fc65e0e65d98bd81e2f91a8a92fa0d243ca Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Mon, 5 Oct 2026 13:46:20 +1100 Subject: [PATCH 1/3] Revert agent generations that cut off the control plane --- DESIGN.md | 4 + cmd/tomswall/agent.go | 3 +- internal/agent/agent.go | 203 +++++++++++++++++++++---- internal/agent/agent_test.go | 4 +- internal/agent/client.go | 30 +++- internal/agent/safeapply_test.go | 249 +++++++++++++++++++++++++++++++ 6 files changed, 458 insertions(+), 35 deletions(-) create mode 100644 internal/agent/safeapply_test.go diff --git a/DESIGN.md b/DESIGN.md index 3253925..cf30d34 100644 --- a/DESIGN.md +++ b/DESIGN.md @@ -430,6 +430,10 @@ report the generation applied, giving a fleet-wide "converged / N behind" view. source/dest disables its rule loudly, never opens it. - **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. +- **A generation that severs the API is reverted.** After applying, the agent + reports `applied` over a fresh connection; if that fails at the transport level + it restores the snapshot, reports `reverted`, and skips that generation (persisted + in `/var/lib/tomswall/reverted.json`) until a newer one arrives. --- diff --git a/cmd/tomswall/agent.go b/cmd/tomswall/agent.go index 6f91379..3cdc2b6 100644 --- a/cmd/tomswall/agent.go +++ b/cmd/tomswall/agent.go @@ -29,7 +29,8 @@ func agentCmd() *cobra.Command { Long: `Agent runs the control-plane pull loop: it fetches this device's compiled config from tomswallapi, differentially applies it, and reports the applied 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 device name defaults to the system hostname.`, diff --git a/internal/agent/agent.go b/internal/agent/agent.go index 868c4e0..89c0870 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -2,8 +2,13 @@ package agent import ( "context" + "encoding/json" + "errors" "fmt" "log/slog" + "net/url" + "os" + "path/filepath" "time" "git.unkin.net/unkin/tomswall/internal/config" @@ -12,9 +17,10 @@ import ( ) // 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. restore returns the firewall to +// its state before the call; it is nil when nothing changed. type Applier interface { - Apply(ctx context.Context, cfg *config.Config) error + Apply(ctx context.Context, cfg *config.Config) (restore func() error, err error) } // Agent runs the pull-apply-report loop for one device. @@ -62,16 +68,26 @@ 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) + rv, err := a.readReverted() + if err != nil { + return err } - return a.applyConfig(ctx, rc, true) + 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 verified by reaching +// the API through the new ruleset and reverted if that fails; only then is it cached. +func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error { resolver := a.Resolver if resolver == nil { resolver = NewResolver(rc.Resolver) @@ -82,46 +98,177 @@ 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 { + + 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() + + restore, err := a.Applier.Apply(ctx, cfg) + if err != nil { + if restore != nil { + if rerr := restore(); rerr != nil { + err = fmt.Errorf("%w; restoring previous ruleset: %v", 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 fmt.Errorf("apply: %w", 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 ctx.Err() != nil || restore == nil { + return err + } + if rerr := restore(); rerr != nil { + return fmt.Errorf("%w; restoring previous ruleset: %v", err, rerr) + } + rv := &reverted{Generation: rc.Generation, Error: err.Error()} + 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 reverted: %w", rc.Generation, err) + } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) - if report { - if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { - slog.Warn("agent: reporting status 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) - } + if err := a.Cache.Write(raw); err != nil { + slog.Warn("agent: caching config failed", "err", err) + } + if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) { + slog.Warn("agent: clearing reverted generation failed", "err", err) + } + // Report the FIB so the control plane can scope router enforcement. + 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 } -// EngineApplier applies via the real nftables differential engine. -type EngineApplier struct{} +var ( + verifyAttempts = 3 + verifyDelay = 2 * time.Second + verifyTimeout = 5 * time.Second +) -// Apply computes and applies the differential change set for cfg. It refuses -// while a 'tomswall try' awaits confirmation. -func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { - unlock, err := tryapply.Acquire() +// 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 for severing the API; persisted so it +// is not re-applied after a restart until a newer generation arrives. +type reverted struct { + Generation int64 `json:"generation"` + Error string `json:"error,omitempty"` + Reported bool `json:"reported"` +} + +func (a *Agent) revertedPath() string { + return filepath.Join(filepath.Dir(a.Cache.Path), "reverted.json") +} + +func (a *Agent) readReverted() (*reverted, error) { + b, err := os.ReadFile(a.revertedPath()) + if os.IsNotExist(err) { + return nil, nil + } + if err != nil { + return nil, err + } + var rv reverted + if err := json.Unmarshal(b, &rv); err != nil { + return nil, fmt.Errorf("parsing %s: %w", a.revertedPath(), err) + } + return &rv, nil +} + +func (a *Agent) writeReverted(rv *reverted) error { + b, err := json.Marshal(rv) if err != nil { return err } - defer unlock() + return Cache{Path: a.revertedPath()}.Write(b) +} + +// reportReverted reports rv until the control plane accepts it. +func (a *Agent) reportReverted(ctx context.Context, rv *reverted) { + if rv.Reported { + return + } + if err := a.Client.ReportStatus(ctx, Status{Status: StatusReverted, 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, snapshotting +// the live table first so the caller can restore it. The caller holds the try lock. +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 nil, nil } - return engine.Apply(changes) + snap, err := engine.Snapshot() + if err != nil { + return nil, fmt.Errorf("snapshotting ruleset: %w", err) + } + restore := func() error { return engine.Restore(snap) } + return restore, engine.Apply(changes) } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 3019815..1cafc57 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -132,10 +132,10 @@ type fakeApplier struct { 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 nil, nil } const renderedYAML = `generation: 7 diff --git a/internal/agent/client.go b/internal/agent/client.go index d385011..0bec573 100644 --- a/internal/agent/client.go +++ b/internal/agent/client.go @@ -26,7 +26,9 @@ func NewClient(baseURL, device, token string) *Client { BaseURL: baseURL, Device: device, 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 } -// ReportStatus tells the control plane which generation this device has applied. -func (c *Client) ReportStatus(ctx context.Context, generation int64) error { +// Status values reported to POST /api/v1/devices/{name}/status. +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) - payload, _ := json.Marshal(map[string]int64{"generation": generation}) + payload, _ := json.Marshal(st) req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload)) if err != nil { return err @@ -116,3 +132,9 @@ func (c *Client) ReportStatus(ctx context.Context, generation int64) error { } return nil } + +func noKeepAlive() http.RoundTripper { + t := http.DefaultTransport.(*http.Transport).Clone() + t.DisableKeepAlives = true + return t +} diff --git a/internal/agent/safeapply_test.go b/internal/agent/safeapply_test.go new file mode 100644 index 0000000..bb38054 --- /dev/null +++ b/internal/agent/safeapply_test.go @@ -0,0 +1,249 @@ +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/tryapply" +) + +func TestMain(m *testing.M) { + dir, err := os.MkdirTemp("", "tomswall-agent-test") + if err != nil { + panic(err) + } + tryapply.Dir = dir + 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 +} + +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; onApply simulates its effect. +type fakeEngine struct { + applies, restores int + err error + onApply func() + onRestore func() +} + +func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, error) { + f.applies++ + if f.onApply != nil { + f.onApply() + } + return func() error { + f.restores++ + if f.onRestore != nil { + f.onRestore() + } + return nil + }, f.err +} + +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 { + 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 { + 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") { + t.Fatalf("restores=%d last=%+v", eng.restores, st) + } +} + +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()) + } +} From 9092b463a09e5e7230ff791d00e188a900bdaa27 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Mon, 5 Oct 2026 13:52:32 +1100 Subject: [PATCH 2/3] Apply agent generations as a pending try with a revert timer --- DESIGN.md | 11 +-- internal/agent/agent.go | 105 ++++++++++++++++++++--------- internal/agent/agent_test.go | 4 +- internal/agent/cache.go | 14 ++-- internal/agent/safeapply_test.go | 89 +++++++++++++++++++++--- internal/tryapply/tryapply.go | 65 +++++++++++------- internal/tryapply/tryapply_test.go | 14 ++-- 7 files changed, 214 insertions(+), 88 deletions(-) diff --git a/DESIGN.md b/DESIGN.md index cf30d34..ea91ed0 100644 --- a/DESIGN.md +++ b/DESIGN.md @@ -430,10 +430,13 @@ report the generation applied, giving a fleet-wide "converged / N behind" view. source/dest disables its rule loudly, never opens it. - **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. -- **A generation that severs the API is reverted.** After applying, the agent - reports `applied` over a fresh connection; if that fails at the transport level - it restores the snapshot, reports `reverted`, and skips that generation (persisted - in `/var/lib/tomswall/reverted.json`) until a newer one arrives. +- **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`. --- diff --git a/internal/agent/agent.go b/internal/agent/agent.go index 89c0870..a5711ac 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -17,12 +17,16 @@ import ( ) // Applier applies a translated config to the firewall. Abstracted so the run -// loop is testable without touching the kernel. restore returns the firewall to -// its state before the call; it is nil when nothing changed. +// loop is testable without touching the kernel. 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. type Applier interface { - Apply(ctx context.Context, cfg *config.Config) (restore func() error, err error) + Apply(ctx context.Context, cfg *config.Config) (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. type Agent struct { Client *Client @@ -109,11 +113,15 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) } defer unlock() - restore, err := a.Applier.Apply(ctx, cfg) + revert, keep, err := a.Applier.Apply(ctx, cfg) if err != nil { - if restore != nil { - if rerr := restore(); rerr != nil { - err = fmt.Errorf("%w; restoring previous ruleset: %v", err, rerr) + 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 { @@ -121,26 +129,34 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) slog.Warn("agent: reporting status failed", "err", rerr) } } - return fmt.Errorf("apply: %w", err) + return err } if raw == nil { + if keep != nil { + if err := keep(); err != nil { + return fmt.Errorf("dropping snapshot: %w", err) + } + } 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 ctx.Err() != nil || restore == nil { - return err + if keep != nil && ctx.Err() == nil { + return a.revertGeneration(ctx, rc.Generation, StatusReverted, err, revert) } - if rerr := restore(); rerr != nil { - return fmt.Errorf("%w; restoring previous ruleset: %v", err, rerr) + // 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) + } } - rv := &reverted{Generation: rc.Generation, Error: err.Error()} - if werr := a.writeReverted(rv); werr != nil { - slog.Error("agent: persisting reverted generation failed", "err", werr) + return err + } + if keep != nil { + if err := keep(); err != nil { + return fmt.Errorf("dropping snapshot: %w", err) } - a.reportReverted(ctx, rv) - return fmt.Errorf("generation %d reverted: %w", rc.Generation, err) } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) @@ -159,6 +175,27 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) return nil } +// revertGeneration marks generation as reverted before restoring, so a failed +// restore can never lead to re-applying it, then reports status. 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()} + if err := a.writeReverted(rv); err != nil { + slog.Error("agent: persisting reverted generation failed", "err", err) + } + suffix := "" + if rerr := revert(); rerr != nil { + suffix = fmt.Sprintf("; restore: %v; revert timer pending", rerr) + rv.Status = StatusFailed + rv.Error += suffix + if werr := a.writeReverted(rv); werr != nil { + slog.Error("agent: persisting reverted generation failed", "err", werr) + } + } + a.reportReverted(ctx, rv) + return fmt.Errorf("generation %d %s: %w%s", generation, rv.Status, cause, suffix) +} + var ( verifyAttempts = 3 verifyDelay = 2 * time.Second @@ -198,10 +235,11 @@ func (a *Agent) confirm(ctx context.Context, generation int64) error { return fmt.Errorf("%w: %v", errUnreachable, err) } -// reverted is a generation rolled back for severing the API; persisted so it -// is not re-applied after a restart until a newer generation arrives. +// 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"` } @@ -230,7 +268,7 @@ func (a *Agent) writeReverted(rv *reverted) error { if err != nil { return err } - return Cache{Path: a.revertedPath()}.Write(b) + return tryapply.WriteFile(a.revertedPath(), b) } // reportReverted reports rv until the control plane accepts it. @@ -238,7 +276,11 @@ func (a *Agent) reportReverted(ctx context.Context, rv *reverted) { if rv.Reported { return } - if err := a.Client.ReportStatus(ctx, Status{Status: StatusReverted, Generation: rv.Generation, Error: rv.Error}); err != nil { + 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 } @@ -251,24 +293,27 @@ func (a *Agent) reportReverted(ctx context.Context, rv *reverted) { // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} -// Apply computes and applies the differential change set for cfg, snapshotting -// the live table first so the caller can restore it. The caller holds the try lock. -func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { +// Apply computes and applies the differential change set for cfg under a +// pending try, as 'tomswall try' does. The caller holds the try lock. +func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (revert, keep func() error, err error) { engine, err := nftables.NewEngine(cfg) if err != nil { - return nil, fmt.Errorf("initializing nftables: %w", err) + return nil, nil, fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { - return nil, fmt.Errorf("computing changes: %w", err) + return nil, nil, fmt.Errorf("computing changes: %w", err) } if changes.Empty() { - return nil, nil + return nil, nil, nil } snap, err := engine.Snapshot() if err != nil { - return nil, fmt.Errorf("snapshotting ruleset: %w", err) + return nil, nil, fmt.Errorf("snapshotting ruleset: %w", err) } - restore := func() error { return engine.Restore(snap) } - return restore, engine.Apply(changes) + // 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) } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 1cafc57..7921bb9 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -132,10 +132,10 @@ type fakeApplier struct { lastGen int } -func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { +func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, func() error, error) { atomic.AddInt32(&f.count, 1) f.lastGen = len(cfg.Rules) - return nil, nil + return nil, nil, nil } const renderedYAML = `generation: 7 diff --git a/internal/agent/cache.go b/internal/agent/cache.go index 1fb439b..533079b 100644 --- a/internal/agent/cache.go +++ b/internal/agent/cache.go @@ -2,7 +2,8 @@ package agent import ( "os" - "path/filepath" + + "git.unkin.net/unkin/tomswall/internal/tryapply" ) // Cache persists the last known-good rendered config to disk so the agent can @@ -11,16 +12,9 @@ type Cache struct { Path string } -// Write atomically stores the raw config bytes. +// Write durably stores the raw config bytes. func (c Cache) Write(raw []byte) error { - if err := os.MkdirAll(filepath.Dir(c.Path), 0o755); err != nil { - return err - } - tmp := c.Path + ".tmp" - if err := os.WriteFile(tmp, raw, 0o600); err != nil { - return err - } - return os.Rename(tmp, c.Path) + return tryapply.WriteFile(c.Path, raw) } // Read returns the cached config, or (nil, nil) when no cache exists yet. diff --git a/internal/agent/safeapply_test.go b/internal/agent/safeapply_test.go index bb38054..2c11791 100644 --- a/internal/agent/safeapply_test.go +++ b/internal/agent/safeapply_test.go @@ -15,6 +15,7 @@ import ( "time" "git.unkin.net/unkin/tomswall/internal/config" + "git.unkin.net/unkin/tomswall/internal/nftables" "git.unkin.net/unkin/tomswall/internal/tryapply" ) @@ -24,6 +25,10 @@ func TestMain(m *testing.M) { 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() @@ -70,6 +75,9 @@ func newFakeAPI(t *testing.T, gen int64) *fakeAPI { 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 { @@ -81,26 +89,42 @@ func (f *fakeAPI) last() Status { return f.reports[len(f.reports)-1] } -// fakeEngine always changes the ruleset; onApply simulates its effect. +// fakeEngine always changes the ruleset under a real tryapply pending try; +// onApply simulates its effect and restoreErr fails the restore. type fakeEngine struct { applies, restores int - err error + err, restoreErr error onApply func() onRestore func() } -func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, error) { +func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, func() error, error) { f.applies++ - if f.onApply != nil { - f.onApply() + if _, err := tryapply.Arm(&nftables.Snapshot{Table: "tomswall"}, 0, time.Minute); err != nil { + return nil, nil, err } - return func() error { + tryapply.Restore = func(*nftables.Snapshot) error { f.restores++ if f.onRestore != nil { f.onRestore() } - return nil - }, f.err + return f.restoreErr + } + 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 { @@ -129,7 +153,7 @@ func TestSafeApplyReachableApplies(t *testing.T) { 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 { + 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)) } } @@ -141,7 +165,7 @@ func TestSafeApplyUnreachableRevertsAndReports(t *testing.T) { 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 { + 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 == "" { @@ -203,9 +227,52 @@ func TestSafeApplyApplyErrorRestoresAndReportsFailed(t *testing.T) { 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") { + 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) { diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go index a68e024..e18dcbc 100644 --- a/internal/tryapply/tryapply.go +++ b/internal/tryapply/tryapply.go @@ -25,15 +25,15 @@ const Unit = "tomswall-try-revert" var ( // Dir holds the lock and the pending snapshot. Dir = "/var/lib/tomswall" - // run executes a systemd command; replaced in tests. - run = func(name string, args ...string) error { + // Run executes a systemd command; replaced in tests. + Run = func(name string, args ...string) error { if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) } return nil } - // restore rolls the live table back to a snapshot; replaced in tests. - restore = func(s *nftables.Snapshot) error { + // Restore rolls the live table back to a snapshot; replaced in tests. + Restore = func(s *nftables.Snapshot) error { engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) if err != nil { return err @@ -94,23 +94,7 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error) if err != nil { return "", err } - f, err := os.CreateTemp(Dir, ".try-snapshot-*") - 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 { + if err := WriteFile(snapshotPath(), b); err != nil { return "", err } @@ -119,13 +103,46 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error) return "", discardWith(err) } _ = disarm() // a leftover timer from an earlier try would block the unit name - if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, + if err := Run("systemd-run", "--quiet", "--collect", "--unit", Unit, fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil { return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) } return id, nil } +// WriteFile durably replaces path with b: temp file, fsync, rename, fsync the directory. +func WriteFile(path string, b []byte) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0o755); err != nil { + return err + } + f, err := os.CreateTemp(dir, "."+filepath.Base(path)+"-*") + if err != nil { + return err + } + defer os.Remove(f.Name()) + if _, err := f.Write(b); err != nil { + f.Close() + return err + } + if err := f.Sync(); err != nil { + f.Close() + return err + } + if err := f.Close(); err != nil { + return err + } + if err := os.Rename(f.Name(), path); err != nil { + return err + } + d, err := os.Open(dir) + if err != nil { + return err + } + defer d.Close() + return d.Sync() +} + // Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. func Discard() error { _ = disarm() @@ -143,7 +160,7 @@ func discardWith(err error) 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. @@ -193,7 +210,7 @@ func Abort() 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 Discard() diff --git a/internal/tryapply/tryapply_test.go b/internal/tryapply/tryapply_test.go index e82dc83..d39de95 100644 --- a/internal/tryapply/tryapply_test.go +++ b/internal/tryapply/tryapply_test.go @@ -15,12 +15,12 @@ func setup(t *testing.T) *[]string { t.Helper() Dir = t.TempDir() var cmds []string - orig := run - run = func(name string, args ...string) error { + orig := Run + Run = func(name string, args ...string) error { cmds = append(cmds, name+" "+strings.Join(args, " ")) return nil } - t.Cleanup(func() { run = orig }) + t.Cleanup(func() { Run = orig }) return &cmds } @@ -91,7 +91,7 @@ func TestAcquireRefusesWhilePending(t *testing.T) { func TestArmFailureDiscardsSnapshot(t *testing.T) { setup(t) - run = func(name string, args ...string) error { + Run = func(name string, args ...string) error { if name == "systemd-run" { return errors.New("no systemd") } @@ -145,12 +145,12 @@ func TestConfirmAfterRevertFails(t *testing.T) { func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot { t.Helper() var got []*nftables.Snapshot - orig := restore - restore = func(s *nftables.Snapshot) error { + orig := Restore + Restore = func(s *nftables.Snapshot) error { got = append(got, s) return err } - t.Cleanup(func() { restore = orig }) + t.Cleanup(func() { Restore = orig }) return &got } From c7e02c089aaa4f172b357d96ff6e646931b44832 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Mon, 5 Oct 2026 13:56:11 +1100 Subject: [PATCH 3/3] Apply the cached config without safe-apply and keep reverted generations in memory --- internal/agent/agent.go | 42 ++++++++------ internal/agent/agent_test.go | 2 +- internal/agent/safeapply_test.go | 96 +++++++++++++++++++++++++++++--- internal/tryapply/tryapply.go | 4 +- 4 files changed, 117 insertions(+), 27 deletions(-) diff --git a/internal/agent/agent.go b/internal/agent/agent.go index a5711ac..2071e58 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -17,11 +17,12 @@ import ( ) // Applier applies a translated config to the firewall. Abstracted so the run -// loop is testable without touching the kernel. 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. +// 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 { - Apply(ctx context.Context, cfg *config.Config) (revert, keep func() error, err 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. @@ -35,6 +36,9 @@ type Agent struct { Applier Applier // Resolver overrides the DNS resolver (tests); nil derives it per-config. 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 @@ -79,6 +83,9 @@ func (a *Agent) RunOnce(ctx context.Context) error { if err != nil { return err } + 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 { @@ -89,8 +96,10 @@ func (a *Agent) RunOnce(ctx context.Context) error { return a.applyConfig(ctx, rc, raw) } -// applyConfig applies rc. A fetched config (raw != nil) is verified by reaching -// the API through the new ruleset and reverted if that fails; only then is it cached. +// 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 if resolver == nil { @@ -113,7 +122,7 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) } defer unlock() - revert, keep, err := a.Applier.Apply(ctx, cfg) + revert, keep, err := a.Applier.Apply(ctx, cfg, raw != nil) if err != nil { err = fmt.Errorf("apply: %w", err) if raw != nil && revert != nil { @@ -132,11 +141,6 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) return err } if raw == nil { - if keep != nil { - if err := keep(); err != nil { - return fmt.Errorf("dropping snapshot: %w", err) - } - } slog.Info("agent: applied cached config", "generation", rc.Generation, "rules", len(cfg.Rules)) return nil } @@ -163,6 +167,7 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) if err := a.Cache.Write(raw); err != nil { slog.Warn("agent: caching config failed", "err", err) } + a.lastReverted = nil if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) { slog.Warn("agent: clearing reverted generation failed", "err", err) } @@ -176,10 +181,12 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) } // revertGeneration marks generation as reverted before restoring, so a failed -// restore can never lead to re-applying it, then reports status. A failed +// restore can never lead to re-applying it, then reports status. It is also kept +// in memory in case persisting fails. A failed // restore leaves the snapshot and timer armed and is reported as failed. func (a *Agent) revertGeneration(ctx context.Context, generation int64, status string, cause error, revert func() error) error { rv := &reverted{Generation: generation, Status: status, Error: cause.Error()} + a.lastReverted = rv if err := a.writeReverted(rv); err != nil { slog.Error("agent: persisting reverted generation failed", "err", err) } @@ -293,9 +300,9 @@ func (a *Agent) reportReverted(ctx context.Context, rv *reverted) { // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} -// Apply computes and applies the differential change set for cfg under a -// pending try, as 'tomswall try' does. The caller holds the try lock. -func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (revert, keep func() error, err error) { +// 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) if err != nil { return nil, nil, fmt.Errorf("initializing nftables: %w", err) @@ -307,6 +314,9 @@ func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (revert, keep if changes.Empty() { return nil, nil, nil } + if !safe { + return nil, nil, engine.Apply(changes) + } snap, err := engine.Snapshot() if err != nil { return nil, nil, fmt.Errorf("snapshotting ruleset: %w", err) diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 7921bb9..8f5beb7 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -132,7 +132,7 @@ type fakeApplier struct { lastGen int } -func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, func() error, error) { +func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config, _ bool) (func() error, func() error, error) { atomic.AddInt32(&f.count, 1) f.lastGen = len(cfg.Rules) return nil, nil, nil diff --git a/internal/agent/safeapply_test.go b/internal/agent/safeapply_test.go index 2c11791..8d83cc8 100644 --- a/internal/agent/safeapply_test.go +++ b/internal/agent/safeapply_test.go @@ -89,17 +89,20 @@ func (f *fakeAPI) last() Status { return f.reports[len(f.reports)-1] } -// fakeEngine always changes the ruleset under a real tryapply pending try; -// onApply simulates its effect and restoreErr fails the restore. +// 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, restores int - err, restoreErr error - onApply func() - onRestore func() + applies, plain, restores int + err, restoreErr error + onApply func() + onRestore func() } -func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, func() error, error) { - f.applies++ +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 } @@ -110,6 +113,7 @@ func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, func( } return f.restoreErr } + f.applies++ if f.onApply != nil { f.onApply() } @@ -314,3 +318,79 @@ func TestSafeApplySkipsWhileTryPending(t *testing.T) { 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) + } +} diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go index e18dcbc..584f4c3 100644 --- a/internal/tryapply/tryapply.go +++ b/internal/tryapply/tryapply.go @@ -25,14 +25,14 @@ const Unit = "tomswall-try-revert" var ( // Dir holds the lock and the pending snapshot. 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 { if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) } return nil } - // Restore rolls the live table back to a snapshot; 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 { engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) if err != nil {