diff --git a/client/internal/peer/worker_relay.go b/client/internal/peer/worker_relay.go index 0402992c9..449f82d8a 100644 --- a/client/internal/peer/worker_relay.go +++ b/client/internal/peer/worker_relay.go @@ -54,19 +54,15 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) { w.relaySupportedOnRemotePeer.Store(true) // the relayManager will return with error in case if the connection has lost with relay server - currentRelayAddress, _, err := w.relayManager.RelayInstanceAddress() + _, _, err := w.relayManager.RelayInstanceAddress() if err != nil { w.log.Errorf("failed to handle new offer: %s", err) return } - srv := w.preferredRelayServer(currentRelayAddress, remoteOfferAnswer.RelaySrvAddress) - var serverIP netip.Addr - if srv == remoteOfferAnswer.RelaySrvAddress { - serverIP = remoteOfferAnswer.RelaySrvIP - } - - relayedConn, err := w.relayManager.OpenConn(w.peerCtx, srv, w.config.Key, serverIP) + preferForeign := !w.isController + remoteRelayServer := relayClient.RelayServer{Addr: remoteOfferAnswer.RelaySrvAddress, IP: remoteOfferAnswer.RelaySrvIP} + relayedConn, err := w.relayManager.OpenConn(w.peerCtx, remoteRelayServer, w.config.Key, preferForeign) if err != nil { if errors.Is(err, relayClient.ErrConnAlreadyExists) { w.log.Debugf("handled offer by reusing existing relay connection") @@ -80,14 +76,13 @@ 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) + if err := w.relayManager.AddCloseListener(relayedConn.RemoteAddr().String(), w.onRelayClientDisconnected); err != nil { + w.log.Errorf("failed to add close listener: %s", err) _ = relayedConn.Close() return } - w.log.Debugf("peer conn opened via Relay: %s", srv) + w.log.Debugf("peer conn opened via Relay: %s", relayedConn.RemoteAddr()) go w.conn.onRelayConnectionIsReady(RelayConnInfo{ relayedConn: relayedConn, rosenpassPubKey: remoteOfferAnswer.RosenpassPubKey, @@ -126,13 +121,6 @@ func (w *WorkerRelay) isRelaySupported(answer *OfferAnswer) bool { return answer.RelaySrvAddress != "" } -func (w *WorkerRelay) preferredRelayServer(myRelayAddress, remoteRelayAddress string) string { - if w.isController { - return myRelayAddress - } - return remoteRelayAddress -} - func (w *WorkerRelay) onRelayClientDisconnected() { go w.conn.onRelayDisconnected() } diff --git a/shared/relay/client/fallback.go b/shared/relay/client/fallback.go new file mode 100644 index 000000000..c3434a46d --- /dev/null +++ b/shared/relay/client/fallback.go @@ -0,0 +1,158 @@ +package client + +import ( + "context" + "errors" + "net" + "time" + + log "github.com/sirupsen/logrus" +) + +const ( + raceTotalTimeout = 60 * time.Second + raceFallbackDelay = 10 * time.Second +) + +type raceAttempt struct { + conn net.Conn + err error +} + +type connRace struct { + racer *ConnRacer + peerKey string + remoteRelayServer RelayServer + preferForeign bool + + raceCtx context.Context + otherCtx context.Context + cancelPreferred context.CancelFunc + cancelOther context.CancelFunc + results chan raceAttempt + fallbackTimer *time.Timer + + otherStarted bool + settled int + lastErr error +} + +type ConnRacer struct { + home *Client + foreignStore *ForeignRelaysStore +} + +func NewConnRacer(home *Client, foreignStore *ForeignRelaysStore) *ConnRacer { + return &ConnRacer{ + home: home, + foreignStore: foreignStore, + } +} + +func (r *ConnRacer) Run(ctx context.Context, peerKey string, remoteRelayServer RelayServer, preferForeign bool) (net.Conn, error) { + raceCtx, cancel := context.WithTimeout(ctx, raceTotalTimeout) + defer cancel() + + preferredCtx, cancelPreferred := context.WithCancel(raceCtx) + otherCtx, cancelOther := context.WithCancel(raceCtx) + + race := &connRace{ + racer: r, + peerKey: peerKey, + remoteRelayServer: remoteRelayServer, + preferForeign: preferForeign, + raceCtx: raceCtx, + otherCtx: otherCtx, + cancelPreferred: cancelPreferred, + cancelOther: cancelOther, + results: make(chan raceAttempt, 2), + fallbackTimer: time.NewTimer(raceFallbackDelay), + } + defer race.fallbackTimer.Stop() + + go func() { + race.results <- r.open(preferredCtx, peerKey, remoteRelayServer, preferForeign) + }() + + for { + select { + case <-race.fallbackTimer.C: + race.startOther() + case res := <-race.results: + if conn, err, done := race.handleResult(res); done { + return conn, err + } + case <-raceCtx.Done(): + return race.onTimeout() + } + } +} + +func (c *connRace) startOther() { + if c.otherStarted { + return + } + c.otherStarted = true + c.fallbackTimer.Stop() + go func() { + c.results <- c.racer.open(c.otherCtx, c.peerKey, c.remoteRelayServer, !c.preferForeign) + }() +} + +func (c *connRace) handleResult(res raceAttempt) (net.Conn, error, bool) { + if (res.err == nil && res.conn != nil) || errors.Is(res.err, ErrConnAlreadyExists) { + c.stop() + return res.conn, res.err, true + } + + c.lastErr = res.err + c.settled++ + if !c.otherStarted { + c.startOther() + return nil, nil, false + } + if c.settled == 2 { + c.cancelPreferred() + c.cancelOther() + return nil, c.lastErr, true + } + return nil, nil, false +} + +func (c *connRace) onTimeout() (net.Conn, error) { + c.stop() + if c.lastErr != nil { + return nil, c.lastErr + } + return nil, c.raceCtx.Err() +} + +func (c *connRace) stop() { + c.cancelPreferred() + c.cancelOther() + go c.racer.drainLoser(c.results, c.settled, c.otherStarted) +} + +func (r *ConnRacer) open(ctx context.Context, peerKey string, remoteRelayServer RelayServer, foreign bool) raceAttempt { + if foreign { + conn, err := r.foreignStore.OpenConn(ctx, peerKey, remoteRelayServer) + return raceAttempt{conn: conn, err: err} + } + conn, err := r.home.OpenConn(ctx, peerKey) + return raceAttempt{conn: conn, err: err} +} + +func (r *ConnRacer) drainLoser(results chan raceAttempt, settled int, otherStarted bool) { + started := 1 + if otherStarted { + started = 2 + } + for i := settled; i < started; i++ { + res := <-results + if res.conn != nil { + if err := res.conn.Close(); err != nil { + log.Debugf("failed to close losing relay connection: %v", err) + } + } + } +} diff --git a/shared/relay/client/foreign_relays.go b/shared/relay/client/foreign_relays.go index fed510c5f..dca449480 100644 --- a/shared/relay/client/foreign_relays.go +++ b/shared/relay/client/foreign_relays.go @@ -3,7 +3,6 @@ package client import ( "context" "net" - "net/netip" "sync" "time" @@ -19,7 +18,7 @@ type foreignRelay struct { inUse int } -type foreignRelays struct { +type ForeignRelaysStore struct { mu sync.RWMutex clients map[string]*foreignRelay @@ -34,8 +33,8 @@ type foreignRelays struct { keepUnusedServerTime time.Duration } -func newForeignRelays(ctx context.Context, tokenStore *relayAuth.TokenStore, peerID string, mtu uint16, transportFallback *transportFallback, onDisconnect func(string), keepUnusedServerTime time.Duration) *foreignRelays { - return &foreignRelays{ +func NewForeignRelaysStore(ctx context.Context, tokenStore *relayAuth.TokenStore, peerID string, mtu uint16, transportFallback *transportFallback, onDisconnect func(string), keepUnusedServerTime time.Duration) *ForeignRelaysStore { + return &ForeignRelaysStore{ clients: make(map[string]*foreignRelay), ctx: ctx, tokenStore: tokenStore, @@ -47,8 +46,8 @@ func newForeignRelays(ctx context.Context, tokenStore *relayAuth.TokenStore, pee } } -func (f *foreignRelays) openConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) { - fr, err := f.acquire(serverAddress, serverIP) +func (f *ForeignRelaysStore) OpenConn(ctx context.Context, peerKey string, remoteRelayServer RelayServer) (net.Conn, error) { + fr, err := f.acquire(remoteRelayServer) if err != nil { return nil, err } @@ -57,24 +56,24 @@ func (f *foreignRelays) openConn(ctx context.Context, serverAddress, peerKey str return fr.client.OpenConn(ctx, peerKey) } -func (f *foreignRelays) acquire(serverAddress string, serverIP netip.Addr) (*foreignRelay, error) { +func (f *ForeignRelaysStore) acquire(remoteRelayServer RelayServer) (*foreignRelay, error) { f.mu.Lock() - if fr, ok := f.clients[serverAddress]; ok { + if fr, ok := f.clients[remoteRelayServer.Addr]; ok { fr.inUse++ f.mu.Unlock() return fr, nil } f.mu.Unlock() - v, err, _ := f.group.Do(serverAddress, func() (any, error) { + v, err, _ := f.group.Do(remoteRelayServer.Addr, func() (any, error) { f.mu.RLock() - fr, ok := f.clients[serverAddress] + fr, ok := f.clients[remoteRelayServer.Addr] f.mu.RUnlock() if ok { return fr, nil } - relayClient := NewClientWithServerIP(serverAddress, serverIP, f.tokenStore, f.peerID, f.mtu) + relayClient := NewClientWithServerIP(remoteRelayServer.Addr, remoteRelayServer.IP, f.tokenStore, f.peerID, f.mtu) relayClient.SetTransportFallback(f.transportFallback) if err := relayClient.Connect(f.ctx); err != nil { return nil, err @@ -83,7 +82,7 @@ func (f *foreignRelays) acquire(serverAddress string, serverIP netip.Addr) (*for f.mu.Lock() fr = &foreignRelay{client: relayClient, created: time.Now()} - f.clients[serverAddress] = fr + f.clients[remoteRelayServer.Addr] = fr f.mu.Unlock() return fr, nil }) @@ -93,22 +92,22 @@ func (f *foreignRelays) acquire(serverAddress string, serverIP netip.Addr) (*for fr := v.(*foreignRelay) f.mu.Lock() - if cur, ok := f.clients[serverAddress]; !ok || cur != fr { + if cur, ok := f.clients[remoteRelayServer.Addr]; !ok || cur != fr { f.mu.Unlock() - return f.acquire(serverAddress, serverIP) + return f.acquire(remoteRelayServer) } fr.inUse++ f.mu.Unlock() return fr, nil } -func (f *foreignRelays) release(fr *foreignRelay) { +func (f *ForeignRelaysStore) release(fr *foreignRelay) { f.mu.Lock() fr.inUse-- f.mu.Unlock() } -func (f *foreignRelays) evict(serverAddress string) { +func (f *ForeignRelaysStore) evict(serverAddress string) { f.mu.Lock() defer f.mu.Unlock() if _, ok := f.clients[serverAddress]; ok { @@ -117,7 +116,7 @@ func (f *foreignRelays) evict(serverAddress string) { } } -func (f *foreignRelays) cleanupUnused() { +func (f *ForeignRelaysStore) cleanupUnused() { f.mu.Lock() defer f.mu.Unlock() @@ -140,7 +139,7 @@ func (f *foreignRelays) cleanupUnused() { } } -func (f *foreignRelays) states() []RelayConnState { +func (f *ForeignRelaysStore) states() []RelayConnState { f.mu.RLock() clients := make([]*Client, 0, len(f.clients)) for _, fr := range f.clients { diff --git a/shared/relay/client/manager.go b/shared/relay/client/manager.go index 05cf2fa26..a550f9675 100644 --- a/shared/relay/client/manager.go +++ b/shared/relay/client/manager.go @@ -38,6 +38,11 @@ type RelayConnState struct { Err error } +type RelayServer struct { + Addr string + IP netip.Addr +} + // WithMaxBackoffInterval caps the exponential backoff between reconnect // attempts to the home relay. A non-positive value keeps the default. func WithMaxBackoffInterval(d time.Duration) ManagerOption { @@ -62,7 +67,7 @@ type Manager struct { relayClientMu sync.RWMutex reconnectGuard *Guard - foreign *foreignRelays + foreign *ForeignRelaysStore onDisconnectedListeners map[string]*list.List onReconnectedListenerFn func() @@ -105,7 +110,7 @@ func NewManager(ctx context.Context, serverURLs []string, peerID string, mtu uin for _, opt := range opts { opt(m) } - m.foreign = newForeignRelays(ctx, tokenStore, peerID, mtu, tf, m.onServerDisconnected, m.keepUnusedServerTime) + m.foreign = NewForeignRelaysStore(ctx, tokenStore, peerID, mtu, tf, m.onServerDisconnected, m.keepUnusedServerTime) m.serverPicker.ServerURLs.Store(serverURLs) m.reconnectGuard = NewGuard(m.serverPicker, m.maxBackoffInterval) return m @@ -137,13 +142,7 @@ func (m *Manager) Serve() error { return err } -// OpenConn opens a connection to the given peer key. If the peer is on the same relay server, the connection will be -// established via the relay server. If the peer is on a different relay server, the manager will establish a new -// connection to the relay server. It returns back with a net.Conn what represent the remote peer connection. -// -// serverIP, when valid and serverAddress is foreign, is used as a dial target if the FQDN-based dial fails. -// Ignored for the local home-server path. TLS verification still uses the FQDN via SNI. -func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) { +func (m *Manager) OpenConn(ctx context.Context, remoteRelayServer RelayServer, peerKey string, preferForeign bool) (net.Conn, error) { m.relayClientMu.RLock() defer m.relayClientMu.RUnlock() @@ -151,26 +150,17 @@ func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, s return nil, ErrRelayClientNotConnected } - foreign, err := m.isForeignServer(serverAddress) + foreign, err := m.isForeignServer(remoteRelayServer.Addr) if err != nil { return nil, err } - var ( - netConn net.Conn - ) if !foreign { - log.Debugf("open peer connection via permanent server: %s", peerKey) - netConn, err = m.relayClient.OpenConn(ctx, peerKey) - } else { - log.Debugf("open peer connection via foreign server: %s", serverAddress) - netConn, err = m.foreign.openConn(ctx, serverAddress, peerKey, serverIP) - } - if err != nil { - return nil, err + return m.relayClient.OpenConn(ctx, peerKey) } - return netConn, err + racer := NewConnRacer(m.relayClient, m.foreign) + return racer.Run(ctx, peerKey, remoteRelayServer, preferForeign) } // Ready returns true if the home Relay client is connected to the relay server.