diff --git a/olm/peer.go b/olm/peer.go index 7c4609b..610fc93 100644 --- a/olm/peer.go +++ b/olm/peer.go @@ -228,28 +228,49 @@ func (o *Olm) handleWgPeerRelay(msg websocket.WSMessage) { var relayData struct { peers.RelayPeerData - ChainId string `json:"chainId"` + ChainId string `json:"chainId"` + ChainIds []string `json:"chainIds"` } if err := json.Unmarshal(jsonData, &relayData); err != nil { logger.Error("Error unmarshaling relay data: %v", err) return } - if monitor := pm.GetPeerMonitor(); monitor != nil { - monitor.CancelRelaySend(relayData.ChainId) + siteIds := relayData.SiteIds + relayEndpoints := relayData.RelayEndpoints + chainIds := relayData.ChainIds + if len(siteIds) == 0 { + siteIds = []int{relayData.SiteId} + relayEndpoints = []string{relayData.RelayEndpoint} + chainIds = []string{relayData.ChainId} } - primaryRelay, err := util.ResolveDomainUpstream(relayData.RelayEndpoint, o.tunnelConfig.PublicDNS) + monitor := pm.GetPeerMonitor() - if err != nil { - logger.Error("Failed to resolve primary relay endpoint: %v", err) - return + for i, siteId := range siteIds { + var relayEndpoint, chainId string + if i < len(relayEndpoints) { + relayEndpoint = relayEndpoints[i] + } + if i < len(chainIds) { + chainId = chainIds[i] + } + + if monitor != nil { + monitor.CancelRelaySend(chainId) + } + + primaryRelay, err := util.ResolveDomainUpstream(relayEndpoint, o.tunnelConfig.PublicDNS) + if err != nil { + logger.Error("Failed to resolve primary relay endpoint for site %d: %v", siteId, err) + continue + } + + // Update HTTP server to mark this peer as using relay + o.apiServer.UpdatePeerRelayStatus(siteId, relayEndpoint, true) + + pm.RelayPeer(siteId, primaryRelay, relayData.RelayPort) } - - // Update HTTP server to mark this peer as using relay - o.apiServer.UpdatePeerRelayStatus(relayData.SiteId, relayData.RelayEndpoint, true) - - pm.RelayPeer(relayData.SiteId, primaryRelay, relayData.RelayPort) } func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) { @@ -270,27 +291,48 @@ func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) { var relayData struct { peers.UnRelayPeerData - ChainId string `json:"chainId"` + ChainId string `json:"chainId"` + ChainIds []string `json:"chainIds"` } if err := json.Unmarshal(jsonData, &relayData); err != nil { logger.Error("Error unmarshaling relay data: %v", err) return } - if monitor := pm.GetPeerMonitor(); monitor != nil { - monitor.CancelRelaySend(relayData.ChainId) + siteIds := relayData.SiteIds + endpoints := relayData.Endpoints + chainIds := relayData.ChainIds + if len(siteIds) == 0 { + siteIds = []int{relayData.SiteId} + endpoints = []string{relayData.Endpoint} + chainIds = []string{relayData.ChainId} } - primaryRelay, err := util.ResolveDomainUpstream(relayData.Endpoint, o.tunnelConfig.PublicDNS) + monitor := pm.GetPeerMonitor() - if err != nil { - logger.Warn("Failed to resolve primary relay endpoint: %v", err) + for i, siteId := range siteIds { + var endpoint, chainId string + if i < len(endpoints) { + endpoint = endpoints[i] + } + if i < len(chainIds) { + chainId = chainIds[i] + } + + if monitor != nil { + monitor.CancelRelaySend(chainId) + } + + primaryRelay, err := util.ResolveDomainUpstream(endpoint, o.tunnelConfig.PublicDNS) + if err != nil { + logger.Warn("Failed to resolve primary relay endpoint for site %d: %v", siteId, err) + } + + // Update HTTP server to mark this peer as using relay + o.apiServer.UpdatePeerRelayStatus(siteId, endpoint, false) + + pm.UnRelayPeer(siteId, primaryRelay) } - - // Update HTTP server to mark this peer as using relay - o.apiServer.UpdatePeerRelayStatus(relayData.SiteId, relayData.Endpoint, false) - - pm.UnRelayPeer(relayData.SiteId, primaryRelay) } // handleWgPeerLocal handles the server's acknowledgement of an "olm/wg/local" message. @@ -314,15 +356,23 @@ func (o *Olm) handleWgPeerLocal(msg websocket.WSMessage) { var localData struct { peers.LocalPeerAckData - ChainId string `json:"chainId"` + ChainId string `json:"chainId"` + ChainIds []string `json:"chainIds"` } if err := json.Unmarshal(jsonData, &localData); err != nil { logger.Error("Error unmarshaling local ack data: %v", err) return } + chainIds := localData.ChainIds + if len(chainIds) == 0 { + chainIds = []string{localData.ChainId} + } + if monitor := pm.GetPeerMonitor(); monitor != nil { - monitor.CancelLocalSend(localData.ChainId) + for _, chainId := range chainIds { + monitor.CancelLocalSend(chainId) + } } } @@ -346,15 +396,23 @@ func (o *Olm) handleWgPeerUnlocal(msg websocket.WSMessage) { var localData struct { peers.LocalPeerAckData - ChainId string `json:"chainId"` + ChainId string `json:"chainId"` + ChainIds []string `json:"chainIds"` } if err := json.Unmarshal(jsonData, &localData); err != nil { logger.Error("Error unmarshaling unlocal ack data: %v", err) return } + chainIds := localData.ChainIds + if len(chainIds) == 0 { + chainIds = []string{localData.ChainId} + } + if monitor := pm.GetPeerMonitor(); monitor != nil { - monitor.CancelLocalSend(localData.ChainId) + for _, chainId := range chainIds { + monitor.CancelLocalSend(chainId) + } } } diff --git a/peers/monitor/batch_sender.go b/peers/monitor/batch_sender.go new file mode 100644 index 0000000..a551732 --- /dev/null +++ b/peers/monitor/batch_sender.go @@ -0,0 +1,212 @@ +package monitor + +import ( + "sync" + "time" + + "github.com/fosrl/newt/logger" + "github.com/fosrl/olm/websocket" +) + +// Batch tuning: each new item added resets a batchDebounce trailing window, so a burst of +// decisions that trickle in over a few hundred ms (e.g. 100 sites each independently +// finishing their own rapid holepunch test after a network-wide blip) still collapses into +// one flush instead of splitting across several. batchMaxWait bounds how long a steady +// trickle of arrivals can keep postponing that flush, so an item is never held back more +// than batchMaxWait past its own arrival. Once sent, items the server hasn't acknowledged +// yet are retried together every batchInterval, up to batchMaxAttempts times, mirroring the +// cadence the previous per-site SendMessageInterval senders used. +const ( + batchDebounce = 200 * time.Millisecond + batchMaxWait = 750 * time.Millisecond + batchInterval = 2 * time.Second + batchMaxAttempts = 10 +) + +// batchSendItem is a single queued site decision awaiting acknowledgement. +type batchSendItem struct { + siteID int + endpoint string // only populated for the local batcher; ignored otherwise + attempts int +} + +// batchSender coalesces per-site "olm/wg/*" notifications of a single message type into +// periodic batched websocket messages, retrying items the server hasn't acknowledged (via +// cancel) until they succeed or exhaust batchMaxAttempts. When the connected server doesn't +// understand the batched wire format (see websocket.Client.SupportsBatchedSiteMessages), +// it falls back to sending each item as its own message in the pre-batching singular form. +type batchSender struct { + mu sync.Mutex + items map[string]*batchSendItem // chainId -> item + messageType string + wsClient *websocket.Client + debounceTimer *time.Timer + burstStarted time.Time + runOnce sync.Once + stopChan chan struct{} +} + +func newBatchSender(wsClient *websocket.Client, messageType string) *batchSender { + return &batchSender{ + items: make(map[string]*batchSendItem), + messageType: messageType, + wsClient: wsClient, + stopChan: make(chan struct{}), + } +} + +// add queues siteID for the next flush and returns the chainId that identifies it for +// cancel/ack purposes. endpoint is only meaningful for the local batcher. +func (b *batchSender) add(siteID int, endpoint string) string { + chainId := generateChainId() + + b.mu.Lock() + b.items[chainId] = &batchSendItem{siteID: siteID, endpoint: endpoint} + + now := time.Now() + if b.debounceTimer == nil { + // First item of a new burst: start the trailing window. + b.burstStarted = now + b.debounceTimer = time.AfterFunc(batchDebounce, b.flush) + } else { + // Extend the window for this new arrival, capped so a burst that keeps trickling + // in still flushes within batchMaxWait of its first item. + delay := batchDebounce + if remaining := batchMaxWait - now.Sub(b.burstStarted); remaining < delay { + if remaining < 0 { + remaining = 0 + } + delay = remaining + } + b.debounceTimer.Reset(delay) + } + b.mu.Unlock() + + b.runOnce.Do(func() { go b.run() }) + + return chainId +} + +// cancel removes chainId from the pending set, e.g. once the server has acknowledged it. +// Returns true if it was pending. +func (b *batchSender) cancel(chainId string) bool { + b.mu.Lock() + defer b.mu.Unlock() + if _, ok := b.items[chainId]; !ok { + return false + } + delete(b.items, chainId) + return true +} + +// cancelAll clears all pending items, e.g. on shutdown. +func (b *batchSender) cancelAll() { + b.mu.Lock() + b.items = make(map[string]*batchSendItem) + b.mu.Unlock() +} + +type readySend struct { + chainId string + siteID int + endpoint string +} + +// flush sends everything currently pending, dropping items that have exhausted +// batchMaxAttempts. If the server supports the batched wire format it goes out as a single +// message; otherwise each item is sent individually in the pre-batching singular form. +func (b *batchSender) flush() { + b.mu.Lock() + b.debounceTimer = nil + + ready := make([]readySend, 0, len(b.items)) + for chainId, item := range b.items { + item.attempts++ + if item.attempts > batchMaxAttempts { + logger.Warn("olm: giving up on %s for site %d (chain %s) after %d attempts", b.messageType, item.siteID, chainId, batchMaxAttempts) + delete(b.items, chainId) + continue + } + ready = append(ready, readySend{chainId: chainId, siteID: item.siteID, endpoint: item.endpoint}) + } + wsClient := b.wsClient + messageType := b.messageType + b.mu.Unlock() + + if len(ready) == 0 || wsClient == nil { + return + } + + if !wsClient.SupportsBatchedSiteMessages() { + b.sendIndividually(wsClient, messageType, ready) + return + } + + siteIds := make([]int, len(ready)) + chainIds := make([]string, len(ready)) + endpoints := make([]string, len(ready)) + hasEndpoints := false + for i, item := range ready { + siteIds[i] = item.siteID + chainIds[i] = item.chainId + endpoints[i] = item.endpoint + if item.endpoint != "" { + hasEndpoints = true + } + } + + data := map[string]interface{}{ + "siteIds": siteIds, + "chainIds": chainIds, + } + if hasEndpoints { + data["endpoints"] = endpoints + } + + if err := wsClient.SendMessage(messageType, data); err != nil { + logger.Error("olm: failed to send batched %s: %v", messageType, err) + } else { + logger.Info("olm: sent batched %s for %d site(s)", messageType, len(siteIds)) + } +} + +// sendIndividually sends each item as its own message in the singular siteId/chainId form, +// for servers that predate the batched wire format. +func (b *batchSender) sendIndividually(wsClient *websocket.Client, messageType string, ready []readySend) { + for _, item := range ready { + data := map[string]interface{}{ + "siteId": item.siteID, + "chainId": item.chainId, + } + if item.endpoint != "" { + data["endpoint"] = item.endpoint + } + if err := wsClient.SendMessage(messageType, data); err != nil { + logger.Error("olm: failed to send %s for site %d: %v", messageType, item.siteID, err) + } + } +} + +// run periodically resends any items still awaiting acknowledgement. +func (b *batchSender) run() { + ticker := time.NewTicker(batchInterval) + defer ticker.Stop() + for { + select { + case <-b.stopChan: + return + case <-ticker.C: + b.flush() + } + } +} + +// close stops the background retry loop. The sender must not be used again afterwards. +func (b *batchSender) close() { + select { + case <-b.stopChan: + // already closed + default: + close(b.stopChan) + } +} diff --git a/peers/monitor/monitor.go b/peers/monitor/monitor.go index 2796b41..2820ca0 100644 --- a/peers/monitor/monitor.go +++ b/peers/monitor/monitor.go @@ -40,9 +40,10 @@ type PeerMonitor struct { wsClient *websocket.Client publicDNS []string - // Relay sender tracking - relaySends map[string]func() - relaySendMu sync.Mutex + // Relay sender tracking — batched per message type so many relay/unrelay decisions + // arriving together (e.g. after a network-wide blip) go out as one websocket message. + relayBatch *batchSender + unrelayBatch *batchSender // Netstack fields middleDev *middleDevice.MiddleDevice @@ -82,9 +83,9 @@ type PeerMonitor struct { localSwitchCallback func(siteId int, endpoint string) // invoked when a local endpoint becomes active localFallbackCallback func(siteId int) // invoked when we fall back from a local endpoint - // Local connection sender tracking, keyed by chainId (informational messages only) - localSends map[string]func() - localSendMu sync.Mutex + // Local connection sender tracking (informational messages only), batched the same way. + localBatch *batchSender + unlocalBatch *batchSender // Exponential backoff fields for holepunch monitor defaultHolepunchMinInterval time.Duration // Minimum interval (initial) @@ -149,14 +150,16 @@ func NewPeerMonitor(wsClient *websocket.Client, middleDev *middleDevice.MiddleDe holepunchEndpoints: make(map[int]string), holepunchStatus: make(map[int]bool), relayedPeers: make(map[int]bool), - relaySends: make(map[string]func()), + relayBatch: newBatchSender(wsClient, "olm/wg/relay"), + unrelayBatch: newBatchSender(wsClient, "olm/wg/unrelay"), holepunchMaxAttempts: 3, // Trigger relay after 3 consecutive failures holepunchFailures: make(map[int]int), localEndpoints: make(map[int][]string), localActiveEndpoint: make(map[int]string), localFailures: make(map[int]int), localTestTimeout: 300 * time.Millisecond, // local network round trips should be fast - localSends: make(map[string]func()), + localBatch: newBatchSender(wsClient, "olm/wg/local"), + unlocalBatch: newBatchSender(wsClient, "olm/wg/unlocal"), // Rapid initial test settings: complete within ~1.5 seconds rapidTestInterval: 200 * time.Millisecond, // 200ms between attempts rapidTestTimeout: 400 * time.Millisecond, // 400ms timeout per attempt @@ -628,23 +631,16 @@ func (pm *PeerMonitor) handleConnectionStatusChange(siteID int, status Connectio } } -// sendRelay sends a relay message to the server with retry, keyed by chainId +// sendRelay queues a relay message for the server, batched with any other relay decisions +// made in the same short window, with retry keyed by chainId. func (pm *PeerMonitor) sendRelay(siteID int) error { if pm.wsClient == nil { return fmt.Errorf("websocket client is nil") } - chainId := generateChainId() - stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/relay", map[string]interface{}{ - "siteId": siteID, - "chainId": chainId, - }, 2*time.Second, 10) + chainId := pm.relayBatch.add(siteID, "") - pm.relaySendMu.Lock() - pm.relaySends[chainId] = stopFunc - pm.relaySendMu.Unlock() - - logger.Info("Sent relay message for site %d (chain %s)", siteID, chainId) + logger.Info("Queued relay message for site %d (chain %s)", siteID, chainId) return nil } @@ -654,23 +650,16 @@ func (pm *PeerMonitor) RequestRelay(siteID int) error { return pm.sendRelay(siteID) } -// sendUnRelay sends an unrelay message to the server with retry, keyed by chainId +// sendUnRelay queues an unrelay message for the server, batched with any other unrelay +// decisions made in the same short window, with retry keyed by chainId. func (pm *PeerMonitor) sendUnRelay(siteID int) error { if pm.wsClient == nil { return fmt.Errorf("websocket client is nil") } - chainId := generateChainId() - stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/unrelay", map[string]interface{}{ - "siteId": siteID, - "chainId": chainId, - }, 2*time.Second, 10) + chainId := pm.unrelayBatch.add(siteID, "") - pm.relaySendMu.Lock() - pm.relaySends[chainId] = stopFunc - pm.relaySendMu.Unlock() - - logger.Info("Sent unrelay message for site %d (chain %s)", siteID, chainId) + logger.Info("Queued unrelay message for site %d (chain %s)", siteID, chainId) return nil } @@ -683,18 +672,9 @@ func (pm *PeerMonitor) sendLocal(siteID int, endpoint string) { return } - chainId := generateChainId() - stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/local", map[string]interface{}{ - "siteId": siteID, - "endpoint": endpoint, - "chainId": chainId, - }, 2*time.Second, 10) + chainId := pm.localBatch.add(siteID, endpoint) - pm.localSendMu.Lock() - pm.localSends[chainId] = stopFunc - pm.localSendMu.Unlock() - - logger.Info("Sent local-connection message for site %d (%s, chain %s)", siteID, endpoint, chainId) + logger.Info("Queued local-connection message for site %d (%s, chain %s)", siteID, endpoint, chainId) } // sendUnLocal notifies the server that this peer fell back from its local network endpoint, @@ -704,65 +684,39 @@ func (pm *PeerMonitor) sendUnLocal(siteID int) { return } - chainId := generateChainId() - stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/unlocal", map[string]interface{}{ - "siteId": siteID, - "chainId": chainId, - }, 2*time.Second, 10) + chainId := pm.unlocalBatch.add(siteID, "") - pm.localSendMu.Lock() - pm.localSends[chainId] = stopFunc - pm.localSendMu.Unlock() - - logger.Info("Sent unlocal-connection message for site %d (chain %s)", siteID, chainId) + logger.Info("Queued unlocal-connection message for site %d (chain %s)", siteID, chainId) } -// CancelLocalSend stops the interval sender for the given chainId, if one exists. -// If chainId is empty, all active local-connection senders are stopped. +// CancelLocalSend removes chainId from the pending local/unlocal batches, if present. +// If chainId is empty, all pending local-connection items are cleared. func (pm *PeerMonitor) CancelLocalSend(chainId string) { - pm.localSendMu.Lock() - defer pm.localSendMu.Unlock() - if chainId == "" { - for id, stop := range pm.localSends { - if stop != nil { - stop() - } - delete(pm.localSends, id) - } + pm.localBatch.cancelAll() + pm.unlocalBatch.cancelAll() logger.Info("Cancelled all local-connection senders") return } - if stop, ok := pm.localSends[chainId]; ok { - stop() - delete(pm.localSends, chainId) + if pm.localBatch.cancel(chainId) || pm.unlocalBatch.cancel(chainId) { logger.Info("Cancelled local-connection sender for chain %s", chainId) } else { logger.Warn("CancelLocalSend: no active sender for chain %s", chainId) } } -// CancelRelaySend stops the interval sender for the given chainId, if one exists. -// If chainId is empty, all active relay senders are stopped. +// CancelRelaySend removes chainId from the pending relay/unrelay batches, if present. +// If chainId is empty, all pending relay items are cleared. func (pm *PeerMonitor) CancelRelaySend(chainId string) { - pm.relaySendMu.Lock() - defer pm.relaySendMu.Unlock() - if chainId == "" { - for id, stop := range pm.relaySends { - if stop != nil { - stop() - } - delete(pm.relaySends, id) - } + pm.relayBatch.cancelAll() + pm.unrelayBatch.cancelAll() logger.Info("Cancelled all relay senders") return } - if stop, ok := pm.relaySends[chainId]; ok { - stop() - delete(pm.relaySends, chainId) + if pm.relayBatch.cancel(chainId) || pm.unrelayBatch.cancel(chainId) { logger.Info("Cancelled relay sender for chain %s", chainId) } else { logger.Warn("CancelRelaySend: no active sender for chain %s", chainId) @@ -1177,25 +1131,17 @@ func (pm *PeerMonitor) Close() { } pm.exitNodeMu.Unlock() - // Stop all pending relay senders - pm.relaySendMu.Lock() - for chainId, stop := range pm.relaySends { - if stop != nil { - stop() - } - delete(pm.relaySends, chainId) - } - pm.relaySendMu.Unlock() + // Stop all pending relay/unrelay batch senders + pm.relayBatch.cancelAll() + pm.relayBatch.close() + pm.unrelayBatch.cancelAll() + pm.unrelayBatch.close() - // Stop all pending local-connection senders - pm.localSendMu.Lock() - for chainId, stop := range pm.localSends { - if stop != nil { - stop() - } - delete(pm.localSends, chainId) - } - pm.localSendMu.Unlock() + // Stop all pending local-connection batch senders + pm.localBatch.cancelAll() + pm.localBatch.close() + pm.unlocalBatch.cancelAll() + pm.unlocalBatch.close() pm.mutex.Lock() defer pm.mutex.Unlock() diff --git a/peers/types.go b/peers/types.go index 8d4f04f..f97b5fb 100644 --- a/peers/types.go +++ b/peers/types.go @@ -35,15 +35,26 @@ type PeerRemove struct { SiteId int `json:"siteId"` } +// RelayPeerData represents the server's acknowledgement of an "olm/wg/relay" message. The +// server replies in kind: SiteId/RelayEndpoint for a single-site request, or the parallel +// SiteIds/RelayEndpoints arrays for a batched one (RelayPort is shared by the whole batch). type RelayPeerData struct { - SiteId int `json:"siteId"` - RelayEndpoint string `json:"relayEndpoint"` + SiteId int `json:"siteId,omitempty"` + RelayEndpoint string `json:"relayEndpoint,omitempty"` RelayPort uint16 `json:"relayPort"` + + SiteIds []int `json:"siteIds,omitempty"` + RelayEndpoints []string `json:"relayEndpoints,omitempty"` } +// UnRelayPeerData represents the server's acknowledgement of an "olm/wg/unrelay" message, +// in either its single-site or batched (parallel-array) form. See RelayPeerData. type UnRelayPeerData struct { - SiteId int `json:"siteId"` - Endpoint string `json:"endpoint"` + SiteId int `json:"siteId,omitempty"` + Endpoint string `json:"endpoint,omitempty"` + + SiteIds []int `json:"siteIds,omitempty"` + Endpoints []string `json:"endpoints,omitempty"` } // LocalPeerAckData represents the server's acknowledgement of an "olm/wg/local" or diff --git a/websocket/client.go b/websocket/client.go index 9e88981..ecf04d8 100644 --- a/websocket/client.go +++ b/websocket/client.go @@ -104,7 +104,8 @@ type Client struct { configVersionMux sync.RWMutex token string // Cached authentication token exitNodes []ExitNode // Cached exit nodes from token response - tokenMux sync.RWMutex // Protects token and exitNodes + serverVersion string // Server version from the last token response + tokenMux sync.RWMutex // Protects token, exitNodes and serverVersion forceNewToken bool // Flag to force fetching a new token on next connection processingMessage bool // Flag to track if a message is currently being processed processingMux sync.RWMutex // Protects processingMessage @@ -423,11 +424,11 @@ func (c *Client) RegisterHandler(messageType string, handler MessageHandler) { c.handlers[messageType] = handler } -func (c *Client) getToken() (string, []ExitNode, error) { +func (c *Client) getToken() (string, []ExitNode, string, error) { // Parse the base URL to ensure we have the correct hostname baseURL, err := url.Parse(c.baseURL) if err != nil { - return "", nil, fmt.Errorf("failed to parse base URL: %w", err) + return "", nil, "", fmt.Errorf("failed to parse base URL: %w", err) } // Ensure we have the base URL without trailing slashes @@ -439,7 +440,7 @@ func (c *Client) getToken() (string, []ExitNode, error) { if c.tlsConfig.ClientCertFile != "" || c.tlsConfig.ClientKeyFile != "" || len(c.tlsConfig.CAFiles) > 0 || c.tlsConfig.PKCS12File != "" { tlsConfig, err = c.setupTLS() if err != nil { - return "", nil, fmt.Errorf("failed to setup TLS configuration: %w", err) + return "", nil, "", fmt.Errorf("failed to setup TLS configuration: %w", err) } } @@ -461,7 +462,7 @@ func (c *Client) getToken() (string, []ExitNode, error) { jsonData, err := json.Marshal(tokenData) if err != nil { - return "", nil, fmt.Errorf("failed to marshal token request data: %w", err) + return "", nil, "", fmt.Errorf("failed to marshal token request data: %w", err) } // Create a new request @@ -471,7 +472,7 @@ func (c *Client) getToken() (string, []ExitNode, error) { bytes.NewBuffer(jsonData), ) if err != nil { - return "", nil, fmt.Errorf("failed to create request: %w", err) + return "", nil, "", fmt.Errorf("failed to create request: %w", err) } // Set headers @@ -490,7 +491,7 @@ func (c *Client) getToken() (string, []ExitNode, error) { } resp, err := client.Do(req) if err != nil { - return "", nil, fmt.Errorf("failed to request new token: %w", err) + return "", nil, "", fmt.Errorf("failed to request new token: %w", err) } defer resp.Body.Close() @@ -500,33 +501,33 @@ func (c *Client) getToken() (string, []ExitNode, error) { // Return AuthError for 401/403 status codes if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { - return "", nil, &AuthError{ + return "", nil, "", &AuthError{ StatusCode: resp.StatusCode, Message: string(body), } } // For other errors (5xx, network issues, etc.), return regular error - return "", nil, fmt.Errorf("failed to get token with status code: %d, body: %s", resp.StatusCode, string(body)) + return "", nil, "", fmt.Errorf("failed to get token with status code: %d, body: %s", resp.StatusCode, string(body)) } var tokenResp TokenResponse if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil { logger.Error("websocket: Failed to decode token response.") - return "", nil, fmt.Errorf("failed to decode token response: %w", err) + return "", nil, "", fmt.Errorf("failed to decode token response: %w", err) } if !tokenResp.Success { - return "", nil, fmt.Errorf("failed to get token: %s", tokenResp.Message) + return "", nil, "", fmt.Errorf("failed to get token: %s", tokenResp.Message) } if tokenResp.Data.Token == "" { - return "", nil, fmt.Errorf("received empty token from server") + return "", nil, "", fmt.Errorf("received empty token from server") } logger.Debug("websocket: Received token: %s", tokenResp.Data.Token) - return tokenResp.Data.Token, tokenResp.Data.ExitNodes, nil + return tokenResp.Data.Token, tokenResp.Data.ExitNodes, tokenResp.Data.ServerVersion, nil } func (c *Client) connectWithRetry() { @@ -564,13 +565,14 @@ func (c *Client) establishConnection() error { c.tokenMux.Lock() needNewToken := c.token == "" || c.forceNewToken if needNewToken { - token, exitNodes, err := c.getToken() + token, exitNodes, serverVersion, err := c.getToken() if err != nil { c.tokenMux.Unlock() return fmt.Errorf("failed to get token: %w", err) } c.token = token c.exitNodes = exitNodes + c.serverVersion = serverVersion c.forceNewToken = false if c.onTokenUpdate != nil { diff --git a/websocket/version.go b/websocket/version.go new file mode 100644 index 0000000..eaba0f4 --- /dev/null +++ b/websocket/version.go @@ -0,0 +1,75 @@ +package websocket + +import ( + "strconv" + "strings" +) + +// minBatchedSiteMessagesVersion is the first pangolin server version that understands the +// batched siteIds/chainIds (and endpoints/relayEndpoints) form of the olm/wg/relay, +// olm/wg/unrelay, olm/wg/local and olm/wg/unlocal messages. Servers older than this only +// understand the singular siteId/chainId form, so olm must fall back to sending those +// messages one at a time. +const minBatchedSiteMessagesVersion = "1.24.0" + +// SupportsBatchedSiteMessages reports whether the server we're connected to is new enough +// to understand batched relay/unrelay/local/unlocal messages, based on the serverVersion +// returned in the last token exchange. +func (c *Client) SupportsBatchedSiteMessages() bool { + return supportsBatchedSiteMessages(c.ServerVersion()) +} + +// ServerVersion returns the version reported by the server during the last token exchange, +// or "" if unknown (e.g. no successful connection has completed yet). +func (c *Client) ServerVersion() string { + c.tokenMux.RLock() + defer c.tokenMux.RUnlock() + return c.serverVersion +} + +func supportsBatchedSiteMessages(serverVersion string) bool { + if serverVersion == "" { + // Unknown server version: assume it predates batching rather than risk sending a + // format an old server can't parse. + return false + } + return compareVersions(baseVersion(serverVersion), minBatchedSiteMessagesVersion) >= 0 +} + +// baseVersion strips any "-suffix" build metadata (e.g. "1.24.0-s.5" -> "1.24.0"). +func baseVersion(v string) string { + if i := strings.IndexByte(v, '-'); i >= 0 { + return v[:i] + } + return v +} + +// compareVersions compares two dotted numeric version strings (e.g. "1.24.0"), returning +// -1 if a < b, 0 if equal, or 1 if a > b. Missing or non-numeric components are treated as +// 0, so a partial or malformed version degrades gracefully instead of panicking. +func compareVersions(a, b string) int { + as := strings.Split(a, ".") + bs := strings.Split(b, ".") + + max := len(as) + if len(bs) > max { + max = len(bs) + } + + for i := 0; i < max; i++ { + var an, bn int + if i < len(as) { + an, _ = strconv.Atoi(as[i]) + } + if i < len(bs) { + bn, _ = strconv.Atoi(bs[i]) + } + if an != bn { + if an < bn { + return -1 + } + return 1 + } + } + return 0 +} diff --git a/websocket/version_test.go b/websocket/version_test.go new file mode 100644 index 0000000..0bcfe38 --- /dev/null +++ b/websocket/version_test.go @@ -0,0 +1,47 @@ +package websocket + +import "testing" + +func TestSupportsBatchedSiteMessages(t *testing.T) { + cases := []struct { + version string + want bool + }{ + {"", false}, + {"1.23.0", false}, + {"1.23.9", false}, + {"1.24.0", true}, + {"1.24.1", true}, + {"1.25.0", true}, + {"2.0.0", true}, + {"1.24.0-s.5", true}, // cloud build of a supported base version + {"1.23.0-s.99", false}, // cloud build of an unsupported base version + {"not-a-version", false}, + } + + for _, c := range cases { + if got := supportsBatchedSiteMessages(c.version); got != c.want { + t.Errorf("supportsBatchedSiteMessages(%q) = %v, want %v", c.version, got, c.want) + } + } +} + +func TestCompareVersions(t *testing.T) { + cases := []struct { + a, b string + want int + }{ + {"1.24.0", "1.24.0", 0}, + {"1.23.0", "1.24.0", -1}, + {"1.24.0", "1.23.0", 1}, + {"1.24", "1.24.0", 0}, + {"1.24.1", "1.24", 1}, + {"2.0.0", "1.99.99", 1}, + } + + for _, c := range cases { + if got := compareVersions(c.a, c.b); got != c.want { + t.Errorf("compareVersions(%q, %q) = %d, want %d", c.a, c.b, got, c.want) + } + } +}