mirror of
https://github.com/fosrl/olm.git
synced 2026-10-02 10:49:10 +02:00
Support batched relay,unrelay,local,unlocal messages
This commit is contained in:
+85
-27
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user