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/nftables" "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 tryapply.Run = func(name string, args ...string) error { timerCmds = append(timerCmds, name) return nil } 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 } // 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 { 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, when safe under a real tryapply pending // try; onApply simulates its effect and restoreErr fails the restore. type fakeEngine struct { applies, plain, restores int err, restoreErr error onApply func() onRestore func() } 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 } tryapply.Restore = func(*nftables.Snapshot) error { f.restores++ if f.onRestore != nil { f.onRestore() } return f.restoreErr } f.applies++ 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 { 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 || pending(t) { 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 || 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 == "" { 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") || 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) { 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()) } } // 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) } }