Files
tomswall/internal/agent/agent.go
T
unkin-agent c7e02c089a
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Apply the cached config without safe-apply and keep reverted generations in memory
2026-10-05 13:56:11 +11:00

330 lines
10 KiB
Go

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)
}