package agent import ( "context" "encoding/json" "errors" "fmt" "log/slog" "net/url" "os" "path/filepath" "time" "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" "git.unkin.net/unkin/tomswall/internal/tryapply" ) // Applier applies a translated config to the firewall. Abstracted so the run // loop is testable without touching the kernel. restore returns the firewall to // its state before the call; it is nil when nothing changed. type Applier interface { Apply(ctx context.Context, cfg *config.Config) (restore 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 } // 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) } rv, err := a.readReverted() if err != nil { return err } if rv != nil { a.reportReverted(ctx, rv) if rc.Generation <= rv.Generation { slog.Info("agent: generation was reverted, waiting for a newer one", "generation", rc.Generation) return nil } } 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. 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) } unlock, err := tryapply.Acquire() if errors.Is(err, tryapply.ErrPending) { slog.Warn("agent: a 'tomswall try' is pending, skipping cycle") return nil } if err != nil { return err } defer unlock() restore, err := a.Applier.Apply(ctx, cfg) if err != nil { if restore != nil { if rerr := restore(); rerr != nil { err = fmt.Errorf("%w; restoring previous ruleset: %v", err, rerr) } } if raw != nil { if rerr := a.Client.ReportStatus(ctx, Status{Status: StatusFailed, Generation: rc.Generation, Error: err.Error()}); rerr != nil { slog.Warn("agent: reporting status failed", "err", rerr) } } return fmt.Errorf("apply: %w", err) } if raw == nil { slog.Info("agent: applied cached config", "generation", rc.Generation, "rules", len(cfg.Rules)) return nil } if err := a.confirm(ctx, rc.Generation); err != nil { if ctx.Err() != nil || restore == nil { return err } if rerr := restore(); rerr != nil { return fmt.Errorf("%w; restoring previous ruleset: %v", err, rerr) } rv := &reverted{Generation: rc.Generation, Error: err.Error()} if werr := a.writeReverted(rv); werr != nil { slog.Error("agent: persisting reverted generation failed", "err", werr) } a.reportReverted(ctx, rv) return fmt.Errorf("generation %d reverted: %w", rc.Generation, err) } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) if err := a.Cache.Write(raw); err != nil { slog.Warn("agent: caching config failed", "err", err) } if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) { slog.Warn("agent: clearing reverted generation 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 } var ( verifyAttempts = 3 verifyDelay = 2 * time.Second verifyTimeout = 5 * time.Second ) // errUnreachable means the API could not be reached through the new ruleset. var errUnreachable = errors.New("control plane unreachable after apply") // confirm reports generation as applied over a fresh connection, which proves // the API is reachable through the new ruleset. Any HTTP response counts as // reachable; only repeated transport failures return errUnreachable. func (a *Agent) confirm(ctx context.Context, generation int64) error { var err error for i := 0; i < verifyAttempts; i++ { if i > 0 { select { case <-ctx.Done(): return ctx.Err() case <-time.After(verifyDelay): } } actx, cancel := context.WithTimeout(ctx, verifyTimeout) err = a.Client.ReportStatus(actx, Status{Status: StatusApplied, Generation: generation}) cancel() if ctx.Err() != nil { return ctx.Err() } var uerr *url.Error if !errors.As(err, &uerr) { if err != nil { slog.Warn("agent: reporting status failed", "err", err) } return nil } } return fmt.Errorf("%w: %v", errUnreachable, err) } // reverted is a generation rolled back for severing the API; persisted so it // is not re-applied after a restart until a newer generation arrives. type reverted struct { Generation int64 `json:"generation"` Error string `json:"error,omitempty"` Reported bool `json:"reported"` } func (a *Agent) revertedPath() string { return filepath.Join(filepath.Dir(a.Cache.Path), "reverted.json") } func (a *Agent) readReverted() (*reverted, error) { b, err := os.ReadFile(a.revertedPath()) if os.IsNotExist(err) { return nil, nil } if err != nil { return nil, err } var rv reverted if err := json.Unmarshal(b, &rv); err != nil { return nil, fmt.Errorf("parsing %s: %w", a.revertedPath(), err) } return &rv, nil } func (a *Agent) writeReverted(rv *reverted) error { b, err := json.Marshal(rv) if err != nil { return err } return Cache{Path: a.revertedPath()}.Write(b) } // reportReverted reports rv until the control plane accepts it. func (a *Agent) reportReverted(ctx context.Context, rv *reverted) { if rv.Reported { return } if err := a.Client.ReportStatus(ctx, Status{Status: StatusReverted, Generation: rv.Generation, Error: rv.Error}); err != nil { slog.Warn("agent: reporting reverted generation failed, retrying next cycle", "generation", rv.Generation, "err", err) return } rv.Reported = true if err := a.writeReverted(rv); err != nil { slog.Warn("agent: persisting reverted generation failed", "err", err) } } // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} // Apply computes and applies the differential change set for cfg, snapshotting // the live table first so the caller can restore it. The caller holds the try lock. 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 nil, nil } snap, err := engine.Snapshot() if err != nil { return nil, fmt.Errorf("snapshotting ruleset: %w", err) } restore := func() error { return engine.Restore(snap) } return restore, engine.Apply(changes) }