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. 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, safe bool) (revert, keep func() error, err error) } // revertDelay is when the revert timer fires if the agent dies mid-apply. const revertDelay = time.Minute // 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 // lastReverted covers a reverted generation whose persistence failed. lastReverted *reverted } // 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 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 { 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 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 { 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() revert, keep, err := a.Applier.Apply(ctx, cfg, raw != nil) if err != nil { err = fmt.Errorf("apply: %w", err) if raw != nil && revert != nil { return a.revertGeneration(ctx, rc.Generation, StatusFailed, err, revert) } if revert != nil { if rerr := revert(); rerr != nil { err = fmt.Errorf("%w; restore: %v; revert timer pending", 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 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 keep != nil && ctx.Err() == nil { return a.revertGeneration(ctx, rc.Generation, StatusReverted, err, revert) } // Shutdown is not a verdict on the generation: keep it. if keep != nil { if kerr := keep(); kerr != nil { slog.Warn("agent: dropping snapshot failed", "err", kerr) } } return err } if keep != nil { if err := keep(); err != nil { return fmt.Errorf("dropping snapshot: %w", 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) } a.lastReverted = nil 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 } // revertGeneration marks generation as reverted before restoring, so 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) } suffix := "" if rerr := revert(); rerr != nil { suffix = fmt.Sprintf("; restore: %v; revert timer pending", rerr) rv.Status = StatusFailed rv.Error += suffix 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 %s: %w%s", generation, rv.Status, cause, suffix) } 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 after a failed apply or for severing the // API; persisted so it is not re-applied until a newer generation arrives. type reverted struct { Generation int64 `json:"generation"` Status string `json:"status,omitempty"` 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 tryapply.WriteFile(a.revertedPath(), b) } // reportReverted reports rv until the control plane accepts it. func (a *Agent) reportReverted(ctx context.Context, rv *reverted) { if rv.Reported { return } status := rv.Status if status == "" { status = StatusReverted } if err := a.Client.ReportStatus(ctx, Status{Status: status, 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, 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) } changes, err := engine.Plan() if err != nil { return nil, nil, fmt.Errorf("computing changes: %w", err) } 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) } // PID 0: 'tomswall confirm' must not signal the agent. if _, err := tryapply.Arm(snap, 0, revertDelay); err != nil { return nil, nil, err } return tryapply.Abort, tryapply.Discard, engine.Apply(changes) }