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:
parent
09268b375c
commit
7c3a736cbe
3 changed files with 126 additions and 8 deletions
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue