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
+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)
}
}
}