Skip to content
Merged
5 changes: 5 additions & 0 deletions .changeset/vw-loopback-probe.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"ftw": patch
---

Let Test connection reuse the saved driver's OAuth tokens and its exact loopback URL.
4 changes: 2 additions & 2 deletions go/internal/api/api_driver_secrets_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
Expand Down
178 changes: 173 additions & 5 deletions go/internal/api/api_drivers_debug.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"context"
"encoding/json"
"fmt"
"log/slog"
"net"
"net/http"
"net/url"
Expand All @@ -21,6 +22,7 @@ import (
"sort"
"strconv"
"strings"
"sync"
"time"

"github.com/srcfl/ftw/go/internal/config"
Expand Down Expand Up @@ -237,7 +239,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
}
Expand Down Expand Up @@ -267,6 +269,7 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) {
if displayName == "" {
displayName = filepath.Base(cfg.Lua)
}
probe := cfg
testName := "__test_" + safeProbeName(displayName) + "_" + strconv.FormatInt(time.Now().UnixNano(), 36)
cfg.Name = testName
if cfg.BatteryCapacityWh <= 0 {
Expand All @@ -280,10 +283,24 @@ 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
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)
}
restartName := probeChangedSharedSecret()
if restartName == "" || s.deps.Registry == nil {
return
}
if err := s.deps.Registry.RestartByName(context.Background(), restartName); err != nil {
slog.Warn("driver probe secret changed but restart failed", "driver", restartName, "err", err)
}
}()
if err := reg.AddProbe(ctx, cfg); err != nil {
writeJSON(w, 200, driverProbeResp{
Name: displayName,
Expand All @@ -293,7 +310,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()
Expand Down Expand Up @@ -321,9 +338,104 @@ func (s *Server) handleDriverTest(w http.ResponseWriter, r *http.Request) {
}
}

func driverSecretStateKey(driverName, key string) string {
return "driver_secret:" + driverName + ":" + key
}

func (s *Server) sameConfiguredProbeDriver(probe config.Driver) (config.Driver, bool) {
if strings.TrimSpace(probe.Name) == "" || strings.TrimSpace(probe.Lua) == "" {
return config.Driver{}, false
}
current, ok := s.configuredDriver(probe.Name)
if !ok || current.Lua == "" ||
filepath.Clean(current.Lua) != filepath.Clean(probe.Lua) {
return config.Driver{}, false
}
return current, true
}

func probeConfigSecret(cfg config.Driver, key string) (string, bool) {
if cfg.Config == nil {
return "", false
}
raw, ok := cfg.Config[key].(string)
if !ok {
return "", false
}
value := strings.TrimSpace(raw)
if value == "" || value == maskedPlaceholder {
return "", false
}
return value, true
}

func probePostedDifferentSecret(probe, live config.Driver, key string) bool {
posted, ok := probeConfigSecret(probe, key)
if !ok {
return false
}
saved, _ := probeConfigSecret(live, key)
return posted != saved
}

func (s *Server) wireDriverProbeSecrets(reg *drivers.Registry, probeName string, probe config.Driver) func() string {
var mu sync.Mutex
changed := false
restartName := ""
changedSharedSecret := func() string {
mu.Lock()
defer mu.Unlock()
if changed {
return restartName
}
return ""
}
if s.deps.State == nil {
return changedSharedSecret
}
live, ok := s.sameConfiguredProbeDriver(probe)
if !ok {
return changedSharedSecret
}
restartName = live.Name
secretOwner := live.SecretOwner()
ownerFor := func(driverName string) string {
if driverName == probeName {
return secretOwner
}
return driverName
}
reg.SecretOverride = func(driverName, key string) (string, bool) {
if driverName == probeName && probePostedDifferentSecret(probe, live, key) {
return "", false
}
return s.deps.State.LoadConfig(driverSecretStateKey(ownerFor(driverName), key))
}
reg.SecretPersister = func(driverName, key, value string) error {
if driverName == probeName && probePostedDifferentSecret(probe, live, key) {
return nil
}
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.SaveDriverSecret(owner, key, 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:
// 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)
Expand All @@ -342,7 +454,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)
}
}
Expand All @@ -353,7 +465,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)
}
}
Expand Down Expand Up @@ -381,6 +493,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 == "" {
Expand Down
Loading
Loading