Apply the cached config without safe-apply and keep reverted generations in memory
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

This commit is contained in:
2026-10-05 13:56:11 +11:00
parent 9092b463a0
commit c7e02c089a
4 changed files with 117 additions and 27 deletions
+26 -16
View File
@@ -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)
+1 -1
View File
@@ -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
+88 -8
View File
@@ -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)
}
}
+2 -2
View File
@@ -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 {