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()) } }