diff --git a/.changeset/vw-loopback-probe.md b/.changeset/vw-loopback-probe.md new file mode 100644 index 00000000..ab795345 --- /dev/null +++ b/.changeset/vw-loopback-probe.md @@ -0,0 +1,5 @@ +--- +"ftw": patch +--- + +Allow Test connection to use the exact loopback URL already saved for an enabled driver. diff --git a/go/internal/api/api_driver_secrets_test.go b/go/internal/api/api_driver_secrets_test.go index 025d4d82..c5cb0ba2 100644 --- a/go/internal/api/api_driver_secrets_test.go +++ b/go/internal/api/api_driver_secrets_test.go @@ -165,12 +165,12 @@ func TestRejectUnsafeProbeTargetsCoversHTTP(t *testing.T) { HTTP: &config.HTTPCapability{AllowedHosts: []string{"127.0.0.1"}}, }, } - if err := rejectUnsafeProbeTargets(cfg); err == nil { + if err := rejectUnsafeProbeTargets(cfg, ""); err == nil { t.Fatal("HTTP loopback probe should be refused") } cfg.Config["host"] = "192.168.1.10" cfg.Capabilities.HTTP.AllowedHosts = []string{"inverter.local"} - if err := rejectUnsafeProbeTargets(cfg); err != nil { + if err := rejectUnsafeProbeTargets(cfg, ""); err != nil { t.Fatalf("LAN HTTP probe refused: %v", err) } } diff --git a/go/internal/api/api_drivers_debug.go b/go/internal/api/api_drivers_debug.go index aff4dd44..943c6e3f 100644 --- a/go/internal/api/api_drivers_debug.go +++ b/go/internal/api/api_drivers_debug.go @@ -237,7 +237,7 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) { resolved.ResolveDriverPaths(baseDir) cfg = resolved.Drivers[0] - if err := rejectUnsafeProbeTargets(cfg); err != nil { + if err := rejectUnsafeProbeTargets(cfg, s.configuredProbeLoopbackHost(cfg)); err != nil { writeJSON(w, 400, map[string]string{"error": err.Error()}) return } @@ -267,6 +267,7 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) { if displayName == "" { displayName = filepath.Base(cfg.Lua) } + secretOwner := displayName testName := "__test_" + safeProbeName(displayName) + "_" + strconv.FormatInt(time.Now().UnixNano(), 36) cfg.Name = testName if cfg.BatteryCapacityWh <= 0 { @@ -280,6 +281,7 @@ 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, secretOwner) ctx, cancel := context.WithTimeout(r.Context(), 12*time.Second) defer cancel() @@ -321,9 +323,31 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) { } } +func driverSecretStateKey(driverName, key string) string { + return "driver_secret:" + driverName + ":" + key +} + +func (s *Server) wireDriverProbeSecrets(reg *drivers.Registry, probeName, secretOwner string) { + if s.deps.State == nil || strings.TrimSpace(secretOwner) == "" { + return + } + ownerFor := func(driverName string) string { + if driverName == probeName { + return secretOwner + } + return driverName + } + reg.SecretOverride = func(driverName, key string) (string, bool) { + return s.deps.State.LoadConfig(driverSecretStateKey(ownerFor(driverName), key)) + } + reg.SecretPersister = func(driverName, key, value string) error { + return s.deps.State.SaveConfig(driverSecretStateKey(ownerFor(driverName), key), value) + } +} + // rejectUnsafeProbeTargets checks every host a driver test might dial: // MQTT, Modbus, config.host / config.url, and HTTP/WS/TCP allowlists. -func rejectUnsafeProbeTargets(cfg config.Driver) error { +func rejectUnsafeProbeTargets(cfg config.Driver, allowedLoopbackHost string) error { if mq := cfg.EffectiveMQTT(); mq != nil { if err := rejectUnsafeProbeHost(mq.Host); err != nil { return fmt.Errorf("mqtt host: %w", err) @@ -342,7 +366,7 @@ func rejectUnsafeProbeTargets(cfg config.Driver) error { } if u, ok := cfg.Config["url"].(string); ok && strings.TrimSpace(u) != "" { if host := hostFromProbeURL(u); host != "" { - if err := rejectUnsafeProbeHost(host); err != nil { + if err := rejectUnsafeProbeHostOrConfiguredLoopback(host, allowedLoopbackHost); err != nil { return fmt.Errorf("config.url: %w", err) } } @@ -353,7 +377,7 @@ func rejectUnsafeProbeTargets(cfg config.Driver) error { if strings.TrimSpace(h) == "" { continue } - if err := rejectUnsafeProbeHost(hostFromAllowlistEntry(h)); err != nil { + if err := rejectUnsafeProbeHostOrConfiguredLoopback(hostFromAllowlistEntry(h), allowedLoopbackHost); err != nil { return fmt.Errorf("http allowlist: %w", err) } } @@ -381,6 +405,62 @@ func rejectUnsafeProbeTargets(cfg config.Driver) error { return nil } +// configuredProbeLoopbackHost permits a test to reach a loopback URL only +// when that exact URL is already saved for the same enabled driver and Lua +// file. A probe cannot introduce a new loopback destination in its request. +func (s *Server) configuredProbeLoopbackHost(probe config.Driver) string { + if probe.Name == "" || probe.Lua == "" || probe.Config == nil { + return "" + } + current, ok := s.configuredDriver(probe.Name) + if !ok || current.Disabled || current.Lua == "" || + filepath.Clean(current.Lua) != filepath.Clean(probe.Lua) || + !sameProbeHTTPAllowlist(current.Capabilities.HTTP, probe.Capabilities.HTTP) { + return "" + } + savedURL, ok := current.Config["url"].(string) + if !ok || savedURL == "" || probe.Config["url"] != savedURL { + return "" + } + u, err := url.Parse(savedURL) + if err != nil || u.Host == "" || + (!strings.EqualFold(u.Scheme, "http") && !strings.EqualFold(u.Scheme, "https")) { + return "" + } + ip := net.ParseIP(u.Hostname()) + if ip == nil || !ip.IsLoopback() { + return "" + } + return ip.String() +} + +func sameProbeHTTPAllowlist(saved, probe *config.HTTPCapability) bool { + if (saved == nil) != (probe == nil) { + return false + } + if saved == nil { + return true + } + if len(saved.AllowedHosts) != len(probe.AllowedHosts) { + return false + } + for i := range saved.AllowedHosts { + if saved.AllowedHosts[i] != probe.AllowedHosts[i] { + return false + } + } + return true +} + +func rejectUnsafeProbeHostOrConfiguredLoopback(host, allowedLoopbackHost string) error { + host = strings.TrimSpace(strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")) + if ip := net.ParseIP(host); ip != nil && ip.IsLoopback() && + allowedLoopbackHost != "" && ip.String() == allowedLoopbackHost { + return nil + } + return rejectUnsafeProbeHost(host) +} + func hostFromProbeURL(raw string) string { u, err := url.Parse(raw) if err != nil || u.Host == "" { diff --git a/go/internal/api/api_drivers_debug_test.go b/go/internal/api/api_drivers_debug_test.go index 45e17c3a..53bd4371 100644 --- a/go/internal/api/api_drivers_debug_test.go +++ b/go/internal/api/api_drivers_debug_test.go @@ -11,6 +11,7 @@ import ( "testing" "github.com/srcfl/ftw/go/internal/config" + "github.com/srcfl/ftw/go/internal/state" ) // /api/drivers/test handler-level coverage. The probe path runs a real @@ -265,6 +266,158 @@ func TestHandleDriverTestRestoresMaskedSecrets(t *testing.T) { } } +func TestHandleDriverTestUsesOriginalDriverSecretState(t *testing.T) { + dir := t.TempDir() + luaPath := filepath.Join(dir, "oauth_probe.lua") + luaSrc := ` +function driver_init(config) + host.set_poll_interval(50) + if config and config.refresh_token == "fresh-token" then + host.emit_metric("used_fresh_token", 1) + host.persist_secret("refresh_token", "rotated-token") + else + host.emit_metric("used_stale_token", 1) + 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) + } + st, err := state.Open(filepath.Join(dir, "state.db")) + if err != nil { + t.Fatalf("open state: %v", err) + } + t.Cleanup(func() { _ = st.Close() }) + if err := st.SaveConfig(driverSecretStateKey("myuplink", "refresh_token"), "fresh-token"); err != nil { + t.Fatalf("save secret override: %v", err) + } + + live := &config.Config{Drivers: []config.Driver{{ + Name: "myuplink", + Lua: luaPath, + Config: map[string]any{ + "refresh_token": "stale-token", + }, + }}} + srv := New(&Deps{ + Cfg: live, + CfgMu: &sync.RWMutex{}, + ConfigPath: filepath.Join(dir, "config.yaml"), + State: st, + }) + body, _ := json.Marshal(map[string]any{ + "name": "myuplink", + "lua": luaPath, + "config": map[string]any{ + "refresh_token": "stale-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 != 200 { + t.Fatalf("status = %d, want 200 (body=%s)", rr.Code, rr.Body.String()) + } + var resp driverProbeResp + if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal: %v (body=%s)", err, rr.Body.String()) + } + if !resp.OK { + t.Fatalf("probe.ok = false, error=%q (body=%s)", resp.Error, rr.Body.String()) + } + if resp.Health == nil || !strings.HasPrefix(resp.Health.Name, "__test_myuplink_") { + t.Fatalf("probe health name = %+v, want temporary myuplink probe", resp.Health) + } + var usedFresh, usedStale bool + for _, m := range resp.Metrics { + switch m.Name { + case "used_fresh_token": + usedFresh = m.Value == 1 + case "used_stale_token": + usedStale = true + } + } + if !usedFresh || usedStale { + t.Fatalf("metrics = %+v, want fresh token metric only", resp.Metrics) + } + if got, ok := st.LoadConfig(driverSecretStateKey("myuplink", "refresh_token")); !ok || got != "rotated-token" { + t.Fatalf("original driver secret = %q ok=%v, want rotated-token", got, ok) + } + if _, ok := st.LoadConfig(driverSecretStateKey(resp.Health.Name, "refresh_token")); ok { + t.Fatalf("probe wrote secret under temporary name %q", resp.Health.Name) + } +} + +func TestConfiguredProbeLoopbackHostRequiresSameEnabledDriverAndURL(t *testing.T) { + driver := config.Driver{ + Name: "audi-vag", + Lua: "/var/lib/ftw/drivers/vw_merged.lua", + Config: map[string]any{ + "url": "http://127.0.0.1:8787", + }, + } + live := &config.Config{Drivers: []config.Driver{driver}} + srv := New(&Deps{Cfg: live}) + + if got := srv.configuredProbeLoopbackHost(driver); got != "127.0.0.1" { + t.Fatalf("configured loopback host = %q, want 127.0.0.1", got) + } + + changed := driver + changed.Lua = "/var/lib/ftw/drivers/other.lua" + if got := srv.configuredProbeLoopbackHost(changed); got != "" { + t.Errorf("different Lua file was trusted: %q", got) + } + changed = driver + changed.Config = map[string]any{"url": "http://127.0.0.1:8080"} + if got := srv.configuredProbeLoopbackHost(changed); got != "" { + t.Errorf("changed URL was trusted: %q", got) + } + changed = driver + changed.Capabilities.HTTP = &config.HTTPCapability{AllowedHosts: []string{"127.0.0.1"}} + if got := srv.configuredProbeLoopbackHost(changed); got != "" { + t.Errorf("changed HTTP allowlist was trusted: %q", got) + } + live.Drivers[0].Disabled = true + if got := srv.configuredProbeLoopbackHost(driver); got != "" { + t.Errorf("disabled configured driver was trusted: %q", got) + } +} + +func TestRejectUnsafeProbeTargetsAllowsOnlyConfiguredLoopbackException(t *testing.T) { + cfg := config.Driver{ + Config: map[string]any{"url": "http://127.0.0.1:8787"}, + Capabilities: config.Capabilities{ + HTTP: &config.HTTPCapability{AllowedHosts: []string{"127.0.0.1:8787"}}, + }, + } + if err := rejectUnsafeProbeTargets(cfg, ""); err == nil { + t.Fatal("unconfigured loopback URL was accepted") + } + if err := rejectUnsafeProbeTargets(cfg, "127.0.0.1"); err != nil { + t.Fatalf("configured loopback URL was rejected: %v", err) + } + + cfg.Config["url"] = "http://127.0.0.2:8787" + if err := rejectUnsafeProbeTargets(cfg, "127.0.0.1"); err == nil { + t.Fatal("different loopback URL was accepted") + } + cfg.Config["url"] = "http://169.254.169.254:8787" + if err := rejectUnsafeProbeTargets(cfg, "127.0.0.1"); err == nil { + t.Fatal("link-local URL was accepted by the loopback exception") + } + cfg.Config["url"] = "http://127.0.0.1:8787" + cfg.MQTT = &config.MQTTConfig{Host: "127.0.0.1"} + if err := rejectUnsafeProbeTargets(cfg, "127.0.0.1"); err == nil { + t.Fatal("loopback MQTT target was accepted by the HTTP URL exception") + } +} + func TestIsSensitiveKey(t *testing.T) { sensitive := []string{ "password", "Password", "mqtt_password", "passwd", "client_secret",