device: add peer session state callback

Add a minimal callback API for observing WireGuard peer session state
changes.

Updates tailscale/corp#42874
This commit is contained in:
Brad Fitzpatrick 2026-06-04 21:19:50 +00:00 committed by 世界
parent 09268b375c
commit 7c3a736cbe
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
3 changed files with 126 additions and 8 deletions

View file

@ -64,6 +64,11 @@ type Device struct {
lookupFunc PeerLookupFunc // or nil if unused lookupFunc PeerLookupFunc // or nil if unused
} }
sessionState struct {
sync.Mutex // serializes PeerSessionStateFunc calls and protects peer.sessionState
fn PeerSessionStateFunc
}
rate struct { rate struct {
underLoadUntil atomic.Int64 underLoadUntil atomic.Int64
limiter ratelimiter.Ratelimiter limiter ratelimiter.Ratelimiter
@ -490,6 +495,34 @@ type PeerLookupFunc func(NoisePublicKey) (_ *NewPeerConfig, ok bool)
// See [Device.SetPeerByIPPacketFunc] and [Device.SetPeerLookupFunc]. // See [Device.SetPeerByIPPacketFunc] and [Device.SetPeerLookupFunc].
type PeerByIPPacketFunc func(src, dst netip.Addr, ipPkt []byte) (_ NoisePublicKey, ok bool) type PeerByIPPacketFunc func(src, dst netip.Addr, ipPkt []byte) (_ NoisePublicKey, ok bool)
// PeerSessionState is the current WireGuard session state for a peer.
type PeerSessionState uint8
const (
// PeerSessionNone means there is no handshake in progress and no session key
// material retained for this peer.
PeerSessionNone PeerSessionState = iota
// PeerSessionHandshake means a handshake is in progress for this peer, but
// there is not currently a usable WireGuard session.
PeerSessionHandshake
// PeerSessionEstablished means the peer has a completed WireGuard session
// with usable session key material.
PeerSessionEstablished
// PeerSessionExpired means the peer's session key material is no longer
// considered usable, but final key cleanup or lazy peer removal may not have
// happened yet.
PeerSessionExpired
)
// PeerSessionStateFunc is called when a peer's WireGuard session state changes.
//
// Calls are serialized per Device and delivered in transition order. The
// callback must be cheap and must not call back into Device.
type PeerSessionStateFunc func(peer NoisePublicKey, state PeerSessionState)
// SetPeerLookupFunc sets the function used to look up peers by public key // SetPeerLookupFunc sets the function used to look up peers by public key
// when receiving packets for unknown peers. // when receiving packets for unknown peers.
func (device *Device) SetPeerLookupFunc(f PeerLookupFunc) { func (device *Device) SetPeerLookupFunc(f PeerLookupFunc) {
@ -507,6 +540,18 @@ func (device *Device) SetPeerByIPPacketFunc(f PeerByIPPacketFunc) {
device.allowedips.device = device device.allowedips.device = device
} }
// SetSessionStateFunc sets the function used to observe peer WireGuard session
// state changes.
//
// It does not replay current state. Callers that need a complete view should set
// it before peers are started or lazily created, and maintain any snapshots,
// sequence numbers, and pubsub state outside wireguard-go.
func (device *Device) SetSessionStateFunc(f PeerSessionStateFunc) {
device.sessionState.Lock()
defer device.sessionState.Unlock()
device.sessionState.fn = f
}
func (device *Device) Close() { func (device *Device) Close() {
device.state.Lock() device.state.Lock()
defer device.state.Unlock() defer device.state.Unlock()

View file

@ -26,6 +26,8 @@ type Peer struct {
txBytes atomic.Uint64 // bytes send to peer (endpoint) txBytes atomic.Uint64 // bytes send to peer (endpoint)
rxBytes atomic.Uint64 // bytes received from peer rxBytes atomic.Uint64 // bytes received from peer
lastHandshakeNano atomic.Int64 // nano seconds since epoch lastHandshakeNano atomic.Int64 // nano seconds since epoch
sessionExpiresNano atomic.Int64 // nano seconds since epoch
sessionState PeerSessionState // guarded by device.sessionState.Mutex
queuedOutboundPackets atomic.Int32 // packets in staged+outbound queues, for input backpressure queuedOutboundPackets atomic.Int32 // packets in staged+outbound queues, for input backpressure
@ -46,6 +48,7 @@ type Peer struct {
retransmitHandshake *Timer retransmitHandshake *Timer
sendKeepalive *Timer sendKeepalive *Timer
newHandshake *Timer newHandshake *Timer
sessionExpired *Timer
zeroKeyMaterial *Timer zeroKeyMaterial *Timer
persistentKeepalive *Timer persistentKeepalive *Timer
handshakeAttempts atomic.Uint32 handshakeAttempts atomic.Uint32
@ -255,6 +258,10 @@ func (peer *Peer) Start() {
func (peer *Peer) ZeroAndFlushAll() { func (peer *Peer) ZeroAndFlushAll() {
device := peer.device device := peer.device
peer.sessionExpiresNano.Store(0)
if peer.timers.sessionExpired != nil {
peer.timers.sessionExpired.Del()
}
// clear key pairs // clear key pairs
@ -277,6 +284,7 @@ func (peer *Peer) ZeroAndFlushAll() {
handshake.mutex.Unlock() handshake.mutex.Unlock()
peer.FlushStagedPackets() peer.FlushStagedPackets()
peer.noteSessionState(PeerSessionNone)
} }
func (peer *Peer) ExpireCurrentKeypairs() { func (peer *Peer) ExpireCurrentKeypairs() {
@ -296,6 +304,9 @@ func (peer *Peer) ExpireCurrentKeypairs() {
next.sendNonce.Store(RejectAfterMessages) next.sendNonce.Store(RejectAfterMessages)
} }
keypairs.Unlock() keypairs.Unlock()
peer.sessionExpiresNano.Store(0)
peer.noteSessionState(PeerSessionExpired)
} }
func (peer *Peer) Stop() { func (peer *Peer) Stop() {
@ -318,6 +329,52 @@ func (peer *Peer) Stop() {
peer.ZeroAndFlushAll() peer.ZeroAndFlushAll()
} }
func (peer *Peer) noteSessionState(state PeerSessionState) {
device := peer.device
device.sessionState.Lock()
defer device.sessionState.Unlock()
if peer.sessionState == state {
return
}
peer.sessionState = state
if f := device.sessionState.fn; f != nil {
f(peer.handshake.remoteStatic, state)
}
}
func (peer *Peer) noteSessionHandshakeStarted() {
device := peer.device
device.sessionState.Lock()
defer device.sessionState.Unlock()
switch peer.sessionState {
case PeerSessionEstablished:
return
case PeerSessionHandshake:
return
}
peer.sessionState = PeerSessionHandshake
if f := device.sessionState.fn; f != nil {
f(peer.handshake.remoteStatic, PeerSessionHandshake)
}
}
func (peer *Peer) noteSessionHandshakeStopped() {
state := PeerSessionNone
if peer.hasKeyMaterial() {
state = PeerSessionExpired
}
peer.noteSessionState(state)
}
func (peer *Peer) hasKeyMaterial() bool {
keypairs := &peer.keypairs
keypairs.RLock()
defer keypairs.RUnlock()
return keypairs.previous != nil || keypairs.current != nil || keypairs.next.Load() != nil
}
func (peer *Peer) SetEndpointFromPacket(endpoint conn.Endpoint) { func (peer *Peer) SetEndpointFromPacket(endpoint conn.Endpoint) {
peer.endpoint.Lock() peer.endpoint.Lock()
defer peer.endpoint.Unlock() defer peer.endpoint.Unlock()

View file

@ -98,6 +98,7 @@ func expiredRetransmitHandshake(peer *Peer) {
if peer.timersActive() && !peer.timers.zeroKeyMaterial.IsPending() { if peer.timersActive() && !peer.timers.zeroKeyMaterial.IsPending() {
peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3) peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3)
} }
peer.noteSessionHandshakeStopped()
} else { } else {
peer.timers.handshakeAttempts.Add(1) peer.timers.handshakeAttempts.Add(1)
peer.device.log.Verbosef("%s - Handshake did not complete after %d seconds, retrying (try %d)", peer, int(RekeyTimeout.Seconds()), peer.timers.handshakeAttempts.Load()+1) peer.device.log.Verbosef("%s - Handshake did not complete after %d seconds, retrying (try %d)", peer, int(RekeyTimeout.Seconds()), peer.timers.handshakeAttempts.Load()+1)
@ -139,6 +140,15 @@ func expiredZeroKeyMaterial(peer *Peer) {
} }
} }
func expiredSession(peer *Peer) {
expires := peer.sessionExpiresNano.Load()
if expires == 0 || time.Now().UnixNano() < expires {
return
}
peer.device.log.Verbosef("%s - Session expired after %d seconds", peer, int(RejectAfterTime.Seconds()))
peer.noteSessionState(PeerSessionExpired)
}
func expiredPersistentKeepalive(peer *Peer) { func expiredPersistentKeepalive(peer *Peer) {
if peer.persistentKeepaliveInterval.Load() > 0 { if peer.persistentKeepaliveInterval.Load() > 0 {
peer.SendKeepalive() peer.SendKeepalive()
@ -182,6 +192,7 @@ func (peer *Peer) timersHandshakeInitiated() {
if peer.timersActive() { if peer.timersActive() {
peer.timers.retransmitHandshake.Mod(RekeyTimeout + time.Millisecond*time.Duration(fastrandn(RekeyTimeoutJitterMaxMs))) peer.timers.retransmitHandshake.Mod(RekeyTimeout + time.Millisecond*time.Duration(fastrandn(RekeyTimeoutJitterMaxMs)))
} }
peer.noteSessionHandshakeStarted()
} }
/* Should be called after a handshake response message is received and processed or when getting key confirmation via the first data message. */ /* Should be called after a handshake response message is received and processed or when getting key confirmation via the first data message. */
@ -197,8 +208,11 @@ func (peer *Peer) timersHandshakeComplete() {
/* Should be called after an ephemeral key is created, which is before sending a handshake response or after receiving a handshake response. */ /* Should be called after an ephemeral key is created, which is before sending a handshake response or after receiving a handshake response. */
func (peer *Peer) timersSessionDerived() { func (peer *Peer) timersSessionDerived() {
if peer.timersActive() { if peer.timersActive() {
peer.sessionExpiresNano.Store(time.Now().Add(RejectAfterTime).UnixNano())
peer.timers.sessionExpired.Mod(RejectAfterTime)
peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3) peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3)
} }
peer.noteSessionState(PeerSessionEstablished)
} }
/* Should be called before a packet with authentication -- keepalive, data, or handshake -- is sent, or after one is received. */ /* Should be called before a packet with authentication -- keepalive, data, or handshake -- is sent, or after one is received. */
@ -213,6 +227,7 @@ func (peer *Peer) timersInit() {
peer.timers.retransmitHandshake = peer.NewTimer(expiredRetransmitHandshake) peer.timers.retransmitHandshake = peer.NewTimer(expiredRetransmitHandshake)
peer.timers.sendKeepalive = peer.NewTimer(expiredSendKeepalive) peer.timers.sendKeepalive = peer.NewTimer(expiredSendKeepalive)
peer.timers.newHandshake = peer.NewTimer(expiredNewHandshake) peer.timers.newHandshake = peer.NewTimer(expiredNewHandshake)
peer.timers.sessionExpired = peer.NewTimer(expiredSession)
peer.timers.zeroKeyMaterial = peer.NewTimer(expiredZeroKeyMaterial) peer.timers.zeroKeyMaterial = peer.NewTimer(expiredZeroKeyMaterial)
peer.timers.persistentKeepalive = peer.NewTimer(expiredPersistentKeepalive) peer.timers.persistentKeepalive = peer.NewTimer(expiredPersistentKeepalive)
} }
@ -227,6 +242,7 @@ func (peer *Peer) timersStop() {
peer.timers.retransmitHandshake.DelSync() peer.timers.retransmitHandshake.DelSync()
peer.timers.sendKeepalive.DelSync() peer.timers.sendKeepalive.DelSync()
peer.timers.newHandshake.DelSync() peer.timers.newHandshake.DelSync()
peer.timers.sessionExpired.DelSync()
peer.timers.zeroKeyMaterial.DelSync() peer.timers.zeroKeyMaterial.DelSync()
peer.timers.persistentKeepalive.DelSync() peer.timers.persistentKeepalive.DelSync()
} }