Apply agent generations as a pending try with a revert timer
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:52:32 +11:00
parent 4ad55fc65e
commit 9092b463a0
7 changed files with 214 additions and 88 deletions
+7 -4
View File
@@ -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`.
---
+75 -30
View File
@@ -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)
}
+2 -2
View File
@@ -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
+4 -10
View File
@@ -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.
+78 -11
View File
@@ -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) {
+41 -24
View File
@@ -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()
+7 -7
View File
@@ -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
}