Apply the cached config without safe-apply and keep reverted generations in memory
This commit is contained in:
+26
-16
@@ -17,11 +17,12 @@ import (
|
||||
)
|
||||
|
||||
// Applier applies a translated config to the firewall. Abstracted so the run
|
||||
// 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.
|
||||
// loop is testable without touching the kernel. With safe, 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 or safe is false.
|
||||
type Applier interface {
|
||||
Apply(ctx context.Context, cfg *config.Config) (revert, keep func() error, err error)
|
||||
Apply(ctx context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error)
|
||||
}
|
||||
|
||||
// revertDelay is when the revert timer fires if the agent dies mid-apply.
|
||||
@@ -35,6 +36,9 @@ type Agent struct {
|
||||
Applier Applier
|
||||
// Resolver overrides the DNS resolver (tests); nil derives it per-config.
|
||||
Resolver *Resolver
|
||||
|
||||
// lastReverted covers a reverted generation whose persistence failed.
|
||||
lastReverted *reverted
|
||||
}
|
||||
|
||||
// Run loops until ctx is cancelled, applying one cycle per Interval (and once
|
||||
@@ -79,6 +83,9 @@ func (a *Agent) RunOnce(ctx context.Context) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if a.lastReverted != nil && (rv == nil || a.lastReverted.Generation > rv.Generation) {
|
||||
rv = a.lastReverted
|
||||
}
|
||||
if rv != nil {
|
||||
a.reportReverted(ctx, rv)
|
||||
if rc.Generation <= rv.Generation {
|
||||
@@ -89,8 +96,10 @@ func (a *Agent) RunOnce(ctx context.Context) error {
|
||||
return a.applyConfig(ctx, rc, raw)
|
||||
}
|
||||
|
||||
// applyConfig applies rc. A fetched config (raw != nil) is verified by reaching
|
||||
// the API through the new ruleset and reverted if that fails; only then is it cached.
|
||||
// applyConfig applies rc. A fetched config (raw != nil) is applied as a pending
|
||||
// try, verified by reaching the API through the new ruleset and reverted if that
|
||||
// fails; only then is it cached. The cached config is the last verified-good one,
|
||||
// so it is applied plainly: there is nothing to verify it against.
|
||||
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error {
|
||||
resolver := a.Resolver
|
||||
if resolver == nil {
|
||||
@@ -113,7 +122,7 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte)
|
||||
}
|
||||
defer unlock()
|
||||
|
||||
revert, keep, err := a.Applier.Apply(ctx, cfg)
|
||||
revert, keep, err := a.Applier.Apply(ctx, cfg, raw != nil)
|
||||
if err != nil {
|
||||
err = fmt.Errorf("apply: %w", err)
|
||||
if raw != nil && revert != nil {
|
||||
@@ -132,11 +141,6 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte)
|
||||
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
|
||||
}
|
||||
@@ -163,6 +167,7 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte)
|
||||
if err := a.Cache.Write(raw); err != nil {
|
||||
slog.Warn("agent: caching config failed", "err", err)
|
||||
}
|
||||
a.lastReverted = nil
|
||||
if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) {
|
||||
slog.Warn("agent: clearing reverted generation failed", "err", err)
|
||||
}
|
||||
@@ -176,10 +181,12 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte)
|
||||
}
|
||||
|
||||
// revertGeneration marks generation as reverted before restoring, so a failed
|
||||
// restore can never lead to re-applying it, then reports status. A failed
|
||||
// restore can never lead to re-applying it, then reports status. It is also kept
|
||||
// in memory in case persisting fails. 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()}
|
||||
a.lastReverted = rv
|
||||
if err := a.writeReverted(rv); err != nil {
|
||||
slog.Error("agent: persisting reverted generation failed", "err", err)
|
||||
}
|
||||
@@ -293,9 +300,9 @@ 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 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) {
|
||||
// Apply computes and applies the differential change set for cfg, with safe
|
||||
// under a pending try as 'tomswall try' does. The caller holds the try lock.
|
||||
func (EngineApplier) Apply(_ context.Context, cfg *config.Config, safe bool) (revert, keep func() error, err error) {
|
||||
engine, err := nftables.NewEngine(cfg)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("initializing nftables: %w", err)
|
||||
@@ -307,6 +314,9 @@ func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (revert, keep
|
||||
if changes.Empty() {
|
||||
return nil, nil, nil
|
||||
}
|
||||
if !safe {
|
||||
return nil, nil, engine.Apply(changes)
|
||||
}
|
||||
snap, err := engine.Snapshot()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("snapshotting ruleset: %w", err)
|
||||
|
||||
@@ -132,7 +132,7 @@ type fakeApplier struct {
|
||||
lastGen int
|
||||
}
|
||||
|
||||
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, func() error, error) {
|
||||
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config, _ bool) (func() error, func() error, error) {
|
||||
atomic.AddInt32(&f.count, 1)
|
||||
f.lastGen = len(cfg.Rules)
|
||||
return nil, nil, nil
|
||||
|
||||
@@ -89,17 +89,20 @@ func (f *fakeAPI) last() Status {
|
||||
return f.reports[len(f.reports)-1]
|
||||
}
|
||||
|
||||
// fakeEngine always changes the ruleset under a real tryapply pending try;
|
||||
// onApply simulates its effect and restoreErr fails the restore.
|
||||
// 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, restores int
|
||||
err, restoreErr error
|
||||
onApply func()
|
||||
onRestore func()
|
||||
applies, plain, restores int
|
||||
err, restoreErr error
|
||||
onApply func()
|
||||
onRestore func()
|
||||
}
|
||||
|
||||
func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, func() error, error) {
|
||||
f.applies++
|
||||
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
|
||||
}
|
||||
@@ -110,6 +113,7 @@ func (f *fakeEngine) Apply(context.Context, *config.Config) (func() error, func(
|
||||
}
|
||||
return f.restoreErr
|
||||
}
|
||||
f.applies++
|
||||
if f.onApply != nil {
|
||||
f.onApply()
|
||||
}
|
||||
@@ -314,3 +318,79 @@ func TestSafeApplySkipsWhileTryPending(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,14 +25,14 @@ 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 executes a systemd command. Test hook; production code must not reassign.
|
||||
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 rolls the live table back to a snapshot. Test hook; production code must not reassign.
|
||||
Restore = func(s *nftables.Snapshot) error {
|
||||
engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}})
|
||||
if err != nil {
|
||||
|
||||
Reference in New Issue
Block a user