package relay import ( "encoding/json" "log" "net/http" "sync" "time" "github.com/gorilla/websocket" ) const ( reconnectDelay = 5 * time.Second writeWait = 10 * time.Second pongWait = 60 * time.Second pingPeriod = (pongWait * 9) / 10 ) // MessageHandler handles incoming messages from the server. type MessageHandler func(msgType string, raw json.RawMessage) // Client maintains a WebSocket connection to the server. type Client struct { url string apiKey string handler MessageHandler conn *websocket.Conn mu sync.Mutex sendCh chan []byte done chan struct{} } // NewClient creates a new WebSocket relay client. func NewClient(url, apiKey string, handler MessageHandler) *Client { return &Client{ url: url, apiKey: apiKey, handler: handler, sendCh: make(chan []byte, 256), done: make(chan struct{}), } } // Start connects to the server and begins read/write loops. func (c *Client) Start() { go c.connectLoop() } // Stop closes the connection. func (c *Client) Stop() { close(c.done) c.mu.Lock() if c.conn != nil { c.conn.Close() } c.mu.Unlock() } // Send sends a message to the server. func (c *Client) Send(msg any) error { data, err := json.Marshal(msg) if err != nil { return err } select { case c.sendCh <- data: return nil case <-c.done: return nil default: log.Printf("relay send buffer full, dropping message") return nil } } // Connected returns true if the client has an active connection. func (c *Client) Connected() bool { c.mu.Lock() defer c.mu.Unlock() return c.conn != nil } func (c *Client) connectLoop() { for { select { case <-c.done: return default: } if err := c.connect(); err != nil { log.Printf("WebSocket connection failed: %v", err) select { case <-time.After(reconnectDelay): case <-c.done: return } continue } // Run read/write loops doneCh := make(chan struct{}) go c.readLoop(doneCh) c.writeLoop(doneCh) c.mu.Lock() if c.conn != nil { c.conn.Close() c.conn = nil } c.mu.Unlock() log.Printf("WebSocket disconnected, reconnecting...") select { case <-time.After(reconnectDelay): case <-c.done: return } } } func (c *Client) connect() error { header := http.Header{} header.Set("X-API-Key", c.apiKey) conn, _, err := websocket.DefaultDialer.Dial(c.url, header) if err != nil { return err } c.mu.Lock() c.conn = conn c.mu.Unlock() log.Printf("Connected to server: %s", c.url) return nil } func (c *Client) readLoop(doneCh chan struct{}) { defer close(doneCh) c.mu.Lock() conn := c.conn c.mu.Unlock() conn.SetReadDeadline(time.Now().Add(pongWait)) conn.SetPongHandler(func(string) error { conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) for { _, message, err := conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) { log.Printf("WebSocket read error: %v", err) } return } var base Message if err := json.Unmarshal(message, &base); err != nil { log.Printf("failed to parse message: %v", err) continue } if c.handler != nil { c.handler(base.Type, json.RawMessage(message)) } } } func (c *Client) writeLoop(doneCh chan struct{}) { ticker := time.NewTicker(pingPeriod) defer ticker.Stop() for { select { case msg := <-c.sendCh: c.mu.Lock() conn := c.conn c.mu.Unlock() if conn == nil { continue } conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := conn.WriteMessage(websocket.TextMessage, msg); err != nil { log.Printf("WebSocket write error: %v", err) return } case <-ticker.C: c.mu.Lock() conn := c.conn c.mu.Unlock() if conn == nil { return } conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } case <-doneCh: return case <-c.done: return } } }