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 {