mirror of
https://github.com/fosrl/olm.git
synced 2026-10-01 18:29:08 +02:00
Support batched relay,unrelay,local,unlocal messages
This commit is contained in:
+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