diff --git a/device/device.go b/device/device.go index 2368383..0cb5735 100644 --- a/device/device.go +++ b/device/device.go @@ -64,6 +64,11 @@ type Device struct { lookupFunc PeerLookupFunc // or nil if unused } + sessionState struct { + sync.Mutex // serializes PeerSessionStateFunc calls and protects peer.sessionState + fn PeerSessionStateFunc + } + rate struct { underLoadUntil atomic.Int64 limiter ratelimiter.Ratelimiter @@ -490,6 +495,34 @@ type PeerLookupFunc func(NoisePublicKey) (_ *NewPeerConfig, ok bool) // See [Device.SetPeerByIPPacketFunc] and [Device.SetPeerLookupFunc]. 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 // when receiving packets for unknown peers. func (device *Device) SetPeerLookupFunc(f PeerLookupFunc) { @@ -507,6 +540,18 @@ func (device *Device) SetPeerByIPPacketFunc(f PeerByIPPacketFunc) { 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() { device.state.Lock() defer device.state.Unlock() diff --git a/device/peer.go b/device/peer.go index b9d45a4..4f90144 100644 --- a/device/peer.go +++ b/device/peer.go @@ -18,14 +18,16 @@ 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 + 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 queuedOutboundPackets atomic.Int32 // packets in staged+outbound queues, for input backpressure @@ -46,6 +48,7 @@ type Peer struct { retransmitHandshake *Timer sendKeepalive *Timer newHandshake *Timer + sessionExpired *Timer zeroKeyMaterial *Timer persistentKeepalive *Timer handshakeAttempts atomic.Uint32 @@ -255,6 +258,10 @@ func (peer *Peer) Start() { func (peer *Peer) ZeroAndFlushAll() { device := peer.device + peer.sessionExpiresNano.Store(0) + if peer.timers.sessionExpired != nil { + peer.timers.sessionExpired.Del() + } // clear key pairs @@ -277,6 +284,7 @@ func (peer *Peer) ZeroAndFlushAll() { handshake.mutex.Unlock() peer.FlushStagedPackets() + peer.noteSessionState(PeerSessionNone) } func (peer *Peer) ExpireCurrentKeypairs() { @@ -296,6 +304,9 @@ func (peer *Peer) ExpireCurrentKeypairs() { next.sendNonce.Store(RejectAfterMessages) } keypairs.Unlock() + + peer.sessionExpiresNano.Store(0) + peer.noteSessionState(PeerSessionExpired) } func (peer *Peer) Stop() { @@ -318,6 +329,52 @@ func (peer *Peer) Stop() { 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) { peer.endpoint.Lock() defer peer.endpoint.Unlock() diff --git a/device/timers.go b/device/timers.go index 97ed451..affa792 100644 --- a/device/timers.go +++ b/device/timers.go @@ -98,6 +98,7 @@ func expiredRetransmitHandshake(peer *Peer) { if peer.timersActive() && !peer.timers.zeroKeyMaterial.IsPending() { peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3) } + peer.noteSessionHandshakeStopped() } else { 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) @@ -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) { if peer.persistentKeepaliveInterval.Load() > 0 { peer.SendKeepalive() @@ -182,6 +192,7 @@ func (peer *Peer) timersHandshakeInitiated() { if peer.timersActive() { 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. */ @@ -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. */ func (peer *Peer) timersSessionDerived() { if peer.timersActive() { + peer.sessionExpiresNano.Store(time.Now().Add(RejectAfterTime).UnixNano()) + peer.timers.sessionExpired.Mod(RejectAfterTime) 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. */ @@ -213,6 +227,7 @@ func (peer *Peer) timersInit() { peer.timers.retransmitHandshake = peer.NewTimer(expiredRetransmitHandshake) peer.timers.sendKeepalive = peer.NewTimer(expiredSendKeepalive) peer.timers.newHandshake = peer.NewTimer(expiredNewHandshake) + peer.timers.sessionExpired = peer.NewTimer(expiredSession) peer.timers.zeroKeyMaterial = peer.NewTimer(expiredZeroKeyMaterial) peer.timers.persistentKeepalive = peer.NewTimer(expiredPersistentKeepalive) } @@ -227,6 +242,7 @@ func (peer *Peer) timersStop() { peer.timers.retransmitHandshake.DelSync() peer.timers.sendKeepalive.DelSync() peer.timers.newHandshake.DelSync() + peer.timers.sessionExpired.DelSync() peer.timers.zeroKeyMaterial.DelSync() peer.timers.persistentKeepalive.DelSync() }