Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/oauth-credential-owner.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"ftw": patch
---

Keep rotated OAuth tokens with a stable driver owner so a rename or reused name cannot attach the wrong account.
13 changes: 7 additions & 6 deletions go/cmd/ftw/driver_registry.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package main

import (
"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"
Expand All @@ -11,14 +12,14 @@ func newDriverRegistry(tel *telemetry.Store, st *state.Store) *drivers.Registry
// Install both callbacks before Add can initialize or poll any driver.
// Rotations keep their own KV rows so they do not apply the whole config
// or restart the driver that just refreshed its credential.
driverSecretKey := func(driverName, key string) string {
return "driver_secret:" + driverName + ":" + key
driverSecretKey := func(owner, key string) string {
return config.DriverSecretStateKey(owner, key)
}
reg.SecretPersister = func(driverName, key, value string) error {
return st.SaveConfig(driverSecretKey(driverName, key), value)
reg.SecretPersister = func(owner, key, value string) error {
return st.SaveDriverSecret(owner, key, value)
}
reg.SecretOverride = func(driverName, key string) (string, bool) {
return st.LoadConfig(driverSecretKey(driverName, key))
reg.SecretOverride = func(owner, key string) (string, bool) {
return st.LoadConfig(driverSecretKey(owner, key))
}
return reg
}
123 changes: 122 additions & 1 deletion go/cmd/ftw/driver_registry_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,7 +96,7 @@ function driver_default_mode() end
if got := metric(tel, "persist_ok"); got != 1 {
t.Errorf("secret persistence during %s failed", scenario.phase)
}
if got, ok := st.LoadConfig("driver_secret:oauth-test:refresh_token"); !ok || got != "synthetic-B" {
if got, ok := st.LoadConfig(config.DriverSecretStateKey(cfg.SecretOwner(), "refresh_token")); !ok || got != "synthetic-B" {
t.Errorf("rotated token B was not stored")
}
}()
Expand All @@ -111,3 +111,124 @@ function driver_default_mode() end
})
}
}

func oauthProbeSource(want string) string {
return fmt.Sprintf(`
function driver_init(config)
host.emit_metric("started_with_target", config.refresh_token == %q and 1 or 0)
host.emit_metric("started_with_B", config.refresh_token == "synthetic-B" and 1 or 0)
host.set_poll_interval(10)
end
function driver_poll() return 60000 end
function driver_command() end
function driver_default_mode() end
`, want)
}

func waitDriverMetric(t *testing.T, tel *telemetry.Store, name, key string) float64 {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if value, _, ok := tel.LatestMetric(name, key); ok {
return value
}
time.Sleep(time.Millisecond)
}
t.Fatalf("driver %s did not emit %s", name, key)
return 0
}

func TestDriverRegistryRotatedSecretFollowsOwnerNotName(t *testing.T) {
dir := t.TempDir()
lua := filepath.Join(dir, "oauth.lua")
if err := os.WriteFile(lua, []byte(oauthProbeSource("synthetic-B")), 0600); err != nil {
t.Fatal(err)
}
path, database := filepath.Join(dir, "config.yaml"), filepath.Join(dir, "state.db")
st, err := state.Open(database)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = st.Close() })
if err := st.SaveConfig(config.DriverSecretStateKey("old-name", "refresh_token"), "synthetic-B"); err != nil {
t.Fatal(err)
}

cfg, err := config.Parse([]byte(`
site:
name: Test
fuse:
max_amps: 16
drivers:
- name: ferroamp
lua: drivers/ferroamp.lua
is_site_meter: true
capabilities:
mqtt:
host: 192.168.1.153
api:
port: 8080
`), dir)
if err != nil {
t.Fatal(err)
}
cfg.ConfigDatabase = database
oauth := config.Driver{
Name: "old-name", Lua: lua, Capabilities: config.Capabilities{Standalone: true},
Config: map[string]any{"refresh_token": "synthetic-A"},
}
cfg.Drivers = append(cfg.Drivers, oauth)
if err := config.SaveStored(st, path, cfg); err != nil {
t.Fatal(err)
}
owner := cfg.Drivers[len(cfg.Drivers)-1].CredentialOwner
if owner == "" || owner == "old-name" {
t.Fatalf("credential_owner = %q", owner)
}

