From 06102ad706fc9c429acb2b615c09919d9eb0837d Mon Sep 17 00:00:00 2001 From: Owen Date: Tue, 1 Sep 2026 12:12:19 -0400 Subject: [PATCH] Add connection management methods and tests for websocket client Fixes #139 --- websocket/client.go | 145 +++++++++++++++++++++++++-------------- websocket/client_test.go | 113 ++++++++++++++++++++++++++++++ 2 files changed, 205 insertions(+), 53 deletions(-) create mode 100644 websocket/client_test.go diff --git a/websocket/client.go b/websocket/client.go index feb795c..9e88981 100644 --- a/websocket/client.go +++ b/websocket/client.go @@ -82,6 +82,7 @@ type Config struct { type Client struct { config *Config conn *websocket.Conn + connMux sync.Mutex // protects conn: reads, writes, and the reconnect compare-and-clear baseURL string handlers map[string]MessageHandler done chan struct{} @@ -101,17 +102,17 @@ type Client struct { configNeedsSave bool // Flag to track if config needs to be saved configVersion int // Latest config version received from server configVersionMux sync.RWMutex - token string // Cached authentication token - exitNodes []ExitNode // Cached exit nodes from token response - tokenMux sync.RWMutex // Protects token and exitNodes - 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 - processingWg sync.WaitGroup // WaitGroup to wait for message processing to complete - getPingData func() map[string]any // Callback to get additional ping data - pingStarted bool // Flag to track if ping monitor has been started - pingStartedMux sync.Mutex // Protects pingStarted - pingDone chan struct{} // Channel to stop the ping monitor independently + token string // Cached authentication token + exitNodes []ExitNode // Cached exit nodes from token response + tokenMux sync.RWMutex // Protects token and exitNodes + 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 + processingWg sync.WaitGroup // WaitGroup to wait for message processing to complete + getPingData func() map[string]any // Callback to get additional ping data + pingStarted bool // Flag to track if ping monitor has been started + pingStartedMux sync.Mutex // Protects pingStarted + pingDone chan struct{} // Channel to stop the ping monitor independently } type ClientOption func(*Client) @@ -214,6 +215,34 @@ func (c *Client) GetConfig() *Config { return c.config } +// getConn returns the current connection, or nil if not connected. +func (c *Client) getConn() *websocket.Conn { + c.connMux.Lock() + defer c.connMux.Unlock() + return c.conn +} + +// setConn replaces the current connection. +func (c *Client) setConn(conn *websocket.Conn) { + c.connMux.Lock() + c.conn = conn + c.connMux.Unlock() +} + +// clearConnIfCurrent nils out c.conn only if it is still set to old, and +// reports whether it did so. Reconnects are triggered independently by the +// read pump (on a read error) and by sendPing (on a write error), and both +// can fire for the same dead connection at once. +func (c *Client) clearConnIfCurrent(old *websocket.Conn) bool { + c.connMux.Lock() + defer c.connMux.Unlock() + if c.conn != old { + return false + } + c.conn = nil + return true +} + // Connect establishes the WebSocket connection func (c *Client) Connect() error { if c.isDisconnected { @@ -242,14 +271,14 @@ func (c *Client) Close() error { c.setConnected(false) // Close the WebSocket connection gracefully - if c.conn != nil { + if conn := c.getConn(); conn != nil { // Send close message c.writeMux.Lock() - c.conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) c.writeMux.Unlock() // Close the connection - return c.conn.Close() + return conn.Close() } return nil @@ -270,12 +299,12 @@ func (c *Client) Disconnect() error { // Wait for any message currently being processed to complete c.processingWg.Wait() - if c.conn != nil { + if conn := c.getConn(); conn != nil { c.writeMux.Lock() - c.conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) c.writeMux.Unlock() - err := c.conn.Close() - c.conn = nil + err := conn.Close() + c.clearConnIfCurrent(conn) return err } return nil @@ -286,7 +315,8 @@ func (c *Client) SendMessage(messageType string, data interface{}) error { if c == nil { return fmt.Errorf("client is nil") } - if c.isDisconnected || c.conn == nil { + conn := c.getConn() + if c.isDisconnected || conn == nil { return fmt.Errorf("not connected") } @@ -299,10 +329,10 @@ func (c *Client) SendMessage(messageType string, data interface{}) error { c.writeMux.Lock() defer c.writeMux.Unlock() - if err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline)); err != nil { + if err := conn.SetWriteDeadline(time.Now().Add(writeDeadline)); err != nil { return err } - return c.conn.WriteJSON(msg) + return conn.WriteJSON(msg) } func (c *Client) SendMessageInterval(messageType string, data interface{}, interval time.Duration, maxAttempts int) (stop func(), update func(newData interface{})) { @@ -319,7 +349,7 @@ func (c *Client) SendMessageInterval(messageType string, data interface{}, inter count := 0 send := func() { - if c.isDisconnected || c.conn == nil { + if c.isDisconnected || c.getConn() == nil { return } err := c.SendMessage(messageType, currentData) @@ -423,10 +453,10 @@ func (c *Client) getToken() (string, []ExitNode, error) { } tokenData := map[string]interface{}{ - "olmId": c.config.ID, - "secret": c.config.Secret, + "olmId": c.config.ID, + "secret": c.config.Secret, "userToken": c.config.UserToken, - "orgId": c.config.OrgID, + "orgId": c.config.OrgID, } jsonData, err := json.Marshal(tokenData) @@ -617,7 +647,7 @@ func (c *Client) establishConnection() error { return fmt.Errorf("failed to connect to WebSocket: %w", err) } - c.conn = conn + c.setConn(conn) c.setConnected(true) // Arm a read deadline and refresh it whenever a pong arrives. Combined with @@ -626,17 +656,19 @@ func (c *Client) establishConnection() error { // on sleep/resume, or total packet loss) that a write-side check alone // misses: small periodic pings fit in the kernel send buffer and keep // "succeeding" even when nothing is actually reaching the peer. - _ = c.conn.SetReadDeadline(time.Now().Add(c.pongWait)) - c.conn.SetPongHandler(func(appData string) error { - _ = c.conn.SetReadDeadline(time.Now().Add(c.pongWait)) + _ = conn.SetReadDeadline(time.Now().Add(c.pongWait)) + conn.SetPongHandler(func(appData string) error { + _ = conn.SetReadDeadline(time.Now().Add(c.pongWait)) return nil }) // Note: ping monitor is NOT started here - it will be started when // StartPingMonitor() is called after registration completes - // Start the read pump with disconnect detection - go c.readPumpWithDisconnectDetection() + // Start the read pump with disconnect detection, bound to this specific + // connection instance (not the mutable c.conn field) so it can never be + // made to read a connection other than the one it was spawned for. + go c.readPumpWithDisconnectDetection(conn) if c.onConnect != nil { if err := c.onConnect(); err != nil { @@ -712,7 +744,8 @@ func (c *Client) setupPKCS12TLS() (*tls.Config, error) { // sendPing sends a single ping message func (c *Client) sendPing() { - if c.isDisconnected || c.conn == nil { + conn := c.getConn() + if c.isDisconnected || conn == nil { return } // Skip ping if a message is currently being processed @@ -747,16 +780,16 @@ func (c *Client) sendPing() { logger.Debug("websocket: Sending ping: %+v", pingMsg) c.writeMux.Lock() - err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline)) + err := conn.SetWriteDeadline(time.Now().Add(writeDeadline)) if err == nil { - err = c.conn.WriteJSON(pingMsg) + err = conn.WriteJSON(pingMsg) } if err == nil { // Protocol-level ping: a standards-compliant server replies with a // PONG, which refreshes the read deadline via SetPongHandler. This is // what actually detects a half-open connection where writes still // "succeed" (buffered by the kernel) but nothing is reaching the peer. - _ = c.conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(writeDeadline)) + _ = conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(writeDeadline)) } c.writeMux.Unlock() if err != nil { @@ -767,7 +800,7 @@ func (c *Client) sendPing() { return default: logger.Error("websocket: Ping failed: %v", err) - c.reconnect() + c.reconnect(conn) return } } @@ -815,15 +848,15 @@ func (c *Client) StartPingMonitor() { c.pingStartedMux.Lock() defer c.pingStartedMux.Unlock() - + if c.pingStarted { return } c.pingStarted = true - + // Create a new pingDone channel for this ping monitor instance c.pingDone = make(chan struct{}) - + // Send an initial ping immediately go func() { c.sendPing() @@ -835,11 +868,11 @@ func (c *Client) StartPingMonitor() { func (c *Client) stopPingMonitor() { c.pingStartedMux.Lock() defer c.pingStartedMux.Unlock() - + if !c.pingStarted { return } - + // Close the pingDone channel to stop the monitor close(c.pingDone) c.pingStarted = false @@ -860,19 +893,21 @@ func (c *Client) setConfigVersion(version int) { c.configVersion = version } -// readPumpWithDisconnectDetection reads messages and triggers reconnect on error -func (c *Client) readPumpWithDisconnectDetection() { +// readPumpWithDisconnectDetection reads messages and triggers reconnect on +// error. conn is the specific connection instance this pump was spawned for +// (captured at spawn time, not re-read from c.conn) so that a concurrent +// reconnect swapping or clearing c.conn can never cause this loop to read +// through a nil or unrelated connection. +func (c *Client) readPumpWithDisconnectDetection(conn *websocket.Conn) { defer func() { - if c.conn != nil { - c.conn.Close() - } + conn.Close() // Only attempt reconnect if we're not shutting down select { case <-c.done: // Shutting down, don't reconnect return default: - c.reconnect() + c.reconnect(conn) } }() @@ -881,13 +916,13 @@ func (c *Client) readPumpWithDisconnectDetection() { case <-c.done: return default: - messageType, p, err := c.conn.ReadMessage() + messageType, p, err := conn.ReadMessage() if err == nil { // Any inbound traffic means the peer is alive — extend the // read deadline (also covers servers that answer the // app-level "olm/ping" with a message rather than a // protocol pong). - _ = c.conn.SetReadDeadline(time.Now().Add(c.pongWait)) + _ = conn.SetReadDeadline(time.Now().Add(c.pongWait)) } if err != nil { // Check if we're shutting down or explicitly disconnected before logging error @@ -961,11 +996,15 @@ func (c *Client) readPumpWithDisconnectDetection() { } } -func (c *Client) reconnect() { +// reconnect tears down old and starts a fresh connectWithRetry loop +func (c *Client) reconnect(old *websocket.Conn) { + if !c.clearConnIfCurrent(old) { + return + } + c.setConnected(false) - if c.conn != nil { - c.conn.Close() - c.conn = nil + if old != nil { + old.Close() } // Don't reconnect if explicitly disconnected diff --git a/websocket/client_test.go b/websocket/client_test.go new file mode 100644 index 0000000..121dfbc --- /dev/null +++ b/websocket/client_test.go @@ -0,0 +1,113 @@ +package websocket + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gorilla/websocket" +) + +// newTestClient returns a bare Client suitable for exercising conn +// bookkeeping directly, without dialing a real server. isDisconnected is set +// so reconnect() performs its compare-and-clear but never spawns a +// connectWithRetry goroutine. +func newTestClient() *Client { + return &Client{ + done: make(chan struct{}), + isDisconnected: true, + } +} + +// dialTestConn spins up a throwaway websocket server and returns a live +// client-side *websocket.Conn, for tests that need a real connection (i.e. +// one that survives a Close() call without panicking). +func dialTestConn(t *testing.T) *websocket.Conn { + t.Helper() + upgrader := websocket.Upgrader{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + select {} + })) + t.Cleanup(srv.Close) + + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("failed to dial test server: %v", err) + } + t.Cleanup(func() { conn.Close() }) + return conn +} + +func TestClearConnIfCurrentClearsMatchingConnection(t *testing.T) { + c := newTestClient() + a := dialTestConn(t) + c.setConn(a) + + if !c.clearConnIfCurrent(a) { + t.Fatal("expected clearConnIfCurrent(a) to succeed when a is the active connection") + } + if got := c.getConn(); got != nil { + t.Fatalf("expected conn to be nil after clearing, got %v", got) + } +} + +func TestClearConnIfCurrentIgnoresStaleConnection(t *testing.T) { + c := newTestClient() + a := dialTestConn(t) + b := dialTestConn(t) + c.setConn(a) + + // Simulate a concurrent winner already having replaced a with b. + c.setConn(b) + + if c.clearConnIfCurrent(a) { + t.Fatal("expected clearConnIfCurrent(a) to fail once b has replaced a as the active connection") + } + if got := c.getConn(); got != b { + t.Fatalf("stale clearConnIfCurrent(a) must not clobber the newer connection: got %v, want %v", got, b) + } +} + +// TestReconnectIgnoresStaleConnection is the direct regression test for +// fosrl/olm#139: the read pump and sendPing can both react to the same dead +// connection and call reconnect() independently. Once one of them has +// already replaced c.conn with a newer connection, the other's reconnect +// call (carrying the old, now-stale connection) must be a no-op rather than +// tearing down or nulling out the newer connection out from under its own +// read pump. +func TestReconnectIgnoresStaleConnection(t *testing.T) { + c := newTestClient() + a := dialTestConn(t) + b := dialTestConn(t) + c.setConn(a) + + // A winning concurrent reconnect already replaced a with b. + c.setConn(b) + + // The stale report for `a` must not touch b, and must not attempt to + // close `a` a second time or otherwise panic. + c.reconnect(a) + + if got := c.getConn(); got != b { + t.Fatalf("stale reconnect(a) must not clobber the newer connection: got %v, want %v", got, b) + } +} + +func TestReconnectClearsCurrentConnection(t *testing.T) { + c := newTestClient() + a := dialTestConn(t) + c.setConn(a) + + c.reconnect(a) + + if got := c.getConn(); got != nil { + t.Fatalf("expected conn to be cleared after reconnect(a) when a was still current, got %v", got) + } +}