From 969f63298581802073bedfbe6a544c8e82e9fd3d Mon Sep 17 00:00:00 2001 From: Segran Date: Sun, 4 Oct 2026 23:18:44 +0200 Subject: [PATCH] fix(api): restart driver after probe rotates shared secret --- go/internal/api/api_drivers_debug.go | 47 +++++- go/internal/api/api_drivers_debug_test.go | 184 ++++++++++++++++++++++ 2 files changed, 225 insertions(+), 6 deletions(-) diff --git a/go/internal/api/api_drivers_debug.go b/go/internal/api/api_drivers_debug.go index 4fd469f0..a3758341 100644 --- a/go/internal/api/api_drivers_debug.go +++ b/go/internal/api/api_drivers_debug.go @@ -11,6 +11,7 @@ import ( "context" "encoding/json" "fmt" + "log/slog" "net" "net/http" "net/url" @@ -21,6 +22,7 @@ import ( "sort" "strconv" "strings" + "sync" "time" "github.com/srcfl/ftw/go/internal/config" @@ -281,11 +283,23 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) { reg.MQTTFactory = s.deps.DriverMQTTFactory reg.ModbusFactory = s.deps.DriverModbusFactory reg.ARPLookup = s.deps.DriverARPLookup - s.wireDriverProbeSecrets(reg, testName, probe) + probeChangedSharedSecret := s.wireDriverProbeSecrets(reg, testName, probe) ctx, cancel := context.WithTimeout(r.Context(), 12*time.Second) defer cancel() started := time.Now() + probeAdded := false + defer func() { + if probeAdded { + reg.RemoveProbe(cfg.Name) + } + if !probeChangedSharedSecret() || s.deps.Registry == nil { + return + } + if err := s.deps.Registry.RestartByName(context.Background(), probe.Name); err != nil { + slog.Warn("driver probe secret changed but restart failed", "driver", probe.Name, "err", err) + } + }() if err := reg.AddProbe(ctx, cfg); err != nil { writeJSON(w, 200, driverProbeResp{ Name: displayName, @@ -295,7 +309,7 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) { }) return } - defer reg.RemoveProbe(cfg.Name) + probeAdded = true ticker := time.NewTicker(250 * time.Millisecond) defer ticker.Stop() @@ -363,13 +377,20 @@ func probePostedDifferentSecret(probe, live config.Driver, key string) bool { return posted != saved } -func (s *Server) wireDriverProbeSecrets(reg *drivers.Registry, probeName string, probe config.Driver) { +func (s *Server) wireDriverProbeSecrets(reg *drivers.Registry, probeName string, probe config.Driver) func() bool { + var mu sync.Mutex + changed := false + changedSharedSecret := func() bool { + mu.Lock() + defer mu.Unlock() + return changed + } if s.deps.State == nil { - return + return changedSharedSecret } live, ok := s.sameConfiguredProbeDriver(probe) if !ok { - return + return changedSharedSecret } secretOwner := live.Name ownerFor := func(driverName string) string { @@ -388,8 +409,22 @@ func (s *Server) wireDriverProbeSecrets(reg *drivers.Registry, probeName string, if driverName == probeName && probePostedDifferentSecret(probe, live, key) { return nil } - return s.deps.State.SaveConfig(driverSecretStateKey(ownerFor(driverName), key), value) + owner := ownerFor(driverName) + stateKey := driverSecretStateKey(owner, key) + if old, ok := s.deps.State.LoadConfig(stateKey); ok && old == value { + return nil + } + if err := s.deps.State.SaveConfig(stateKey, value); err != nil { + return err + } + if driverName == probeName && owner == secretOwner { + mu.Lock() + changed = true + mu.Unlock() + } + return nil } + return changedSharedSecret } // rejectUnsafeProbeTargets checks every host a driver test might dial: diff --git a/go/internal/api/api_drivers_debug_test.go b/go/internal/api/api_drivers_debug_test.go index 3d9f4ff7..ca53615b 100644 --- a/go/internal/api/api_drivers_debug_test.go +++ b/go/internal/api/api_drivers_debug_test.go @@ -1,6 +1,7 @@ package api import ( + "context" "encoding/json" "net/http" "net/http/httptest" @@ -11,7 +12,9 @@ import ( "testing" "github.com/srcfl/ftw/go/internal/config" + "github.com/srcfl/ftw/go/internal/drivers" "github.com/srcfl/ftw/go/internal/state" + "github.com/srcfl/ftw/go/internal/telemetry" ) // /api/drivers/test handler-level coverage. The probe path runs a real @@ -591,3 +594,184 @@ func TestRedactDumpLog(t *testing.T) { t.Errorf("redactDumpLog dropped benign text: %q", got) } } + +func writeProbeRestartLua(t *testing.T, dir string) string { + t.Helper() + luaPath := filepath.Join(dir, "probe_restart.lua") + luaSrc := ` +function driver_init(config) + host.set_poll_interval(50) + if config and config.rotate_secret then + host.persist_secret("refresh_token", config.persist_value) + end +end +function driver_poll() end +function driver_command() end +function driver_default_mode() end +function driver_cleanup() end +` + if err := os.WriteFile(luaPath, []byte(luaSrc), 0o644); err != nil { + t.Fatalf("write lua: %v", err) + } + return luaPath +} + +func TestHandleDriverTestRestartsRunningDriverAfterRefreshTokenRotation(t *testing.T) { + dir := t.TempDir() + luaPath := writeProbeRestartLua(t, dir) + + st, err := state.Open(filepath.Join(dir, "state.db")) + if err != nil { + t.Fatalf("open state: %v", err) + } + t.Cleanup(func() { _ = st.Close() }) + + const secretKey = "refresh_token" + if err := st.SaveConfig(driverSecretStateKey("myuplink", secretKey), "fresh-token"); err != nil { + t.Fatalf("save secret: %v", err) + } + + tel := telemetry.NewStore() + reg := drivers.NewRegistry(tel) + reg.SecretOverride = func(driverName, key string) (string, bool) { + return st.LoadConfig(driverSecretStateKey(driverName, key)) + } + reg.SecretPersister = func(driverName, key, value string) error { + return st.SaveConfig(driverSecretStateKey(driverName, key), value) + } + t.Cleanup(reg.ShutdownAll) + + liveDriver := config.Driver{ + Name: "myuplink", + Lua: luaPath, + Config: map[string]any{ + "refresh_token": "stale-token", + }, + } + if err := reg.Add(context.Background(), liveDriver); err != nil { + t.Fatalf("add live driver: %v", err) + } + + before, ok := reg.ControlStatus("myuplink") + if !ok { + t.Fatal("live driver missing before probe") + } + + live := &config.Config{Drivers: []config.Driver{liveDriver}} + srv := New(&Deps{ + Cfg: live, + CfgMu: &sync.RWMutex{}, + ConfigPath: filepath.Join(dir, "config.yaml"), + State: st, + Registry: reg, + }) + + body, _ := json.Marshal(map[string]any{ + "name": "myuplink", + "lua": luaPath, + "config": map[string]any{ + "refresh_token": "stale-token", + "rotate_secret": true, + "persist_value": "rotated-token", + }, + }) + req := httptest.NewRequest(http.MethodPost, "/api/drivers/test", strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + srv.Handler().ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200 (body=%s)", rr.Code, rr.Body.String()) + } + + if got, ok := st.LoadConfig(driverSecretStateKey("myuplink", secretKey)); !ok || got != "rotated-token" { + t.Fatalf("persisted secret = %q ok=%v, want rotated-token", got, ok) + } + + after, ok := reg.ControlStatus("myuplink") + if !ok { + t.Fatal("live driver missing after probe") + } + if after.Generation <= before.Generation { + t.Fatalf("generation = %d after probe, want greater than %d after rotated shared secret", + after.Generation, before.Generation) + } +} + +func TestHandleDriverTestDoesNotRestartRunningDriverWhenSecretUnchanged(t *testing.T) { + dir := t.TempDir() + luaPath := writeProbeRestartLua(t, dir) + + st, err := state.Open(filepath.Join(dir, "state.db")) + if err != nil { + t.Fatalf("open state: %v", err) + } + t.Cleanup(func() { _ = st.Close() }) + + const secretKey = "refresh_token" + if err := st.SaveConfig(driverSecretStateKey("myuplink", secretKey), "same-token"); err != nil { + t.Fatalf("save secret: %v", err) + } + + tel := telemetry.NewStore() + reg := drivers.NewRegistry(tel) + reg.SecretOverride = func(driverName, key string) (string, bool) { + return st.LoadConfig(driverSecretStateKey(driverName, key)) + } + reg.SecretPersister = func(driverName, key, value string) error { + return st.SaveConfig(driverSecretStateKey(driverName, key), value) + } + t.Cleanup(reg.ShutdownAll) + + liveDriver := config.Driver{ + Name: "myuplink", + Lua: luaPath, + Config: map[string]any{ + "refresh_token": "stale-token", + }, + } + if err := reg.Add(context.Background(), liveDriver); err != nil { + t.Fatalf("add live driver: %v", err) + } + + before, ok := reg.ControlStatus("myuplink") + if !ok { + t.Fatal("live driver missing before probe") + } + + live := &config.Config{Drivers: []config.Driver{liveDriver}} + srv := New(&Deps{ + Cfg: live, + CfgMu: &sync.RWMutex{}, + ConfigPath: filepath.Join(dir, "config.yaml"), + State: st, + Registry: reg, + }) + + body, _ := json.Marshal(map[string]any{ + "name": "myuplink", + "lua": luaPath, + "config": map[string]any{ + "refresh_token": "stale-token", + "rotate_secret": true, + "persist_value": "same-token", + }, + }) + req := httptest.NewRequest(http.MethodPost, "/api/drivers/test", strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + srv.Handler().ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200 (body=%s)", rr.Code, rr.Body.String()) + } + + after, ok := reg.ControlStatus("myuplink") + if !ok { + t.Fatal("live driver missing after probe") + } + if after.Generation != before.Generation { + t.Fatalf("generation changed from %d to %d even though shared secret was unchanged", + before.Generation, after.Generation) + } +}