Compare commits
10 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| afa056b454 | |||
| f593f7625d | |||
| 7c8bd87ec0 | |||
| 9854b0e7b6 | |||
| c7e02c089a | |||
| 9092b463a0 | |||
| 502d06bdda | |||
| 4ad55fc65e | |||
| 7fcbb5fad8 | |||
| dc406c4f56 |
@@ -59,7 +59,7 @@ steps:
|
|||||||
cpu: 2
|
cpu: 2
|
||||||
|
|
||||||
- name: package
|
- name: package
|
||||||
image: git.unkin.net/unkin/almalinux9-rpmbuilder:latest
|
image: artifactapi.k8s.syd1.au.unkin.net/docker-internal/rpmbuilder:0.1.0-alma9
|
||||||
commands:
|
commands:
|
||||||
- ./scripts/build-rpm.sh ${CI_COMMIT_TAG}
|
- ./scripts/build-rpm.sh ${CI_COMMIT_TAG}
|
||||||
depends_on: [build]
|
depends_on: [build]
|
||||||
|
|||||||
@@ -430,6 +430,13 @@ report the generation applied, giving a fleet-wide "converged / N behind" view.
|
|||||||
source/dest disables its rule loudly, never opens it.
|
source/dest disables its rule loudly, never opens it.
|
||||||
- **Adds fail closed, the control plane fails open.** Partial rollout blocks new
|
- **Adds fail closed, the control plane fails open.** Partial rollout blocks new
|
||||||
flows until every hop converges; a dead API leaves the last-good posture running.
|
flows until every hop converges; a dead API leaves the last-good posture running.
|
||||||
|
- **A generation that severs the API is reverted.** The agent applies as a
|
||||||
|
`tomswall try` does (on-disk snapshot, systemd revert timer, shared lock), then
|
||||||
|
reports `applied` over a fresh connection. If that fails at the transport level,
|
||||||
|
or the apply errors, it records the generation in
|
||||||
|
`/var/lib/tomswall/reverted.json` (skipped until a newer one arrives), restores
|
||||||
|
the snapshot and reports `reverted`/`failed`. A failed restore leaves the timer
|
||||||
|
to retry it and is reported `failed`.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -29,7 +29,8 @@ func agentCmd() *cobra.Command {
|
|||||||
Long: `Agent runs the control-plane pull loop: it fetches this device's compiled
|
Long: `Agent runs the control-plane pull loop: it fetches this device's compiled
|
||||||
config from tomswallapi, differentially applies it, and reports the applied
|
config from tomswallapi, differentially applies it, and reports the applied
|
||||||
generation back. It caches the last known-good config and, if the control plane
|
generation back. It caches the last known-good config and, if the control plane
|
||||||
is unreachable, keeps applying that cache — it never fails closed.
|
is unreachable, keeps applying that cache — it never fails closed. A new
|
||||||
|
generation that cuts the agent off from the API is reverted and reported as such.
|
||||||
|
|
||||||
The agent token defaults to the TOMSWALL_AGENT_TOKEN environment variable, and
|
The agent token defaults to the TOMSWALL_AGENT_TOKEN environment variable, and
|
||||||
the device name defaults to the system hostname.`,
|
the device name defaults to the system hostname.`,
|
||||||
|
|||||||
+231
-29
@@ -2,8 +2,13 @@ package agent
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"git.unkin.net/unkin/tomswall/internal/config"
|
"git.unkin.net/unkin/tomswall/internal/config"
|
||||||
@@ -12,11 +17,17 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// Applier applies a translated config to the firewall. Abstracted so the run
|
// Applier applies a translated config to the firewall. Abstracted so the run
|
||||||
// loop is testable without touching the kernel.
|
// 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 {
|
type Applier interface {
|
||||||
Apply(ctx context.Context, cfg *config.Config) 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.
|
||||||
|
const revertDelay = time.Minute
|
||||||
|
|
||||||
// Agent runs the pull-apply-report loop for one device.
|
// Agent runs the pull-apply-report loop for one device.
|
||||||
type Agent struct {
|
type Agent struct {
|
||||||
Client *Client
|
Client *Client
|
||||||
@@ -25,6 +36,9 @@ type Agent struct {
|
|||||||
Applier Applier
|
Applier Applier
|
||||||
// Resolver overrides the DNS resolver (tests); nil derives it per-config.
|
// Resolver overrides the DNS resolver (tests); nil derives it per-config.
|
||||||
Resolver *Resolver
|
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
|
// Run loops until ctx is cancelled, applying one cycle per Interval (and once
|
||||||
@@ -62,16 +76,31 @@ func (a *Agent) RunOnce(ctx context.Context) error {
|
|||||||
return fmt.Errorf("control plane unreachable and no cached config: %w", err)
|
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.
|
// Re-apply last known-good; do not report a generation we didn't fetch.
|
||||||
return a.applyConfig(ctx, cached, false)
|
return a.applyConfig(ctx, cached, nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := a.Cache.Write(raw); err != nil {
|
rv, err := a.readReverted()
|
||||||
slog.Warn("agent: caching config failed", "err", err)
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
return a.applyConfig(ctx, rc, true)
|
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)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error {
|
// 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
|
resolver := a.Resolver
|
||||||
if resolver == nil {
|
if resolver == nil {
|
||||||
resolver = NewResolver(rc.Resolver)
|
resolver = NewResolver(rc.Resolver)
|
||||||
@@ -82,46 +111,219 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("translate: %w", err)
|
return fmt.Errorf("translate: %w", err)
|
||||||
}
|
}
|
||||||
if err := a.Applier.Apply(ctx, cfg); err != nil {
|
|
||||||
return fmt.Errorf("apply: %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))
|
slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules))
|
||||||
|
|
||||||
if report {
|
if err := a.Cache.Write(raw); err != nil {
|
||||||
if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil {
|
slog.Warn("agent: caching config failed", "err", err)
|
||||||
slog.Warn("agent: reporting status failed", "err", err)
|
}
|
||||||
}
|
a.lastReverted = nil
|
||||||
// Report the FIB so the control plane can scope router enforcement.
|
if err := os.Remove(a.revertedPath()); err != nil && !os.IsNotExist(err) {
|
||||||
if fib := CollectFIB(ctx); len(fib) > 0 {
|
slog.Warn("agent: clearing reverted generation failed", "err", err)
|
||||||
if err := a.Client.ReportRoutes(ctx, fib); err != nil {
|
}
|
||||||
slog.Warn("agent: reporting routes 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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// EngineApplier applies via the real nftables differential engine.
|
// revertGeneration marks generation as reverted before restoring, so a failed
|
||||||
type EngineApplier struct{}
|
// 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)
|
||||||
|
}
|
||||||
|
|
||||||
// Apply computes and applies the differential change set for cfg. It refuses
|
var (
|
||||||
// while a 'tomswall try' awaits confirmation.
|
verifyAttempts = 3
|
||||||
func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error {
|
verifyDelay = 2 * time.Second
|
||||||
unlock, err := tryapply.Acquire()
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer unlock()
|
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)
|
engine, err := nftables.NewEngine(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("initializing nftables: %w", err)
|
return nil, nil, fmt.Errorf("initializing nftables: %w", err)
|
||||||
}
|
}
|
||||||
changes, err := engine.Plan()
|
changes, err := engine.Plan()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("computing changes: %w", err)
|
return nil, nil, fmt.Errorf("computing changes: %w", err)
|
||||||
}
|
}
|
||||||
if changes.Empty() {
|
if changes.Empty() {
|
||||||
return nil
|
return nil, nil, nil
|
||||||
}
|
}
|
||||||
return engine.Apply(changes)
|
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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -132,10 +132,10 @@ type fakeApplier struct {
|
|||||||
lastGen int
|
lastGen int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error {
|
func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config, _ bool) (func() error, func() error, error) {
|
||||||
atomic.AddInt32(&f.count, 1)
|
atomic.AddInt32(&f.count, 1)
|
||||||
f.lastGen = len(cfg.Rules)
|
f.lastGen = len(cfg.Rules)
|
||||||
return nil
|
return nil, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
const renderedYAML = `generation: 7
|
const renderedYAML = `generation: 7
|
||||||
|
|||||||
+4
-10
@@ -2,7 +2,8 @@ package agent
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
|
||||||
|
"git.unkin.net/unkin/tomswall/internal/tryapply"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Cache persists the last known-good rendered config to disk so the agent can
|
// Cache persists the last known-good rendered config to disk so the agent can
|
||||||
@@ -11,16 +12,9 @@ type Cache struct {
|
|||||||
Path string
|
Path string
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write atomically stores the raw config bytes.
|
// Write durably stores the raw config bytes.
|
||||||
func (c Cache) Write(raw []byte) error {
|
func (c Cache) Write(raw []byte) error {
|
||||||
if err := os.MkdirAll(filepath.Dir(c.Path), 0o755); err != nil {
|
return tryapply.WriteFile(c.Path, raw)
|
||||||
return err
|
|
||||||
}
|
|
||||||
tmp := c.Path + ".tmp"
|
|
||||||
if err := os.WriteFile(tmp, raw, 0o600); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return os.Rename(tmp, c.Path)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Read returns the cached config, or (nil, nil) when no cache exists yet.
|
// Read returns the cached config, or (nil, nil) when no cache exists yet.
|
||||||
|
|||||||
@@ -26,7 +26,9 @@ func NewClient(baseURL, device, token string) *Client {
|
|||||||
BaseURL: baseURL,
|
BaseURL: baseURL,
|
||||||
Device: device,
|
Device: device,
|
||||||
Token: token,
|
Token: token,
|
||||||
HTTP: &http.Client{Timeout: 30 * time.Second},
|
// No keep-alives: every request, the post-apply check included, opens a
|
||||||
|
// fresh connection that must pass the current ruleset.
|
||||||
|
HTTP: &http.Client{Timeout: 30 * time.Second, Transport: noKeepAlive()},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -94,10 +96,24 @@ func (c *Client) ReportRoutes(ctx context.Context, prefixes []string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ReportStatus tells the control plane which generation this device has applied.
|
// Status values reported to POST /api/v1/devices/{name}/status.
|
||||||
func (c *Client) ReportStatus(ctx context.Context, generation int64) error {
|
const (
|
||||||
|
StatusApplied = "applied"
|
||||||
|
StatusReverted = "reverted"
|
||||||
|
StatusFailed = "failed"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Status is the outcome of applying one generation.
|
||||||
|
type Status struct {
|
||||||
|
Status string `json:"status"`
|
||||||
|
Generation int64 `json:"generation"`
|
||||||
|
Error string `json:"error,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReportStatus tells the control plane the outcome of applying a generation.
|
||||||
|
func (c *Client) ReportStatus(ctx context.Context, st Status) error {
|
||||||
url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device)
|
url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device)
|
||||||
payload, _ := json.Marshal(map[string]int64{"generation": generation})
|
payload, _ := json.Marshal(st)
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -116,3 +132,9 @@ func (c *Client) ReportStatus(ctx context.Context, generation int64) error {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func noKeepAlive() http.RoundTripper {
|
||||||
|
t := http.DefaultTransport.(*http.Transport).Clone()
|
||||||
|
t.DisableKeepAlives = true
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,396 @@
|
|||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.unkin.net/unkin/tomswall/internal/config"
|
||||||
|
"git.unkin.net/unkin/tomswall/internal/nftables"
|
||||||
|
"git.unkin.net/unkin/tomswall/internal/tryapply"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMain(m *testing.M) {
|
||||||
|
dir, err := os.MkdirTemp("", "tomswall-agent-test")
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
tryapply.Dir = dir
|
||||||
|
tryapply.Run = func(name string, args ...string) error {
|
||||||
|
timerCmds = append(timerCmds, name)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
verifyDelay = time.Millisecond
|
||||||
|
verifyTimeout = time.Second
|
||||||
|
code := m.Run()
|
||||||
|
os.RemoveAll(dir)
|
||||||
|
os.Exit(code)
|
||||||
|
}
|
||||||
|
|
||||||
|
// fakeAPI serves a config generation and records status reports; while cut it
|
||||||
|
// drops connections to the status endpoint, as a severing ruleset would.
|
||||||
|
type fakeAPI struct {
|
||||||
|
*httptest.Server
|
||||||
|
gen atomic.Int64
|
||||||
|
cut atomic.Bool
|
||||||
|
code atomic.Int32
|
||||||
|
mu sync.Mutex
|
||||||
|
reports []Status
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeAPI(t *testing.T, gen int64) *fakeAPI {
|
||||||
|
f := &fakeAPI{}
|
||||||
|
f.gen.Store(gen)
|
||||||
|
f.code.Store(http.StatusNoContent)
|
||||||
|
f.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch r.URL.Path {
|
||||||
|
case "/api/v1/devices/fw-a/config":
|
||||||
|
_, _ = w.Write([]byte(strings.Replace(renderedYAML, "generation: 7", "generation: "+itoa(f.gen.Load()), 1)))
|
||||||
|
case "/api/v1/devices/fw-a/status":
|
||||||
|
if f.cut.Load() {
|
||||||
|
conn, _, _ := w.(http.Hijacker).Hijack()
|
||||||
|
conn.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var st Status
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&st)
|
||||||
|
f.mu.Lock()
|
||||||
|
f.reports = append(f.reports, st)
|
||||||
|
f.mu.Unlock()
|
||||||
|
w.WriteHeader(int(f.code.Load()))
|
||||||
|
default:
|
||||||
|
w.WriteHeader(http.StatusNotFound)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
t.Cleanup(f.Close)
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
|
||||||
|
// timerCmds records the systemd commands tryapply runs.
|
||||||
|
var timerCmds []string
|
||||||
|
|
||||||
|
func itoa(n int64) string { b, _ := json.Marshal(n); return string(b) }
|
||||||
|
|
||||||
|
func (f *fakeAPI) last() Status {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
if len(f.reports) == 0 {
|
||||||
|
return Status{}
|
||||||
|
}
|
||||||
|
return f.reports[len(f.reports)-1]
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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, plain, restores int
|
||||||
|
err, restoreErr error
|
||||||
|
onApply func()
|
||||||
|
onRestore func()
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
}
|
||||||
|
tryapply.Restore = func(*nftables.Snapshot) error {
|
||||||
|
f.restores++
|
||||||
|
if f.onRestore != nil {
|
||||||
|
f.onRestore()
|
||||||
|
}
|
||||||
|
return f.restoreErr
|
||||||
|
}
|
||||||
|
f.applies++
|
||||||
|
if f.onApply != nil {
|
||||||
|
f.onApply()
|
||||||
|
}
|
||||||
|
return tryapply.Abort, tryapply.Discard, f.err
|
||||||
|
}
|
||||||
|
|
||||||
|
// pending reports whether a snapshot is still armed and its timer not stopped since.
|
||||||
|
func pending(t *testing.T) bool {
|
||||||
|
t.Helper()
|
||||||
|
_, err := os.Stat(filepath.Join(tryapply.Dir, "try-snapshot.json"))
|
||||||
|
armed := len(timerCmds) > 0 && timerCmds[len(timerCmds)-1] == "systemd-run"
|
||||||
|
if (err == nil) != armed {
|
||||||
|
t.Fatalf("snapshot present=%v but timer armed=%v", err == nil, armed)
|
||||||
|
}
|
||||||
|
return armed
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAgent(t *testing.T, api *fakeAPI, eng *fakeEngine) *Agent {
|
||||||
|
return &Agent{
|
||||||
|
Client: NewClient(api.URL, "fw-a", "tok"),
|
||||||
|
Cache: Cache{Path: filepath.Join(t.TempDir(), "rendered.yaml")},
|
||||||
|
Applier: eng,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cachedGen(t *testing.T, a *Agent) int64 {
|
||||||
|
rc, err := a.Cache.Read()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if rc == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return rc.Generation
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyReachableApplies(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if eng.restores != 0 || api.last() != (Status{Status: StatusApplied, Generation: 7}) || cachedGen(t, a) != 7 || pending(t) {
|
||||||
|
t.Fatalf("restores=%d last=%+v cache=%d", eng.restores, api.last(), cachedGen(t, a))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyUnreachableRevertsAndReports(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, onRestore: func() { api.cut.Store(false) }}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) {
|
||||||
|
t.Fatalf("want errUnreachable, got %v", err)
|
||||||
|
}
|
||||||
|
if eng.restores != 1 || cachedGen(t, a) != 0 || pending(t) {
|
||||||
|
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
|
||||||
|
}
|
||||||
|
if st := api.last(); st.Status != StatusReverted || st.Generation != 7 || st.Error == "" {
|
||||||
|
t.Fatalf("last report %+v", st)
|
||||||
|
}
|
||||||
|
rv, _ := a.readReverted()
|
||||||
|
if rv == nil || rv.Generation != 7 || !rv.Reported {
|
||||||
|
t.Fatalf("persisted %+v", rv)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyRevertReportedOnceReachable(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
_ = a.RunOnce(context.Background())
|
||||||
|
if eng.restores != 1 || api.last().Status != "" {
|
||||||
|
t.Fatalf("restores=%d last=%+v", eng.restores, api.last())
|
||||||
|
}
|
||||||
|
api.cut.Store(false)
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if eng.applies != 1 || api.last() != (Status{Status: StatusReverted, Generation: 7, Error: api.last().Error}) {
|
||||||
|
t.Fatalf("applies=%d last=%+v", eng.applies, api.last())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyShutdownDoesNotRevert(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
eng := &fakeEngine{onApply: func() { api.cut.Store(true); cancel() }}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(ctx); !errors.Is(err, context.Canceled) {
|
||||||
|
t.Fatalf("want context.Canceled, got %v", err)
|
||||||
|
}
|
||||||
|
if rv, _ := a.readReverted(); eng.restores != 0 || rv != nil || cachedGen(t, a) != 0 {
|
||||||
|
t.Fatalf("restores=%d reverted=%+v cache=%d", eng.restores, rv, cachedGen(t, a))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyHTTPErrorDoesNotRevert(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
api.code.Store(http.StatusInternalServerError)
|
||||||
|
eng := &fakeEngine{}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if eng.restores != 0 || cachedGen(t, a) != 7 {
|
||||||
|
t.Fatalf("restores=%d cache=%d", eng.restores, cachedGen(t, a))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyApplyErrorRestoresAndReportsFailed(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{err: errors.New("netlink: boom")}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); err == nil {
|
||||||
|
t.Fatal("want error")
|
||||||
|
}
|
||||||
|
if st := api.last(); eng.restores != 1 || st.Status != StatusFailed || !strings.Contains(st.Error, "boom") || pending(t) {
|
||||||
|
t.Fatalf("restores=%d last=%+v", eng.restores, st)
|
||||||
|
}
|
||||||
|
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported {
|
||||||
|
t.Fatalf("persisted %+v", rv)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyApplyErrorRestoreFailsKeepsTimer(t *testing.T) {
|
||||||
|
t.Cleanup(func() { _ = tryapply.Discard() })
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{err: errors.New("netlink: boom"), restoreErr: errors.New("netlink: stuck")}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); err == nil || !strings.Contains(err.Error(), "restore: restoring snapshot: netlink: stuck") {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
want := "apply: netlink: boom; restore: restoring snapshot: netlink: stuck; revert timer pending"
|
||||||
|
if st := api.last(); st != (Status{Status: StatusFailed, Generation: 7, Error: want}) || !pending(t) {
|
||||||
|
t.Fatalf("last=%+v", st)
|
||||||
|
}
|
||||||
|
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || rv.Status != StatusFailed {
|
||||||
|
t.Fatalf("persisted %+v", rv)
|
||||||
|
}
|
||||||
|
// The next cycle waits for the timer instead of re-applying.
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil || eng.applies != 1 {
|
||||||
|
t.Fatalf("err=%v applies=%d", err, eng.applies)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyUnreachableRestoreFailsKeepsTimer(t *testing.T) {
|
||||||
|
t.Cleanup(func() { _ = tryapply.Discard() })
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{onApply: func() { api.cut.Store(true) }, restoreErr: errors.New("netlink: stuck")}
|
||||||
|
eng.onRestore = func() { api.cut.Store(false) }
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); !errors.Is(err, errUnreachable) || !strings.Contains(err.Error(), "revert timer pending") {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
st := api.last()
|
||||||
|
if st.Status != StatusFailed || st.Generation != 7 || !strings.HasPrefix(st.Error, errUnreachable.Error()) ||
|
||||||
|
!strings.HasSuffix(st.Error, "; restore: restoring snapshot: netlink: stuck; revert timer pending") || !pending(t) {
|
||||||
|
t.Fatalf("last=%+v", st)
|
||||||
|
}
|
||||||
|
if rv, _ := a.readReverted(); rv == nil || rv.Generation != 7 || !rv.Reported || cachedGen(t, a) != 0 {
|
||||||
|
t.Fatalf("persisted %+v cache=%d", rv, cachedGen(t, a))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplyRevertedGenerationSkippedAfterRestart(t *testing.T) {
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.writeReverted(&reverted{Generation: 7, Reported: true}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if eng.applies != 0 {
|
||||||
|
t.Fatalf("reverted generation re-applied")
|
||||||
|
}
|
||||||
|
|
||||||
|
api.gen.Store(8)
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if rv, _ := a.readReverted(); eng.applies != 1 || api.last().Generation != 8 || rv != nil {
|
||||||
|
t.Fatalf("applies=%d last=%+v reverted=%+v", eng.applies, api.last(), rv)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeApplySkipsWhileTryPending(t *testing.T) {
|
||||||
|
marker := filepath.Join(tryapply.Dir, "try-snapshot.json")
|
||||||
|
if err := os.WriteFile(marker, []byte("{}"), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer os.Remove(marker)
|
||||||
|
api := newFakeAPI(t, 7)
|
||||||
|
eng := &fakeEngine{}
|
||||||
|
a := newAgent(t, api, eng)
|
||||||
|
if err := a.RunOnce(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if eng.applies != 0 || api.last().Status != "" {
|
||||||
|
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,15 +25,15 @@ const Unit = "tomswall-try-revert"
|
|||||||
var (
|
var (
|
||||||
// Dir holds the lock and the pending snapshot.
|
// Dir holds the lock and the pending snapshot.
|
||||||
Dir = "/var/lib/tomswall"
|
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 {
|
Run = func(name string, args ...string) error {
|
||||||
if out, err := exec.Command(name, args...).CombinedOutput(); err != nil {
|
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 fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out)))
|
||||||
}
|
}
|
||||||
return nil
|
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 {
|
Restore = func(s *nftables.Snapshot) error {
|
||||||
engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}})
|
engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -94,23 +94,7 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error)
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
f, err := os.CreateTemp(Dir, ".try-snapshot-*")
|
if err := WriteFile(snapshotPath(), b); err != nil {
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
defer os.Remove(f.Name())
|
|
||||||
if _, err := f.Write(b); err != nil {
|
|
||||||
f.Close()
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
if err := f.Sync(); err != nil {
|
|
||||||
f.Close()
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
if err := f.Close(); err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
if err := os.Rename(f.Name(), snapshotPath()); err != nil {
|
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -119,13 +103,46 @@ func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error)
|
|||||||
return "", discardWith(err)
|
return "", discardWith(err)
|
||||||
}
|
}
|
||||||
_ = disarm() // a leftover timer from an earlier try would block the unit name
|
_ = disarm() // a leftover timer from an earlier try would block the unit name
|
||||||
if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit,
|
if err := Run("systemd-run", "--quiet", "--collect", "--unit", Unit,
|
||||||
fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil {
|
fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil {
|
||||||
return "", discardWith(fmt.Errorf("arming revert timer: %w", err))
|
return "", discardWith(fmt.Errorf("arming revert timer: %w", err))
|
||||||
}
|
}
|
||||||
return id, nil
|
return id, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WriteFile durably replaces path with b: temp file, fsync, rename, fsync the directory.
|
||||||
|
func WriteFile(path string, b []byte) error {
|
||||||
|
dir := filepath.Dir(path)
|
||||||
|
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
f, err := os.CreateTemp(dir, "."+filepath.Base(path)+"-*")
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer os.Remove(f.Name())
|
||||||
|
if _, err := f.Write(b); err != nil {
|
||||||
|
f.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := f.Sync(); err != nil {
|
||||||
|
f.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := f.Close(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := os.Rename(f.Name(), path); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
d, err := os.Open(dir)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer d.Close()
|
||||||
|
return d.Sync()
|
||||||
|
}
|
||||||
|
|
||||||
// Discard drops the pending snapshot and timer without restoring. The caller must hold the lock.
|
// Discard drops the pending snapshot and timer without restoring. The caller must hold the lock.
|
||||||
func Discard() error {
|
func Discard() error {
|
||||||
_ = disarm()
|
_ = disarm()
|
||||||
@@ -143,7 +160,7 @@ func discardWith(err error) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func disarm() error {
|
func disarm() error {
|
||||||
return run("systemctl", "stop", Unit+".timer")
|
return Run("systemctl", "stop", Unit+".timer")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e.
|
// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e.
|
||||||
@@ -193,7 +210,7 @@ func Abort() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func restorePending(p *pending) error {
|
func restorePending(p *pending) error {
|
||||||
if err := restore(p.Snapshot); err != nil {
|
if err := Restore(p.Snapshot); err != nil {
|
||||||
return fmt.Errorf("restoring snapshot: %w", err)
|
return fmt.Errorf("restoring snapshot: %w", err)
|
||||||
}
|
}
|
||||||
return Discard()
|
return Discard()
|
||||||
|
|||||||
@@ -15,12 +15,12 @@ func setup(t *testing.T) *[]string {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
Dir = t.TempDir()
|
Dir = t.TempDir()
|
||||||
var cmds []string
|
var cmds []string
|
||||||
orig := run
|
orig := Run
|
||||||
run = func(name string, args ...string) error {
|
Run = func(name string, args ...string) error {
|
||||||
cmds = append(cmds, name+" "+strings.Join(args, " "))
|
cmds = append(cmds, name+" "+strings.Join(args, " "))
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { run = orig })
|
t.Cleanup(func() { Run = orig })
|
||||||
return &cmds
|
return &cmds
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -91,7 +91,7 @@ func TestAcquireRefusesWhilePending(t *testing.T) {
|
|||||||
|
|
||||||
func TestArmFailureDiscardsSnapshot(t *testing.T) {
|
func TestArmFailureDiscardsSnapshot(t *testing.T) {
|
||||||
setup(t)
|
setup(t)
|
||||||
run = func(name string, args ...string) error {
|
Run = func(name string, args ...string) error {
|
||||||
if name == "systemd-run" {
|
if name == "systemd-run" {
|
||||||
return errors.New("no systemd")
|
return errors.New("no systemd")
|
||||||
}
|
}
|
||||||
@@ -145,12 +145,12 @@ func TestConfirmAfterRevertFails(t *testing.T) {
|
|||||||
func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot {
|
func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
var got []*nftables.Snapshot
|
var got []*nftables.Snapshot
|
||||||
orig := restore
|
orig := Restore
|
||||||
restore = func(s *nftables.Snapshot) error {
|
Restore = func(s *nftables.Snapshot) error {
|
||||||
got = append(got, s)
|
got = append(got, s)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { restore = orig })
|
t.Cleanup(func() { Restore = orig })
|
||||||
return &got
|
return &got
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -47,6 +47,17 @@ contents:
|
|||||||
file_info:
|
file_info:
|
||||||
mode: 0640
|
mode: 0640
|
||||||
|
|
||||||
|
# systemd unit + environment file for applying a local config at boot.
|
||||||
|
- src: packaging/tomswall.service
|
||||||
|
dst: /usr/lib/systemd/system/tomswall.service
|
||||||
|
file_info:
|
||||||
|
mode: 0644
|
||||||
|
- src: packaging/tomswall.env
|
||||||
|
dst: /etc/tomswall/tomswall.env
|
||||||
|
type: config|noreplace
|
||||||
|
file_info:
|
||||||
|
mode: 0644
|
||||||
|
|
||||||
# Shell completions (generated by scripts/build-rpm.sh before packaging).
|
# Shell completions (generated by scripts/build-rpm.sh before packaging).
|
||||||
- src: dist/completions/tomswall.bash
|
- src: dist/completions/tomswall.bash
|
||||||
dst: /usr/share/bash-completion/completions/tomswall
|
dst: /usr/share/bash-completion/completions/tomswall
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ Description=tomswall control-plane agent (pull and apply firewall config)
|
|||||||
Documentation=https://git.unkin.net/unkin/tomswall
|
Documentation=https://git.unkin.net/unkin/tomswall
|
||||||
After=network-online.target
|
After=network-online.target
|
||||||
Wants=network-online.target
|
Wants=network-online.target
|
||||||
|
Conflicts=tomswall.service
|
||||||
|
|
||||||
[Service]
|
[Service]
|
||||||
Type=simple
|
Type=simple
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
# Config applied by tomswall.service: a tomswall YAML file or a shorewall directory.
|
||||||
|
TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
|
||||||
|
#TOMSWALL_CONFIG=/etc/shorewall
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
[Unit]
|
||||||
|
Description=tomswall firewall (apply local config at boot)
|
||||||
|
Documentation=https://git.unkin.net/unkin/tomswall
|
||||||
|
DefaultDependencies=no
|
||||||
|
Wants=network-pre.target
|
||||||
|
Before=network-pre.target shutdown.target
|
||||||
|
After=local-fs.target systemd-sysctl.service
|
||||||
|
Conflicts=shutdown.target tomswall-agent.service
|
||||||
|
StartLimitIntervalSec=60
|
||||||
|
StartLimitBurst=5
|
||||||
|
|
||||||
|
[Service]
|
||||||
|
Type=oneshot
|
||||||
|
RemainAfterExit=yes
|
||||||
|
Environment=TOMSWALL_CONFIG=/etc/tomswall/tomswall.yaml
|
||||||
|
EnvironmentFile=-/etc/tomswall/tomswall.env
|
||||||
|
ExecStart=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
|
||||||
|
ExecReload=/usr/sbin/tomswall apply -c ${TOMSWALL_CONFIG}
|
||||||
|
# Fails open: after StartLimitBurst failures within StartLimitIntervalSec, boot continues without the ruleset.
|
||||||
|
Restart=on-failure
|
||||||
|
RestartSec=5
|
||||||
|
# No ExecStop: stopping the unit leaves the ruleset in place (flush would open the firewall).
|
||||||
|
|
||||||
|
[Install]
|
||||||
|
WantedBy=sysinit.target
|
||||||
Reference in New Issue
Block a user