diff --git a/hscontrol/state/persist_test.go b/hscontrol/state/persist_test.go index 3bf5f9709..5d503e6b1 100644 --- a/hscontrol/state/persist_test.go +++ b/hscontrol/state/persist_test.go @@ -613,10 +613,13 @@ func TestPreAuthKeyReauthRevertsNodeStoreOnDBFailure(t *testing.T) { origNodeKey := node.NodeKey() // Fail the node row update so the re-registration's database write errors - // after the NodeStore has already been mutated. + // after the NodeStore has already been mutated. A session connects during + // the write; the rollback must not undo it. require.NoError(t, s.db.DB.Callback().Update().Before("gorm:update"). Register("fail_node_update", func(tx *gorm.DB) { if tx.Statement.Table == "nodes" { + s.Connect(node.ID()) + _ = tx.AddError(errInjectedNodeUpdate) } })) @@ -631,6 +634,10 @@ func TestPreAuthKeyReauthRevertsNodeStoreOnDBFailure(t *testing.T) { require.True(t, ok) require.Equal(t, origNodeKey, got.NodeKey(), "NodeStore must revert to the persisted node key when the write fails") + + online, known := got.IsOnline().GetOk() + require.True(t, known) + require.True(t, online, "rollback dropped the session that connected during the write") } // TestConcurrentPreAuthKeyRegistrationSameMachineKey ensures concurrent diff --git a/hscontrol/state/state.go b/hscontrol/state/state.go index dd8e1dce8..90d5b0a3f 100644 --- a/hscontrol/state/state.go +++ b/hscontrol/state/state.go @@ -2908,9 +2908,23 @@ func (s *State) HandleNodeFromPreAuthKey( if err != nil { // The NodeStore was updated before the database write. Roll it back // so it does not advertise a registration the database rejected - // (e.g. a node key that a restart would not reload). + // (e.g. a node key that a restart would not reload). Restore only + // the fields the update above wrote: sessions, endpoints and health + // may have moved since priorNode was taken. LastSeen stays, as the + // node did contact us. if priorNode != nil { - s.nodeStore.PutNode(*priorNode) + s.nodeStore.UpdateNode(priorNode.ID, func(n *types.Node) { + n.NodeKey = priorNode.NodeKey + n.Hostname = priorNode.Hostname + n.Hostinfo = priorNode.Hostinfo + n.RegisterMethod = priorNode.RegisterMethod + n.Tags = priorNode.Tags + n.UserID = priorNode.UserID + n.User = priorNode.User + n.Expiry = priorNode.Expiry + n.AuthKey = priorNode.AuthKey + n.AuthKeyID = priorNode.AuthKeyID + }) } return types.NodeView{}, s.policyChangeSince(genBefore), fmt.Errorf("writing node to database: %w", err)