250 lines
6.7 KiB
Go
250 lines
6.7 KiB
Go
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())
|
|
}
|
|
}
|