Accept reverted/failed device status reports #26

Merged
benvin merged 2 commits from benvin/status-reverted into main 2026-10-05 21:44:15 +11:00
6 changed files with 199 additions and 11 deletions
@@ -0,0 +1,7 @@
-- Last agent apply report per device. reported_generation keeps its meaning
-- (last applied generation); these columns hold the latest report of any kind.
ALTER TABLE devices
ADD COLUMN status TEXT NOT NULL DEFAULT '' CHECK (status IN ('', 'applied', 'reverted', 'failed')),
ADD COLUMN status_generation BIGINT NOT NULL DEFAULT 0,
ADD COLUMN status_error TEXT NOT NULL DEFAULT '',
ADD COLUMN status_at TIMESTAMPTZ;
+42
View File
@@ -103,6 +103,48 @@ type Device struct {
// ReachablePrefixes is the device's FIB as last reported by its agent
// (server-managed). The compiler uses it to scope router enforcement.
ReachablePrefixes []string `json:"reachable_prefixes,omitempty"`
// AgentStatus is the agent's last apply report (server-managed).
AgentStatus *AgentStatus `json:"agent_status,omitempty"`
}
// ApplyStatus is the outcome an agent reports for a generation.
type ApplyStatus string
const (
StatusApplied ApplyStatus = "applied"
StatusReverted ApplyStatus = "reverted"
StatusFailed ApplyStatus = "failed"
)
// StatusReport is the agent's POST /devices/{name}/status body. An empty
// Status (pre-status agents sending only generation) means applied.
type StatusReport struct {
Status ApplyStatus `json:"status,omitempty"`
Generation int64 `json:"generation"`
Error string `json:"error,omitempty"`
}
// Normalize defaults an empty Status to applied and rejects unknown values.
func (r *StatusReport) Normalize() error {
switch r.Status {
case "":
r.Status = StatusApplied
case StatusApplied, StatusReverted, StatusFailed:
default:
return fmt.Errorf("invalid status %q: want applied, reverted or failed", r.Status)
}
return nil
}
// AgentStatus is a device's stored apply state. AppliedGeneration only advances
// on applied reports; Status/Generation/Error/ReportedAt mirror the last report.
type AgentStatus struct {
AppliedGeneration int64 `json:"applied_generation"`
Status ApplyStatus `json:"status"`
Generation int64 `json:"generation"`
Error string `json:"error,omitempty"`
ReportedAt time.Time `json:"reported_at"`
}
// Binding maps a global zone to one device's local interface(s).
+6 -4
View File
@@ -383,13 +383,15 @@ func (s *Server) handleDeviceRoutes(w http.ResponseWriter, r *http.Request) {
}
func (s *Server) handleDeviceStatus(w http.ResponseWriter, r *http.Request) {
var body struct {
Generation int64 `json:"generation"`
}
var body model.StatusReport
if !decode(w, r, &body) {
return
}
if err := s.store.RecordDeviceStatus(r.Context(), chi.URLParam(r, "name"), body.Generation); err != nil {
if err := body.Normalize(); err != nil {
writeError(w, http.StatusBadRequest, err.Error())
return
}
if err := s.store.RecordDeviceStatus(r.Context(), chi.URLParam(r, "name"), body); err != nil {
if errors.Is(err, store.ErrNotFound) {
writeError(w, http.StatusNotFound, "device not found")
return
+45
View File
@@ -0,0 +1,45 @@
package server
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"git.unkin.net/unkin/tomswallapi/internal/model"
)
func TestStatusReportNormalize(t *testing.T) {
tests := []struct {
in model.ApplyStatus
want model.ApplyStatus
wantErr bool
}{
{in: "", want: model.StatusApplied},
{in: model.StatusApplied, want: model.StatusApplied},
{in: model.StatusReverted, want: model.StatusReverted},
{in: model.StatusFailed, want: model.StatusFailed},
{in: "rolledback", wantErr: true},
}
for _, tt := range tests {
r := model.StatusReport{Status: tt.in, Generation: 3}
err := r.Normalize()
if (err != nil) != tt.wantErr || (!tt.wantErr && r.Status != tt.want) {
t.Errorf("Normalize(%q) = %q, %v", tt.in, r.Status, err)
}
}
}
func TestHandleDeviceStatusRejectsBadPayload(t *testing.T) {
s := &Server{} // nil store: a 400 must return before any DB access
for _, body := range []string{
`{"status":"rolledback","generation":3}`,
`{"generation":3,"bogus":1}`,
} {
rec := httptest.NewRecorder()
s.handleDeviceStatus(rec, httptest.NewRequest(http.MethodPost, "/", strings.NewReader(body)))
if rec.Code != http.StatusBadRequest {
t.Errorf("%s: got %d, want 400", body, rec.Code)
}
}
}
+25 -7
View File
@@ -7,6 +7,7 @@ import (
"encoding/json"
"errors"
"fmt"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
@@ -279,11 +280,20 @@ func (s *Store) UpsertDevice(ctx context.Context, d model.Device) error {
})
}
// RecordDeviceStatus stores the generation an agent reports as applied.
func (s *Store) RecordDeviceStatus(ctx context.Context, name string, generation int64) error {
tag, err := s.pool.Exec(ctx,
`UPDATE devices SET reported_generation = $2, last_seen = now() WHERE name = $1`,
name, generation)
// RecordDeviceStatus stores an agent's apply report. reported_generation only
// advances on applied; the last-status columns only take reports at or above
// the stored status generation, so late or replayed reports never clobber newer state.
func (s *Store) RecordDeviceStatus(ctx context.Context, name string, r model.StatusReport) error {
tag, err := s.pool.Exec(ctx, `
UPDATE devices SET
reported_generation = CASE WHEN $2 = 'applied' THEN GREATEST(reported_generation, $3) ELSE reported_generation END,
status = CASE WHEN $3 >= status_generation THEN $2 ELSE status END,
status_error = CASE WHEN $3 >= status_generation THEN $4 ELSE status_error END,
status_at = CASE WHEN $3 >= status_generation THEN now() ELSE status_at END,
status_generation = GREATEST(status_generation, $3),
last_seen = now()
WHERE name = $1`,
name, string(r.Status), r.Generation, r.Error)
if err != nil {
return err
}
@@ -296,10 +306,14 @@ func (s *Store) RecordDeviceStatus(ctx context.Context, name string, generation
func (s *Store) GetDevice(ctx context.Context, name string) (model.Device, error) {
var d model.Device
var resolver, settings, reachable []byte
var st model.AgentStatus
var statusAt *time.Time
err := s.pool.QueryRow(ctx,
`SELECT name, class, COALESCE(fabric, ''), resolver, settings, reachable_prefixes
`SELECT name, class, COALESCE(fabric, ''), resolver, settings, reachable_prefixes,
reported_generation, status, status_generation, status_error, status_at
FROM devices WHERE name = $1`, name,
).Scan(&d.Name, &d.Class, &d.Fabric, &resolver, &settings, &reachable)
).Scan(&d.Name, &d.Class, &d.Fabric, &resolver, &settings, &reachable,
&st.AppliedGeneration, &st.Status, &st.Generation, &st.Error, &statusAt)
if errors.Is(err, pgx.ErrNoRows) {
return d, ErrNotFound
}
@@ -312,6 +326,10 @@ func (s *Store) GetDevice(ctx context.Context, name string) (model.Device, error
if err := json.Unmarshal(reachable, &d.ReachablePrefixes); err != nil {
return d, err
}
if statusAt != nil {
st.ReportedAt = *statusAt
d.AgentStatus = &st
}
return d, json.Unmarshal(settings, &d.Settings)
}
+74
View File
@@ -149,3 +149,77 @@ func TestNATTierRoundTrip(t *testing.T) {
t.Errorf("expected nat rows to cascade-delete with device, got %d", len(nats))
}
}
func TestRecordDeviceStatus(t *testing.T) {
s := newTestStore(t)
ctx := context.Background()
if err := s.UpsertDevice(ctx, model.Device{Name: "fw1", Class: model.ClassFirewall}); err != nil {
t.Fatalf("upsert device: %v", err)
}
if d, _ := s.GetDevice(ctx, "fw1"); d.AgentStatus != nil {
t.Fatalf("unreported device should have no agent_status: %+v", d.AgentStatus)
}
report := func(r model.StatusReport) model.AgentStatus {
t.Helper()
if err := r.Normalize(); err != nil {
t.Fatalf("normalize: %v", err)
}
if err := s.RecordDeviceStatus(ctx, "fw1", r); err != nil {
t.Fatalf("record %+v: %v", r, err)
}
d, err := s.GetDevice(ctx, "fw1")
if err != nil || d.AgentStatus == nil {
t.Fatalf("get device: %v %+v", err, d.AgentStatus)
}
return *d.AgentStatus
}
// Old agent payload: generation only.
got := report(model.StatusReport{Generation: 5})
if got.AppliedGeneration != 5 || got.Status != model.StatusApplied || got.Generation != 5 || got.ReportedAt.IsZero() {
t.Errorf("old payload: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusApplied, Generation: 6})
if got.AppliedGeneration != 6 || got.Status != model.StatusApplied {
t.Errorf("applied: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusReverted, Generation: 7, Error: "api unreachable"})
if got.AppliedGeneration != 6 || got.Status != model.StatusReverted || got.Generation != 7 || got.Error != "api unreachable" {
t.Errorf("reverted must not advance applied generation: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusFailed, Generation: 8, Error: "nft: syntax error"})
if got.AppliedGeneration != 6 || got.Status != model.StatusFailed || got.Generation != 8 || got.Error != "nft: syntax error" {
t.Errorf("failed must not advance applied generation: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusApplied, Generation: 9})
if got.AppliedGeneration != 9 || got.Status != model.StatusApplied || got.Generation != 9 {
t.Errorf("applied 9: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusApplied, Generation: 4})
if got.AppliedGeneration != 9 || got.Status != model.StatusApplied || got.Generation != 9 {
t.Errorf("older applied must not regress: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusReverted, Generation: 8, Error: "stale"})
if got.AppliedGeneration != 9 || got.Status != model.StatusApplied || got.Generation != 9 || got.Error != "" {
t.Errorf("older reverted must not overwrite newer status: %+v", got)
}
got = report(model.StatusReport{Status: model.StatusReverted, Generation: 9, Error: "lost api"})
if got.AppliedGeneration != 9 || got.Status != model.StatusReverted || got.Generation != 9 || got.Error != "lost api" {
t.Errorf("same-generation reverted after applied must update: %+v", got)
}
if err := s.RecordDeviceStatus(ctx, "fw1", model.StatusReport{Status: "bogus", Generation: 10}); err == nil {
t.Error("db should reject unknown status")
}
if err := s.RecordDeviceStatus(ctx, "nope", model.StatusReport{Status: model.StatusApplied}); err != store.ErrNotFound {
t.Errorf("unknown device: want ErrNotFound, got %v", err)
}
}