mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-06 21:49:08 +02:00
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:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user