package wsclient import ( "encoding/json" "fmt" "log" "net/http" "net/url" "strings" "time" "github.com/gorilla/websocket" ) // Client keeps a streamd /ws/agent connection alive and sends heartbeats // so claims auto-renew; on disconnect streamd releases this agent's claims. type Client struct { BaseURL string Token string AgentID string Log *log.Logger stop chan struct{} } func New(baseURL, token, agentID string, lg *log.Logger) *Client { if lg == nil { lg = log.Default() } return &Client{ BaseURL: strings.TrimRight(baseURL, "/"), Token: token, AgentID: agentID, Log: lg, stop: make(chan struct{}), } } func (c *Client) Stop() { close(c.stop) } // Loop reconnects forever until Stop. func (c *Client) Loop() { backoff := time.Second for { select { case <-c.stop: return default: } if err := c.session(); err != nil { c.Log.Printf("ws: session ended: %v — retry in %s", err, backoff) } select { case <-c.stop: return case <-time.After(backoff): } if backoff < 30*time.Second { backoff *= 2 } } } func (c *Client) session() error { u, err := url.Parse(c.BaseURL) if err != nil { return err } switch u.Scheme { case "https": u.Scheme = "wss" default: u.Scheme = "ws" } u.Path = "/ws/agent" q := u.Query() q.Set("agent_id", c.AgentID) if c.Token != "" { q.Set("token", c.Token) } u.RawQuery = q.Encode() hdr := http.Header{} if c.Token != "" { hdr.Set("Authorization", "Bearer "+c.Token) } dialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second} conn, _, err := dialer.Dial(u.String(), hdr) if err != nil { return fmt.Errorf("dial: %w", err) } defer conn.Close() c.Log.Printf("ws: connected to %s as %s", u.Host, c.AgentID) _ = conn.SetReadDeadline(time.Now().Add(60 * time.Second)) conn.SetPongHandler(func(string) error { _ = conn.SetReadDeadline(time.Now().Add(60 * time.Second)) return nil }) done := make(chan struct{}) go func() { defer close(done) for { _, data, err := conn.ReadMessage() if err != nil { return } _ = conn.SetReadDeadline(time.Now().Add(60 * time.Second)) var msg map[string]any if json.Unmarshal(data, &msg) == nil { if t, _ := msg["type"].(string); t == "heartbeat_ack" { if ok, _ := msg["ok"].(bool); !ok { c.Log.Printf("ws: heartbeat nack: %v", msg["error"]) } } } } }() ticker := time.NewTicker(20 * time.Second) defer ticker.Stop() // Immediate heartbeat so existing claims renew right away on reconnect. if err := c.writeHeartbeat(conn); err != nil { return err } for { select { case <-c.stop: _ = conn.WriteJSON(map[string]string{"type": "release_all"}) return nil case <-done: return fmt.Errorf("read closed") case <-ticker.C: if err := c.writeHeartbeat(conn); err != nil { return err } } } } func (c *Client) writeHeartbeat(conn *websocket.Conn) error { _ = conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) return conn.WriteJSON(map[string]string{ "type": "heartbeat", "agent_id": c.AgentID, }) }