package agent import ( "context" "errors" "fmt" "log/slog" "net/url" "time" "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" ) // Applier applies a translated config to the firewall and returns a func that // restores the previous ruleset. Abstracted so the run loop is testable without // touching the kernel. type Applier interface { Apply(ctx context.Context, cfg *config.Config) (revert func() error, err error) } // Agent runs the pull-apply-report loop for one device. type Agent struct { Client *Client Cache Cache Interval time.Duration Applier Applier // Resolver overrides the DNS resolver (tests); nil derives it per-config. Resolver *Resolver // reverted is the last generation rolled back for cutting off the API; it // is not re-applied until a newer generation is published. reverted int64 } // Run loops until ctx is cancelled, applying one cycle per Interval (and once // immediately). A failed cycle is logged and retried on the next tick — the loop // never exits on transient errors. func (a *Agent) Run(ctx context.Context) error { if a.Interval <= 0 { a.Interval = time.Minute } t := time.NewTicker(a.Interval) defer t.Stop() for { if err := a.RunOnce(ctx); err != nil { slog.Error("agent: apply cycle failed", "err", err) } select { case <-ctx.Done(): return ctx.Err() case <-t.C: } } } // RunOnce performs a single pull-apply-report cycle. On a fetch failure it falls // back to the on-disk cache and re-applies it — it never fails closed. func (a *Agent) RunOnce(ctx context.Context) error { rc, raw, err := a.Client.FetchConfig(ctx) if err != nil { slog.Warn("agent: control plane unreachable, using cached config", "err", err) cached, cerr := a.Cache.Read() if cerr != nil { return fmt.Errorf("read cache: %w", cerr) } if cached == nil { return fmt.Errorf("control plane unreachable and no cached config: %w", err) } // Re-apply last known-good; do not report a generation we didn't fetch. return a.applyConfig(ctx, cached, nil) } if a.reverted != 0 && rc.Generation == a.reverted { slog.Warn("agent: skipping reverted generation", "generation", rc.Generation) return nil } return a.applyConfig(ctx, rc, raw) } // applyConfig applies rc. A freshly fetched config (raw != nil) is verified by // reporting status over a new connection through the new ruleset; if the API is // unreachable the previous ruleset is restored. Only verified configs are cached. func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error { resolver := a.Resolver if resolver == nil { resolver = NewResolver(rc.Resolver) } resolver.ExpandDNSSets(ctx, rc) cfg, err := Translate(rc) if err != nil { return fmt.Errorf("translate: %w", err) } revert, err := a.Applier.Apply(ctx, cfg) if err != nil { return fmt.Errorf("apply: %w", err) } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) if raw != nil { a.Client.HTTP.CloseIdleConnections() if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { var uerr *url.Error if errors.As(err, &uerr) { if rerr := revert(); rerr != nil { return fmt.Errorf("API unreachable after applying generation %d (%v); revert failed: %w", rc.Generation, err, rerr) } a.reverted = rc.Generation return fmt.Errorf("reverted generation %d: API unreachable through new ruleset: %w", rc.Generation, err) } slog.Warn("agent: reporting status failed", "err", err) } if err := a.Cache.Write(raw); err != nil { slog.Warn("agent: caching config failed", "err", err) } // Report the FIB so the control plane can scope router enforcement. if fib := CollectFIB(ctx); len(fib) > 0 { if err := a.Client.ReportRoutes(ctx, fib); err != nil { slog.Warn("agent: reporting routes failed", "err", err) } } } return nil } // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} // Apply snapshots the live table, then computes and applies the differential // change set for cfg. The returned func atomically restores the snapshot. func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { engine, err := nftables.NewEngine(cfg) if err != nil { return nil, fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { return nil, fmt.Errorf("computing changes: %w", err) } if changes.Empty() { return func() error { return nil }, nil } snap, err := engine.Snapshot() if err != nil { return nil, fmt.Errorf("snapshotting ruleset: %w", err) } if err := engine.Apply(changes); err != nil { return nil, err } return func() error { return engine.Restore(snap) }, nil }