diff --git a/hscontrol/state/persist_test.go b/hscontrol/state/persist_test.go index 8e7133b6..3bf5f970 100644 --- a/hscontrol/state/persist_test.go +++ b/hscontrol/state/persist_test.go @@ -5,6 +5,7 @@ import ( "net/netip" "slices" "sync" + "sync/atomic" "testing" "time" @@ -1090,3 +1091,61 @@ func TestNodeWriteChangeWhenPolicyRefreshFails(t *testing.T) { }) } } + +// TestBackfillNodeIPsSurvivesConcurrentPersist guards the backfill against a +// node persist landing between its database write and the NodeStore reload: +// that persist writes the old addresses back and the backfill is lost. +func TestBackfillNodeIPsSurvivesConcurrentPersist(t *testing.T) { + dbPath, s, nodeID := persistTestSetup(t) + require.NoError(t, s.Close()) + + // Dropping the IPv6 prefix gives the backfill work to do. + cfg := persistTestConfig(dbPath) + cfg.PrefixV6 = nil + + s, err := NewState(cfg) + require.NoError(t, err) + t.Cleanup(func() { _ = s.Close() }) + + nv, ok := s.GetNodeByID(nodeID) + require.True(t, ok) + require.True(t, nv.IPv6().Valid(), "precondition: node has an IPv6 address") + + // The backfill reads the nodes table twice: inside its write, then to + // reload NodeStore. Persist the node just before the reload. + var reads atomic.Int32 + + persisted := make(chan struct{}) + + require.NoError(t, s.db.DB.Callback().Query().Before("gorm:query"). + Register("persist_before_reload", func(tx *gorm.DB) { + if tx.Statement.Table != "nodes" || reads.Add(1) != 2 { + return + } + + go func() { + _, _ = s.persistNode(nv) + + close(persisted) + }() + + // Give an unserialised persist time to land before the reload. + select { + case <-persisted: + case <-time.After(200 * time.Millisecond): + } + })) + t.Cleanup(func() { _ = s.db.DB.Callback().Query().Remove("persist_before_reload") }) + + _, _, err = s.BackfillNodeIPs() + require.NoError(t, err) + <-persisted + + nv, ok = s.GetNodeByID(nodeID) + require.True(t, ok) + assert.False(t, nv.IPv6().Valid(), "NodeStore kept the removed IPv6 address") + + dbNode, err := s.db.GetNodeByID(nodeID) + require.NoError(t, err) + assert.Nil(t, dbNode.IPv6, "database kept the removed IPv6 address") +} diff --git a/hscontrol/state/state.go b/hscontrol/state/state.go index d56e3505..dd8e1dce 100644 --- a/hscontrol/state/state.go +++ b/hscontrol/state/state.go @@ -1142,8 +1142,14 @@ func (s *State) RenameNode(nodeID types.NodeID, newName string) (types.NodeView, func (s *State) BackfillNodeIPs() ([]string, []change.Change, error) { genBefore := s.polMan.NodesGeneration() + // Hold persistMu until NodeStore has the new addresses, or a concurrent + // persist writes the old ones back from NodeStore in between. + s.persistMu.Lock() + changes, err := s.db.BackfillNodeIPs(s.ipAlloc) if err != nil { + s.persistMu.Unlock() + return nil, nil, err } @@ -1154,6 +1160,8 @@ func (s *State) BackfillNodeIPs() ([]string, []change.Change, error) { if len(changes) > 0 { nodes, err := s.db.ListNodes() if err != nil { + s.persistMu.Unlock() + return changes, nil, fmt.Errorf("refreshing NodeStore after IP backfill: %w", err) } @@ -1173,6 +1181,8 @@ func (s *State) BackfillNodeIPs() ([]string, []change.Change, error) { s.nodeStore.UpdateNodes(updates) } + s.persistMu.Unlock() + // IPs are policy inputs: without this, clients only learned the new // addresses from whichever unrelated write next refreshed the policy. c, err := s.updatePolicyManagerNodes(genBefore)