Merge branch 'main' into feat-post_quantum_ml_kem

# Conflicts:
#	client/internal/peer/conn.go
#	client/internal/peer/handshaker.go
This commit is contained in:
riccardom
2026-10-05 15:33:58 +02:00
456 changed files with 33601 additions and 14646 deletions
+25 -9
View File
@@ -176,9 +176,10 @@ type Conn struct {
// used to store the remote Rosenpass key for Relayed connection in case of connection update from ice
rosenpassRemoteKey []byte
wgProxyICE wgproxy.Proxy
wgProxyRelay wgproxy.Proxy
handshaker *Handshaker
wgProxyICE wgproxy.Proxy
wgProxyRelay wgproxy.Proxy
relayedConnRef *relayClient.Conn
handshaker *Handshaker
guard *guard.Guard
wg sync.WaitGroup
@@ -625,7 +626,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.mu.Lock()
defer conn.mu.Unlock()
if conn.ctx.Err() != nil {
if conn.ctx.Err() != nil || rci.relayedConn.Context().Err() != nil {
if err := rci.relayedConn.Close(); err != nil {
conn.Log.Warnf("failed to close unnecessary relayed connection: %v", err)
}
@@ -640,7 +641,9 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err)
return
}
wgProxy.SetDisconnectListener(conn.onRelayDisconnected)
wgProxy.SetDisconnectListener(func() {
conn.onRelayDisconnected(rci.relayedConn)
})
conn.dumpState.NewLocalProxy()
@@ -648,7 +651,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
if conn.isICEActive() {
conn.Log.Debugf("do not switch to relay because current priority is: %s", conn.currentConnPriority.String())
conn.setRelayedProxy(wgProxy)
conn.setRelayedProxy(wgProxy, rci.relayedConn)
conn.statusRelay.SetConnected()
conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, pqOK, time.Now())
return
@@ -679,15 +682,26 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.rosenpassRemoteKey = rci.rosenpassPubKey
conn.currentConnPriority = conntype.Relay
conn.statusRelay.SetConnected()
conn.setRelayedProxy(wgProxy)
conn.setRelayedProxy(wgProxy, rci.relayedConn)
conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, pqOK, updateTime)
conn.Log.Infof("start to communicate with peer via relay")
conn.doOnConnected(rci.rosenpassPubKey, rci.rosenpassAddr, updateTime)
}
func (conn *Conn) onRelayDisconnected() {
// onRelayDisconnected reports the teardown of a relayed connection. relayedConn
// names the connection the signal belongs to, so a signal that arrives after
// its connection was replaced is ignored instead of tearing down its successor.
// A nil relayedConn means the caller does not track generations and the current
// connection is always torn down.
func (conn *Conn) onRelayDisconnected(relayedConn *relayClient.Conn) {
conn.mu.Lock()
defer conn.mu.Unlock()
if relayedConn != nil && conn.relayedConnRef != relayedConn {
conn.Log.Debugf("ignoring relay disconnect of a superseded connection")
return
}
conn.handleRelayDisconnectedLocked()
}
@@ -711,6 +725,7 @@ func (conn *Conn) handleRelayDisconnectedLocked() {
_ = conn.wgProxyRelay.CloseConn()
conn.wgProxyRelay = nil
}
conn.relayedConnRef = nil
changed := conn.statusRelay.Get() != worker.StatusDisconnected
if changed {
@@ -1020,13 +1035,14 @@ func (conn *Conn) logTraceConnState() {
}
}
func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy) {
func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy, relayedConn *relayClient.Conn) {
if conn.wgProxyRelay != nil {
if err := conn.wgProxyRelay.CloseConn(); err != nil {
conn.Log.Warnf("failed to close deprecated wg proxy conn: %v", err)
}
}
conn.wgProxyRelay = proxy
conn.relayedConnRef = relayedConn
}
// onWGHandshakeSuccess is called when the first WireGuard handshake is detected
+4 -4
View File
@@ -159,7 +159,7 @@ func (h *Handshaker) notifyListeners(remoteOfferAnswer *OfferAnswer) {
}
func (h *Handshaker) handleRemoteOffer(remoteOfferAnswer OfferAnswer) {
h.log.Infof("received offer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials())
h.log.Infof("received offer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t, relay server: %s, relay IP: %s", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials(), remoteOfferAnswer.RelaySrvAddress, remoteOfferAnswer.RelaySrvIP)
h.onSignalReceived(&remoteOfferAnswer)
// If we are the controller running the KEM, a responder's offer is handled by
@@ -180,7 +180,7 @@ func (h *Handshaker) handleRemoteOffer(remoteOfferAnswer OfferAnswer) {
}
func (h *Handshaker) handleRemoteAnswer(remoteOfferAnswer OfferAnswer) {
h.log.Infof("received answer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials())
h.log.Infof("received answer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t, relay server: %s, relay IP: %s", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials(), remoteOfferAnswer.RelaySrvAddress, remoteOfferAnswer.RelaySrvIP)
h.onSignalReceived(&remoteOfferAnswer)
// Feed the KEM answer (derive+store PSK) BEFORE bringing up the connection so the WG
@@ -309,7 +309,7 @@ func (h *Handshaker) sendOffer() error {
if h.config.PQ != nil {
offer.MlkemPayload, offer.MlkemPort = h.config.PQ.OfferPayload(h.config.Key)
}
h.log.Debugf("sending offer with serial: %s", offer.SessionIDString())
h.log.Debugf("sending offer with serial: %s, relay server: %s, relay IP: %s", offer.SessionIDString(), offer.RelaySrvAddress, offer.RelaySrvIP)
return h.signaler.SignalOffer(offer, h.config.Key)
}
@@ -323,7 +323,7 @@ func (h *Handshaker) sendAnswer(remoteOffer *OfferAnswer) error {
}
answer.MlkemPayload, answer.MlkemPort = h.config.PQ.AnswerPayload(h.config.Key, recvOffer)
}
h.log.Debugf("sending answer with serial: %s", answer.SessionIDString())
h.log.Debugf("sending answer with serial: %s, relay server: %s, relay IP: %s", answer.SessionIDString(), answer.RelaySrvAddress, answer.RelaySrvIP)
return h.signaler.SignalAnswer(answer, h.config.Key)
}
+2 -2
View File
@@ -830,8 +830,8 @@ func (d *Status) SetSessionExpiresAt(deadline time.Time) {
// "none" would blank the UI at the exact moment it should say the session
// ended.
func (d *Status) GetSessionExpiresAt() time.Time {
d.mux.Lock()
defer d.mux.Unlock()
d.mux.RLock()
defer d.mux.RUnlock()
return d.sessionExpiresAt
}
+17 -10
View File
@@ -121,11 +121,8 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
}
}
sessionID, err := NewICESessionID()
if err != nil {
w.log.Errorf("failed to create new session ID: %s", err)
}
w.sessionID = sessionID
// Keep the ID already advertised to the remote. Answers do not get a
// reply, so changing it here makes the next offer restart both sides.
w.abandonNegotiation()
}
@@ -205,6 +202,9 @@ func (w *WorkerICE) Close() {
w.muxAgent.Lock()
defer w.muxAgent.Unlock()
if w.agent != nil || w.agentConnecting {
w.renewSessionID()
}
if w.agent != nil {
w.agentDialerCancel()
if err := w.agent.Close(); err != nil {
@@ -366,16 +366,23 @@ func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.C
// Only the owner of the current session may reset its state: a stale dial
// goroutine waking after a newer attempt must not clobber it.
if w.agent == agent {
sessionID, err := NewICESessionID()
if err != nil {
w.log.Errorf("failed to create new session ID: %s", err)
}
w.sessionID = sessionID
w.renewSessionID()
w.abandonNegotiation()
}
return sessionChanged
}
// renewSessionID starts a new local session, so the remote treats our next offer
// or answer as a restart. Caller holds muxAgent.
func (w *WorkerICE) renewSessionID() {
sessionID, err := NewICESessionID()
if err != nil {
w.log.Errorf("failed to create new session ID: %s", err)
return
}
w.sessionID = sessionID
}
// abandonNegotiation drops all recorded ICE session state so the worker treats the
// next offer as a fresh start instead of a duplicate of a dead negotiation. The
// agent and agentConnecting flags must change together: leaving one stale wedges
@@ -0,0 +1,375 @@
package peer
import (
"context"
"fmt"
"net"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
icemaker "github.com/netbirdio/netbird/client/internal/peer/ice"
)
func TestWorkerICE_RemoteRestartPreservesAdvertisedSession(t *testing.T) {
w := newTestWorkerICE(t)
t.Cleanup(w.Close)
w.dialFunc = parkDial
advertised := w.SessionID()
remoteSession := ICESessionID("remote-first")
offer := OfferAnswer{
IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"},
SessionID: &remoteSession,
}
w.OnNewOffer(&offer)
require.True(t, w.InProgress(), "the first remote session must start ICE")
w.muxAgent.Lock()
firstAgent := w.agent
w.muxAgent.Unlock()
// The same callback handles answers. A changed remote ID must not create
// an unannounced local ID that makes the remote restart on our next offer.
secondSession := ICESessionID("remote-restarted")
answer := offer
answer.SessionID = &secondSession
w.OnNewOffer(&answer)
assert.Equal(t, advertised, w.SessionID(), "following a remote restart must keep our advertised ID")
w.muxAgent.Lock()
secondAgent := w.agent
w.muxAgent.Unlock()
assert.NotSame(t, firstAgent, secondAgent, "the changed remote session must still rebuild ICE")
w.OnNewOffer(&answer)
w.muxAgent.Lock()
defer w.muxAgent.Unlock()
assert.Same(t, secondAgent, w.agent, "a repeated answer must keep the replacement agent")
}
func TestWorkerICE_LocalCloseChangesAdvertisedSession(t *testing.T) {
w := newTestWorkerICE(t)
dialStarted := make(chan struct{})
dialDone := make(chan struct{})
w.dialFunc = func(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) {
close(dialStarted)
defer close(dialDone)
<-ctx.Done()
return nil, ctx.Err()
}
session := ICESessionID("remote-session")
w.OnNewOffer(&OfferAnswer{
IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"},
SessionID: &session,
})
<-dialStarted
advertised := w.SessionID()
w.Close()
assert.NotEqual(t, advertised, w.SessionID(), "a local teardown must tell the remote to restart")
closedSession := w.SessionID()
// The abandoned dial goroutine cleans up after Close returned.
<-dialDone
assert.Never(t, func() bool { return w.SessionID() != closedSession }, 200*time.Millisecond, 10*time.Millisecond,
"the late cleanup of a closed negotiation must not restart again")
w.Close()
assert.Equal(t, closedSession, w.SessionID(), "closing an idle worker must not restart again")
}
// parkDial stands in for the ICE dial. It never connects and returns once the
// negotiation is abandoned, so a test decides when a negotiation fails.
func parkDial(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) {
<-ctx.Done()
return nil, ctx.Err()
}
func newTestSessionID(t *testing.T) ICESessionID {
t.Helper()
sid, err := NewICESessionID()
require.NoError(t, err)
return sid
}
// handshakeSide is one end of a simulated signaling exchange.
type handshakeSide interface {
// message builds the offer or answer the side would send now.
message() OfferAnswer
// receive hands a remote offer or answer to the side's ICE logic.
receive(msg OfferAnswer)
// teardowns counts negotiations the side tore down to follow a remote restart.
teardowns() int
// failAgent ends the side's current negotiation as an ICE failure does.
failAgent()
}
// workerSide drives a real WorkerICE.
type workerSide struct {
t *testing.T
w *WorkerICE
replaced int
}
func newWorkerSide(t *testing.T) *workerSide {
t.Helper()
w := newTestWorkerICE(t)
w.dialFunc = parkDial
t.Cleanup(w.Close)
return &workerSide{t: t, w: w}
}
func (s *workerSide) message() OfferAnswer {
sid := s.w.SessionID()
ufrag, pwd := s.w.GetLocalUserCredentials()
return OfferAnswer{IceCredentials: IceCredentials{UFrag: ufrag, Pwd: pwd}, SessionID: &sid}
}
func (s *workerSide) receive(msg OfferAnswer) {
before := s.agent()
s.w.OnNewOffer(&msg)
if after := s.agent(); before != nil && after != before {
s.replaced++
}
}
func (s *workerSide) teardowns() int { return s.replaced }
func (s *workerSide) agent() *icemaker.ThreadSafeAgent {
s.w.muxAgent.Lock()
defer s.w.muxAgent.Unlock()
return s.w.agent
}
// failAgent runs the cleanup the dial goroutine or the Failed state callback
// performs when the current negotiation dies.
func (s *workerSide) failAgent() {
s.t.Helper()
s.w.muxAgent.Lock()
agent, cancel := s.w.agent, s.w.agentDialerCancel
s.w.muxAgent.Unlock()
require.NotNil(s.t, agent, "failing requires a running negotiation")
s.w.closeAgent(agent, cancel)
}
// legacySide models a remote peer running a release from before this change:
// when it follows a remote restart it also picks a new session ID of its own,
// which it announces only with its next offer or answer.
type legacySide struct {
t *testing.T
sessionID ICESessionID
remoteID ICESessionID
hasAgent bool
replaced int
}
func newLegacySide(t *testing.T) *legacySide {
return &legacySide{t: t, sessionID: newTestSessionID(t)}
}
func (s *legacySide) message() OfferAnswer {
sid := s.sessionID
return OfferAnswer{
IceCredentials: IceCredentials{UFrag: "legacyufrag", Pwd: "legacy-password-long-enough"},
SessionID: &sid,
}
}
func (s *legacySide) receive(msg OfferAnswer) {
if msg.SessionID == nil {
s.hasAgent = true
return
}
if s.hasAgent {
if *msg.SessionID == s.remoteID {
return
}
s.replaced++
s.sessionID = newTestSessionID(s.t)
}
s.hasAgent = true
s.remoteID = *msg.SessionID
}
func (s *legacySide) teardowns() int { return s.replaced }
func (s *legacySide) failAgent() {
s.hasAgent = false
s.remoteID = ""
s.sessionID = newTestSessionID(s.t)
}
// exchange runs one guard-driven round in the order Handshaker.Listen uses: the
// answerer handles the offer and answers with the session ID it holds
// afterwards, and the offerer handles the answer without replying.
func exchange(offerer, answerer handshakeSide) {
answerer.receive(offerer.message())
offerer.receive(answerer.message())
}
// offerPattern decides which side's guard sends the offer in a round.
type offerPattern struct {
name string
picker func(round int, local, remote handshakeSide) (offerer, answerer handshakeSide)
}
var offerPatterns = []offerPattern{
{
// A routing peer whose relay is down keeps offering on its own.
name: "local peer offers",
picker: func(_ int, local, remote handshakeSide) (handshakeSide, handshakeSide) {
return local, remote
},
},
{
name: "both peers offer",
picker: func(round int, local, remote handshakeSide) (handshakeSide, handshakeSide) {
if round%2 == 0 {
return local, remote
}
return remote, local
},
},
}
// assertSettles runs guard rounds and requires the pair to stop restarting
// each other: at most maxTeardowns in total, and none once half the rounds ran.
func assertSettles(t *testing.T, pattern offerPattern, local, remote handshakeSide, maxTeardowns int) {
t.Helper()
const rounds = 10
total := func() int { return local.teardowns() + remote.teardowns() }
start := total()
var halfway int
for round := range rounds {
if round == rounds/2 {
halfway = total()
}
offerer, answerer := pattern.picker(round, local, remote)
exchange(offerer, answerer)
}
assert.LessOrEqual(t, total()-start, maxTeardowns, "the peers must not keep restarting each other")
assert.Equal(t, halfway, total(), "the negotiation must be stable in the later rounds")
}
// establish runs the first offer and answer, so both sides negotiate.
func establish(t *testing.T, local, remote handshakeSide) {
t.Helper()
exchange(local, remote)
require.Zero(t, local.teardowns()+remote.teardowns(), "the first exchange must not restart anything")
}
func TestICESession_SettlesAfterAgentFailure(t *testing.T) {
sides := []struct {
name string
remote func(t *testing.T) handshakeSide
}{
{name: "current remote", remote: func(t *testing.T) handshakeSide { return newWorkerSide(t) }},
{name: "legacy remote", remote: func(t *testing.T) handshakeSide { return newLegacySide(t) }},
}
failures := []struct {
name string
fail func(local, remote handshakeSide)
}{
{name: "remote agent fails", fail: func(_, remote handshakeSide) { remote.failAgent() }},
{name: "local agent fails", fail: func(local, _ handshakeSide) { local.failAgent() }},
{name: "both agents fail", fail: func(local, remote handshakeSide) {
local.failAgent()
remote.failAgent()
}},
}
for _, side := range sides {
for _, failure := range failures {
for _, pattern := range offerPatterns {
t.Run(fmt.Sprintf("%s/%s/%s", side.name, failure.name, pattern.name), func(t *testing.T) {
local := newWorkerSide(t)
remote := side.remote(t)
establish(t, local, remote)
failure.fail(local, remote)
assertSettles(t, pattern, local, remote, 2)
})
}
}
}
}
// TestICESession_LocalCloseRestartsRemote covers an explicit teardown, as on a
// WireGuard handshake timeout. The remote must start over as well, or it keeps
// answering from the negotiation this side just abandoned.
func TestICESession_LocalCloseRestartsRemote(t *testing.T) {
for _, pattern := range offerPatterns {
t.Run(pattern.name, func(t *testing.T) {
local := newWorkerSide(t)
remote := newWorkerSide(t)
establish(t, local, remote)
local.w.Close()
assertSettles(t, pattern, local, remote, 1)
assert.Equal(t, 1, remote.teardowns(), "the remote must restart its negotiation exactly once")
})
}
}
func TestICESession_DuplicateMessagesKeepNegotiation(t *testing.T) {
local := newWorkerSide(t)
remote := newWorkerSide(t)
offer := local.message()
remote.receive(offer)
answer := remote.message()
local.receive(answer)
// Signaling may deliver the same message again, and a peer answers every
// offer, including repeats of one it already handled.
remote.receive(offer)
local.receive(answer)
local.receive(remote.message())
assert.Zero(t, local.teardowns(), "a repeated answer must not restart the negotiation")
assert.Zero(t, remote.teardowns(), "a repeated offer must not restart the negotiation")
}
// TestICESession_RemoteWithoutSessionIDKeepsNegotiation covers remote peers
// too old to send session IDs: once negotiating, their messages cannot tell a
// restart from a repeat, so they must not tear anything down.
func TestICESession_RemoteWithoutSessionIDKeepsNegotiation(t *testing.T) {
local := newWorkerSide(t)
unversioned := OfferAnswer{IceCredentials: IceCredentials{UFrag: "oldufrag", Pwd: "old-password-long-enough"}}
local.receive(unversioned)
require.NotNil(t, local.agent(), "a message without a session ID must still start ICE")
advertised := local.w.SessionID()
for range 3 {
local.receive(unversioned)
}
assert.Zero(t, local.teardowns(), "messages without a session ID must not restart the negotiation")
assert.Equal(t, advertised, local.w.SessionID(), "the advertised session must not change")
}
// TestWorkerICE_StaleCleanupKeepsAdvertisedSession covers the cleanup of a
// replaced negotiation finishing late, from its dial goroutine or its Closed
// state callback. It must neither pick a new session ID, an unannounced local
// restart, nor disturb the negotiation that replaced it.
func TestWorkerICE_StaleCleanupKeepsAdvertisedSession(t *testing.T) {
local := newWorkerSide(t)
remote := newWorkerSide(t)
establish(t, local, remote)
local.w.muxAgent.Lock()
oldAgent, oldCancel := local.w.agent, local.w.agentDialerCancel
local.w.muxAgent.Unlock()
remote.failAgent()
exchange(local, remote)
require.Equal(t, 1, local.teardowns(), "the local side must follow the remote restart")
advertised := local.w.SessionID()
current := local.agent()
local.w.closeAgent(oldAgent, oldCancel)
assert.Equal(t, advertised, local.w.SessionID(), "a stale cleanup must not change the advertised session")
assert.Same(t, current, local.agent(), "a stale cleanup must keep the current negotiation")
assertSettles(t, offerPatterns[1], local, remote, 0)
}
+13 -14
View File
@@ -3,7 +3,6 @@ package peer
import (
"context"
"errors"
"net"
"net/netip"
"sync"
"sync/atomic"
@@ -14,7 +13,7 @@ import (
)
type RelayConnInfo struct {
relayedConn net.Conn
relayedConn *relayClient.Conn
rosenpassPubKey []byte
rosenpassAddr string
}
@@ -27,7 +26,7 @@ type WorkerRelay struct {
conn *Conn
relayManager *relayClient.Manager
relayedConn net.Conn
relayedConn *relayClient.Conn
relayLock sync.Mutex
relaySupportedOnRemotePeer atomic.Bool
@@ -80,12 +79,7 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
w.relayedConn = relayedConn
w.relayLock.Unlock()
err = w.relayManager.AddCloseListener(srv, w.onRelayClientDisconnected)
if err != nil {
log.Errorf("failed to add close listener: %s", err)
_ = relayedConn.Close()
return
}
go w.watchRelayedConn(relayedConn)
w.log.Debugf("peer conn opened via Relay: %s", srv)
go w.conn.onRelayConnectionIsReady(RelayConnInfo{
@@ -109,12 +103,15 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool {
func (w *WorkerRelay) CloseConn() {
w.relayLock.Lock()
defer w.relayLock.Unlock()
if w.relayedConn == nil {
conn := w.relayedConn
w.relayedConn = nil
w.relayLock.Unlock()
if conn == nil {
return
}
if err := w.relayedConn.Close(); err != nil {
if err := conn.Close(); err != nil {
w.log.Warnf("failed to close relay connection: %v", err)
}
}
@@ -133,6 +130,8 @@ func (w *WorkerRelay) preferredRelayServer(myRelayAddress, remoteRelayAddress st
return remoteRelayAddress
}
func (w *WorkerRelay) onRelayClientDisconnected() {
go w.conn.onRelayDisconnected()
func (w *WorkerRelay) watchRelayedConn(relayedConn *relayClient.Conn) {
<-relayedConn.Context().Done()
w.conn.onRelayDisconnected(relayedConn)
}