397 lines
12 KiB
Go
397 lines
12 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/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)
|
|
}
|
|
}
|