diff --git a/client/android/client.go b/client/android/client.go index b870337d1..6f5eaacf3 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -498,6 +498,7 @@ func (c *Client) Networks() *NetworkArray { routesMap := routeManager.GetClientRoutesWithNetID() v6Merged := route.V6ExitMergeSet(routesMap) resolvedDomains := c.recorder.GetResolvedDomainsStates() + activeRoutePeers := c.recorder.GetActiveRoutePeers() networkArray := &NetworkArray{ items: make([]Network, 0), @@ -511,7 +512,7 @@ func (c *Client) Networks() *NetworkArray { continue } - network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged) + network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers) if network == nil { continue } @@ -520,14 +521,14 @@ func (c *Client) Networks() *NetworkArray { return networkArray } -func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network { +func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network { r := routes[0] netStr := r.Network.String() if r.IsDynamic() { netStr = r.Domains.SafeString() } - routePeer, err := c.findBestRoutePeer(routes) + routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers) if err != nil { log.Errorf("could not get peer info for route %s: %v", id, err) return nil @@ -551,12 +552,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo // findBestRoutePeer returns the peer actively routing traffic for the given // HA route group. Falls back to the first connected peer, then the first peer. -func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) { - netStr := routes[0].Network.String() - - fullStatus := c.recorder.GetFullStatus() - for _, p := range fullStatus.Peers { - if _, ok := p.GetRoutes()[netStr]; ok { +func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) { + if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok { + if p, err := c.recorder.GetPeer(peerKey); err == nil { return p, nil } } diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index d753ee43e..826bf6fe0 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -196,6 +196,7 @@ type Status struct { muxRelays sync.RWMutex peers map[string]State ipToKey map[string]string + activeRoutePeers map[route.HAUniqueID]string changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription signalState bool signalError error @@ -257,6 +258,7 @@ func NewRecorder(mgmAddress string) *Status { return &Status{ peers: make(map[string]State), ipToKey: make(map[string]string), + activeRoutePeers: make(map[route.HAUniqueID]string), changeNotify: make(map[string]map[string]*StatusChangeSubscription), eventStreams: make(map[string]chan *proto.SystemEvent), eventQueue: NewEventQueue(eventQueueSize), @@ -481,6 +483,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error { return nil } +func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) { + d.mux.Lock() + defer d.mux.Unlock() + d.activeRoutePeers[haID] = peer +} + +func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) { + d.mux.Lock() + defer d.mux.Unlock() + delete(d.activeRoutePeers, haID) +} + +func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string { + d.mux.RLock() + defer d.mux.RUnlock() + return maps.Clone(d.activeRoutePeers) +} + // CheckRoutes checks if the source and destination addresses are within the same route // and returns the resource ID of the route that contains the addresses func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) { diff --git a/client/internal/peer/status_test.go b/client/internal/peer/status_test.go index 82dff0d6f..b3f01b217 100644 --- a/client/internal/peer/status_test.go +++ b/client/internal/peer/status_test.go @@ -9,6 +9,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/route" ) func TestAddPeer(t *testing.T) { @@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) { status.MarkManagementDisconnected(err) assert.False(t, notified(ch), "redundant disconnect should not notify") } + +func TestActiveRoutePeers(t *testing.T) { + status := NewRecorder("https://mgm") + netA := route.HAUniqueID("net-a-10.0.0.0/24") + netB := route.HAUniqueID("net-b-10.0.0.0/24") + + status.AddActiveRoutePeer(netA, "peerA") + status.AddActiveRoutePeer(netB, "peerB") + + active := status.GetActiveRoutePeers() + assert.Equal(t, "peerA", active[netA]) + assert.Equal(t, "peerB", active[netB]) + + status.RemoveActiveRoutePeer(netA) + delete(active, netB) + + active = status.GetActiveRoutePeers() + _, ok := active[netA] + assert.False(t, ok) + assert.Equal(t, "peerB", active[netB]) +} diff --git a/client/internal/routemanager/client/client.go b/client/internal/routemanager/client/client.go index c691c54f8..973cf1ab8 100644 --- a/client/internal/routemanager/client/client.go +++ b/client/internal/routemanager/client/client.go @@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err) } + w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer) if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil { log.Warnf("Failed to update peer state: %v", err) } @@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { } func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error { + w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID()) if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil { log.Warnf("Failed to update peer state: %v", err) }