Support batched relay,unrelay,local,unlocal messages

This commit is contained in:
Owen
2026-09-24 11:55:53 -04:00
parent c543f39986
commit e04e7e0bcf
7 changed files with 493 additions and 142 deletions
+85 -27
View File
@@ -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)
}
}
}
+212
View File
@@ -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)
}
}
+43 -97
View File
@@ -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()
+15 -4
View File
@@ -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
+16 -14
View File
@@ -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 {
+75
View File
@@ -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
}
+47
View File
@@ -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)
}
}
}