device: fix some lock ordering violations, add a test for a deadlock we hit

Discovered by a tool + test that will come in a future change.

Updates tailscale/tailscale#19513

Signed-off-by: Brad Fitzpatrick <bradfitz@tailscale.com>
This commit is contained in:
Brad Fitzpatrick 2026-04-24 20:49:35 +00:00 committed by 世界
parent f69b24781e
commit 010dd5c6f2
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
3 changed files with 63 additions and 11 deletions

View file

@ -186,14 +186,22 @@ func (device *Device) upLocked() error {
device.ipcMutex.Lock() device.ipcMutex.Lock()
defer device.ipcMutex.Unlock() defer device.ipcMutex.Unlock()
// Collect peers under RLock and then release before calling into them,
// because SendKeepalive can reach CreateMessageInitiation which acquires
// staticIdentity.RLock; holding peers.RLock across that path would
// invert the staticIdentity < peers hierarchy (see lock-ordering.md).
device.peers.RLock() device.peers.RLock()
peers := make([]*Peer, 0, len(device.peers.keyMap))
for _, peer := range device.peers.keyMap { for _, peer := range device.peers.keyMap {
peers = append(peers, peer)
}
device.peers.RUnlock()
for _, peer := range peers {
peer.Start() peer.Start()
if peer.persistentKeepaliveInterval.Load() > 0 { if peer.persistentKeepaliveInterval.Load() > 0 {
peer.SendKeepalive() peer.SendKeepalive()
} }
} }
device.peers.RUnlock()
return nil return nil
} }
@ -540,16 +548,25 @@ func (device *Device) SendKeepalivesToPeersWithCurrentKeypair() {
return return
} }
// Collect the set of peers to keepalive under peers.RLock, then release
// before invoking SendKeepalive. SendKeepalive can reach
// CreateMessageInitiation which acquires staticIdentity.RLock; holding
// peers.RLock across that path would invert the
// staticIdentity < peers hierarchy (see lock-ordering.md).
var peers []*Peer
device.peers.RLock() device.peers.RLock()
for _, peer := range device.peers.keyMap { for _, peer := range device.peers.keyMap {
peer.keypairs.RLock() peer.keypairs.RLock()
sendKeepalive := peer.keypairs.current != nil && !peer.keypairs.current.created.Add(RejectAfterTime).Before(time.Now()) sendKeepalive := peer.keypairs.current != nil && !peer.keypairs.current.created.Add(RejectAfterTime).Before(time.Now())
peer.keypairs.RUnlock() peer.keypairs.RUnlock()
if sendKeepalive { if sendKeepalive {
peer.SendKeepalive() peers = append(peers, peer)
} }
} }
device.peers.RUnlock() device.peers.RUnlock()
for _, peer := range peers {
peer.SendKeepalive()
}
} }
// closeBindLocked closes the device's net.bind. // closeBindLocked closes the device's net.bind.

27
device/lock-ordering.md Normal file
View file

@ -0,0 +1,27 @@
# Lock Ordering in wireguard-go/device
## Lock hierarchy
Locks must be acquired in the order listed below. A goroutine holding a
lock with a higher number must never attempt to acquire a lock with a
lower number.
```
Level 0 device.state.Mutex
Level 1 device.ipcMutex (sync.RWMutex)
Level 2 device.net.RWMutex
Level 3 device.staticIdentity.RWMutex
Level 4 device.peers.RWMutex
Level 5 peer.state.Mutex
Level 6 peer.handshake.mutex (sync.RWMutex)
Level 7 peer.keypairs.RWMutex
Level 8 device.allowedips.mu (sync.RWMutex)
Level 9 device.indexTable.RWMutex
Level 10 peer.endpoint.Mutex
Level 11 device.cookieChecker.RWMutex
Level 12 peer.cookieGenerator.RWMutex
Level 13 Timer.modifyingLock / Timer.runningLock
```
Not every pair of locks appears in practice; the ordering above is the
transitive closure of the pairs that do.

View file

@ -348,17 +348,22 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation, endpoint
return nil return nil
} }
// Snapshot staticIdentity so we don't hold the RLock across LookupPeer,
// which may call NewPeer (reentrant RLock deadlocks against a pending
// SetPrivateKey writer; see lock-ordering.md).
device.staticIdentity.RLock() device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock() publicKey := device.staticIdentity.publicKey
privateKey := device.staticIdentity.privateKey
device.staticIdentity.RUnlock()
mixHash(&hash, &InitialHash, device.staticIdentity.publicKey[:]) mixHash(&hash, &InitialHash, publicKey[:])
mixHash(&hash, &hash, msg.Ephemeral[:]) mixHash(&hash, &hash, msg.Ephemeral[:])
mixKey(&chainKey, &InitialChainKey, msg.Ephemeral[:]) mixKey(&chainKey, &InitialChainKey, msg.Ephemeral[:])
// decrypt static key // decrypt static key
var peerPK NoisePublicKey var peerPK NoisePublicKey
var key [chacha20poly1305.KeySize]byte var key [chacha20poly1305.KeySize]byte
ss, err := device.staticIdentity.privateKey.sharedSecret(msg.Ephemeral) ss, err := privateKey.sharedSecret(msg.Ephemeral)
if err != nil { if err != nil {
return nil return nil
} }
@ -533,6 +538,14 @@ func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
chainKey [blake2s.Size]byte chainKey [blake2s.Size]byte
) )
// Snapshot the static private key before acquiring handshake.mutex so
// that handshake.mutex is never held while acquiring staticIdentity
// (which would invert the staticIdentity < handshake.mutex hierarchy;
// see lock-ordering.md).
device.staticIdentity.RLock()
privateKey := device.staticIdentity.privateKey
device.staticIdentity.RUnlock()
ok := func() bool { ok := func() bool {
// lock handshake state // lock handshake state
@ -543,11 +556,6 @@ func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
return false return false
} }
// lock private key for reading
device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock()
// finish 3-way DH // finish 3-way DH
mixHash(&hash, &handshake.hash, msg.Ephemeral[:]) mixHash(&hash, &handshake.hash, msg.Ephemeral[:])
@ -560,7 +568,7 @@ func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
mixKey(&chainKey, &chainKey, ss[:]) mixKey(&chainKey, &chainKey, ss[:])
setZero(ss[:]) setZero(ss[:])
ss, err = device.staticIdentity.privateKey.sharedSecret(msg.Ephemeral) ss, err = privateKey.sharedSecret(msg.Ephemeral)
if err != nil { if err != nil {
return false return false
} }