start := func(d config.Driver) (*telemetry.Store, func()) {
t.Helper()
tel := telemetry.NewStore()
reg := newDriverRegistry(tel, st)
if err := reg.Add(context.Background(), d); err != nil {
reg.ShutdownAll()
t.Fatal(err)
}
return tel, func() { reg.ShutdownAll() }
}

renamed := cfg.Drivers[len(cfg.Drivers)-1]
renamed.Name = "renamed"
tel, stop := start(renamed)
if got := waitDriverMetric(t, tel, "renamed", "started_with_target"); got != 1 {
t.Error("rename discarded rotated token B")
}
stop()

if err := os.WriteFile(lua, []byte(oauthProbeSource("synthetic-C")), 0600); err != nil {
t.Fatal(err)
}
replaced := config.Driver{
Name: "old-name", Lua: lua, Capabilities: config.Capabilities{Standalone: true},
Config: map[string]any{"refresh_token": "synthetic-C"},
}
replacedCfg := *cfg
replacedCfg.Drivers = append([]config.Driver(nil), cfg.Drivers[:len(cfg.Drivers)-1]...)
replacedCfg.Drivers = append(replacedCfg.Drivers, replaced)
replacedCfg.Revision = cfg.Revision
if err := config.SaveStored(st, path, &replacedCfg); err != nil {
t.Fatal(err)
}
got := replacedCfg.Drivers[len(replacedCfg.Drivers)-1]
if got.CredentialOwner == owner {
t.Fatal("reused name kept the previous credential_owner")
}
tel, stop = start(got)
defer stop()
if got := waitDriverMetric(t, tel, "old-name", "started_with_B"); got != 0 {
t.Error("reused name received leftover rotated token B")
}
if got := waitDriverMetric(t, tel, "old-name", "started_with_target"); got != 1 {
t.Error("reused name did not start with the new account token C")
}
}
6 changes: 6 additions & 0 deletions go/cmd/ftw/forecast_site.go
Original file line number Diff line number Diff line change
Expand Up @@ -293,6 +293,12 @@ func forecastScriptDigests(inputs []config.Driver) map[string]string {
// Keep this encoding compatible with beta.3 so an upgrade can prove that only
// the hash policy changed. Never infer compatibility from a driver name.
func forecastRevisionOf(v forecastSite, weather *config.Weather, hashed []config.Driver, scripts map[string]string) (string, error) {
hashed = append([]config.Driver(nil), hashed...)
for i := range hashed {
// credential_owner is bookkeeping for OAuth rotation. It must not
// reset learned PV/load models when Core mints or migrates it.
hashed[i].CredentialOwner = ""
}
data, err := json.Marshal(struct {
Meter, Timezone string
Options telemetry.ForecastOptions
Expand Down
4 changes: 4 additions & 0 deletions go/cmd/ftw/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -521,6 +521,10 @@ func main() {
} else if len(removed) > 0 {
slog.Info("deleted the stored settings of removed features", "removed", strings.Join(removed, "; "))
}
if err := config.BindCredentialOwners(st, *configPath, cfg); err != nil {
slog.Error("bind driver credential owners", "err", err)
os.Exit(1)
}

if cfg.State != nil && cfg.State.ColdRetentionDays != 0 {
slog.Warn("state.cold_retention_days is retired; fixed EMS history retention applies", "previous_days", cfg.State.ColdRetentionDays)
Expand Down
24 changes: 20 additions & 4 deletions go/internal/api/api_myuplink_oauth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,22 @@ func buildMyUplinkOAuthServer(t *testing.T) (*Server, *config.Config, *state.Sto
return srv, cfg, st
}

func myUplinkRefreshSecret(t *testing.T, cfg *config.Config, st *state.Store) string {
t.Helper()
for _, d := range cfg.Drivers {
if d.Name != "myuplink" {
continue
}
v, ok := st.LoadConfig(config.DriverSecretStateKey(d.SecretOwner(), "refresh_token"))
if !ok {
t.Fatalf("missing refresh secret for owner %q", d.SecretOwner())
}
return v
}
t.Fatal("myuplink driver missing")
return ""
}

func TestNewPKCEPair(t *testing.T) {
v, c, err := newPKCEPair()
if err != nil {
Expand Down Expand Up @@ -255,8 +271,8 @@ func TestMyUplinkOAuthCallbackExchangesAndPersists(t *testing.T) {
t.Errorf("config refresh_token = %v, want RT-from-consent", got)
}
// ...and in the unwatched KV so SecretOverride supersedes any stale value.
if v, ok := st.LoadConfig("driver_secret:myuplink:refresh_token"); !ok || v != "RT-from-consent" {
t.Errorf("KV refresh_token = %q (ok=%v), want RT-from-consent", v, ok)
if v := myUplinkRefreshSecret(t, cfg, st); v != "RT-from-consent" {
t.Errorf("KV refresh_token = %q, want RT-from-consent", v)
}
}

Expand Down Expand Up @@ -304,8 +320,8 @@ func TestMyUplinkOAuthManualExchange(t *testing.T) {
if got := cfg.Drivers[0].Config["refresh_token"]; got != "RT-manual" {
t.Errorf("config refresh_token = %v, want RT-manual", got)
}
if v, ok := st.LoadConfig("driver_secret:myuplink:refresh_token"); !ok || v != "RT-manual" {
t.Errorf("KV refresh_token = %q (ok=%v), want RT-manual", v, ok)
if v := myUplinkRefreshSecret(t, cfg, st); v != "RT-manual" {
t.Errorf("KV refresh_token = %q, want RT-manual", v)
}
}

Expand Down
17 changes: 16 additions & 1 deletion go/internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -903,7 +903,12 @@ func (f Fuse) EffectiveSafetyMarginA() float64 {
// Driver is one driver entry. Each driver is a Lua script loaded by
// the driver host at startup (or on hot-reload via the file watcher).
type Driver struct {
Name string `yaml:"name" json:"name"`
Name string `yaml:"name" json:"name"`
// CredentialOwner is a stable id for rotated OAuth secrets. It is not a
// household setting: Core mints it on first save and the UI posts it
// back unchanged. Secrets are stored as driver_secret:<owner>:<key> so
// a rename keeps the rotation and a reused name cannot inherit it.
CredentialOwner string `yaml:"credential_owner,omitempty" json:"credential_owner,omitempty"`
Lua string `yaml:"lua,omitempty" json:"lua,omitempty"` // path to .lua file
IsSiteMeter bool `yaml:"is_site_meter,omitempty" json:"is_site_meter,omitempty"`
BatteryCapacityWh float64 `yaml:"battery_capacity_wh,omitempty" json:"battery_capacity_wh,omitempty"`
Expand Down Expand Up @@ -1939,6 +1944,7 @@ func (c *Config) Validate() error {
// exists.
siteMeters := 0
names := make(map[string]bool, len(c.Drivers))
owners := make(map[string]string, len(c.Drivers))
for _, d := range c.Drivers {
if d.Name == "" {
return errors.New("driver: name is required")
Expand All @@ -1947,6 +1953,15 @@ func (c *Config) Validate() error {
return fmt.Errorf("driver %q: duplicate name", d.Name)
}
names[d.Name] = true
if owner := strings.TrimSpace(d.CredentialOwner); owner != "" {
if strings.Contains(owner, ":") {
return fmt.Errorf("driver %q: credential_owner must not contain ':'", d.Name)
}
if other, ok := owners[owner]; ok {
return fmt.Errorf("drivers %q and %q share credential_owner", other, d.Name)
}
owners[owner] = d.Name
}

if d.IsSiteMeter {
siteMeters++
Expand Down
Loading
Loading