diff --git a/device/device.go b/device/device.go index 0cb5735..a826c6d 100644 --- a/device/device.go +++ b/device/device.go @@ -64,10 +64,7 @@ type Device struct { lookupFunc PeerLookupFunc // or nil if unused } - sessionState struct { - sync.Mutex // serializes PeerSessionStateFunc calls and protects peer.sessionState - fn PeerSessionStateFunc - } + peerStateFn atomic.Pointer[PeerSessionStateFunc] // observes peer session state changes, nil if unset rate struct { underLoadUntil atomic.Int64 @@ -519,7 +516,7 @@ const ( // PeerSessionStateFunc is called when a peer's WireGuard session state changes. // -// Calls are serialized per Device and delivered in transition order. The +// Calls are serialized per peer and delivered in that peer's transition order. The // callback must be cheap and must not call back into Device. type PeerSessionStateFunc func(peer NoisePublicKey, state PeerSessionState) @@ -546,10 +543,14 @@ func (device *Device) SetPeerByIPPacketFunc(f PeerByIPPacketFunc) { // 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. +// +// The callback must be concurrent-safe and must not call back into Device. func (device *Device) SetSessionStateFunc(f PeerSessionStateFunc) { - device.sessionState.Lock() - defer device.sessionState.Unlock() - device.sessionState.fn = f + if f == nil { + device.peerStateFn.Store(nil) + return + } + device.peerStateFn.Store(&f) } func (device *Device) Close() { diff --git a/device/peer.go b/device/peer.go index 14ecdc2..9726f90 100644 --- a/device/peer.go +++ b/device/peer.go @@ -18,16 +18,20 @@ import ( ) type Peer struct { - isRunning atomic.Bool - keypairs Keypairs - handshake Handshake - device *Device - stopping sync.WaitGroup // routines pending stop - txBytes atomic.Uint64 // bytes send to peer (endpoint) - rxBytes atomic.Uint64 // bytes received from peer - lastHandshakeNano atomic.Int64 // nano seconds since epoch - sessionExpiresNano atomic.Int64 // nano seconds since epoch - sessionState PeerSessionState // guarded by device.sessionState.Mutex + isRunning atomic.Bool + keypairs Keypairs + handshake Handshake + device *Device + stopping sync.WaitGroup // routines pending stop + txBytes atomic.Uint64 // bytes send to peer (endpoint) + rxBytes atomic.Uint64 // bytes received from peer + lastHandshakeNano atomic.Int64 // nano seconds since epoch + + sessionState struct { + sync.Mutex + current PeerSessionState + sessionExpires time.Time + } queuedOutboundPackets atomic.Int32 // packets in staged+outbound queues, for input backpressure @@ -266,7 +270,6 @@ func (peer *Peer) Start() { func (peer *Peer) ZeroAndFlushAll() { device := peer.device - peer.sessionExpiresNano.Store(0) if peer.timers.sessionExpired != nil { peer.timers.sessionExpired.Del() } @@ -292,7 +295,11 @@ func (peer *Peer) ZeroAndFlushAll() { handshake.mutex.Unlock() peer.FlushStagedPackets() - peer.noteSessionState(PeerSessionNone) + + peer.sessionState.Lock() + peer.sessionState.sessionExpires = time.Time{} + peer.noteSessionStateLocked(PeerSessionNone) + peer.sessionState.Unlock() } func (peer *Peer) ExpireCurrentKeypairs() { @@ -313,8 +320,10 @@ func (peer *Peer) ExpireCurrentKeypairs() { } keypairs.Unlock() - peer.sessionExpiresNano.Store(0) - peer.noteSessionState(PeerSessionExpired) + peer.sessionState.Lock() + peer.sessionState.sessionExpires = time.Time{} + peer.noteSessionStateLocked(PeerSessionExpired) + peer.sessionState.Unlock() } func (peer *Peer) Stop() { @@ -338,42 +347,41 @@ func (peer *Peer) Stop() { } func (peer *Peer) noteSessionState(state PeerSessionState) { - device := peer.device - device.sessionState.Lock() - defer device.sessionState.Unlock() + peer.sessionState.Lock() + defer peer.sessionState.Unlock() + peer.noteSessionStateLocked(state) +} - if peer.sessionState == state { +// noteSessionStateLocked records a session state transition and delivers the +// callback. The caller must hold peer.sessionState.Mutex during the +// state determination and transition. +func (peer *Peer) noteSessionStateLocked(state PeerSessionState) { + if peer.sessionState.current == state { return } - peer.sessionState = state - if f := device.sessionState.fn; f != nil { - f(peer.handshake.remoteStatic, state) + peer.sessionState.current = state + if f := peer.device.peerStateFn.Load(); 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: + peer.sessionState.Lock() + defer peer.sessionState.Unlock() + if peer.sessionState.current == PeerSessionEstablished { return } - peer.sessionState = PeerSessionHandshake - if f := device.sessionState.fn; f != nil { - f(peer.handshake.remoteStatic, PeerSessionHandshake) - } + peer.noteSessionStateLocked(PeerSessionHandshake) } func (peer *Peer) noteSessionHandshakeStopped() { + peer.sessionState.Lock() + defer peer.sessionState.Unlock() state := PeerSessionNone if peer.hasKeyMaterial() { state = PeerSessionExpired } - peer.noteSessionState(state) + peer.noteSessionStateLocked(state) } func (peer *Peer) hasKeyMaterial() bool { diff --git a/device/timers.go b/device/timers.go index affa792..d30f26b 100644 --- a/device/timers.go +++ b/device/timers.go @@ -141,12 +141,13 @@ func expiredZeroKeyMaterial(peer *Peer) { } func expiredSession(peer *Peer) { - expires := peer.sessionExpiresNano.Load() - if expires == 0 || time.Now().UnixNano() < expires { + peer.sessionState.Lock() + defer peer.sessionState.Unlock() + if peer.sessionState.sessionExpires.IsZero() || time.Now().Before(peer.sessionState.sessionExpires) { return } peer.device.log.Verbosef("%s - Session expired after %d seconds", peer, int(RejectAfterTime.Seconds())) - peer.noteSessionState(PeerSessionExpired) + peer.noteSessionStateLocked(PeerSessionExpired) } func expiredPersistentKeepalive(peer *Peer) { @@ -208,11 +209,15 @@ 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. */ func (peer *Peer) timersSessionDerived() { if peer.timersActive() { - peer.sessionExpiresNano.Store(time.Now().Add(RejectAfterTime).UnixNano()) + peer.sessionState.Lock() + peer.sessionState.sessionExpires = time.Now().Add(RejectAfterTime) + peer.noteSessionStateLocked(PeerSessionEstablished) + peer.sessionState.Unlock() peer.timers.sessionExpired.Mod(RejectAfterTime) peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3) + } else { + peer.noteSessionState(PeerSessionEstablished) } - peer.noteSessionState(PeerSessionEstablished) } /* Should be called before a packet with authentication -- keepalive, data, or handshake -- is sent, or after one is received. */