package ws import ( "encoding/json" "log" "net/http" "strings" "sync" "time" "github.com/gorilla/websocket" "drmdecryption/apps/streamd/internal/db" ) const ( agentHeartbeatEvery = 20 * time.Second agentStaleAfter = 45 * time.Second claimRenewTTL = 90 * time.Second writeWait = 10 * time.Second pongWait = 60 * time.Second ) var upgrader = websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, } // Hub tracks agent WebSocket sessions and UI listeners. type Hub struct { Store *db.Store Token string Log *log.Logger mu sync.Mutex agents map[string]*agentConn // agent_id -> conn uis map[*websocket.Conn]struct{} } type agentConn struct { conn *websocket.Conn agentID string lastSeen time.Time send chan []byte } type inbound struct { Type string `json:"type"` AgentID string `json:"agent_id,omitempty"` } type outbound struct { Type string `json:"type"` AgentID string `json:"agent_id,omitempty"` OK bool `json:"ok,omitempty"` Error string `json:"error,omitempty"` Agents []any `json:"agents,omitempty"` Message string `json:"message,omitempty"` } func NewHub(store *db.Store, token string, lg *log.Logger) *Hub { if lg == nil { lg = log.Default() } return &Hub{ Store: store, Token: token, Log: lg, agents: map[string]*agentConn{}, uis: map[*websocket.Conn]struct{}{}, } } func (h *Hub) Mount(mux *http.ServeMux) { mux.HandleFunc("/ws/agent", h.handleAgent) mux.HandleFunc("/ws/ui", h.handleUI) mux.HandleFunc("/api/agents", h.handleAgentsAPI) } // RunExpireLoop periodically drops expired claims and dead agent sessions. func (h *Hub) RunExpireLoop(stop <-chan struct{}) { t := time.NewTicker(10 * time.Second) defer t.Stop() for { select { case <-stop: return case <-t.C: h.tick() } } } func (h *Hub) tick() { n, err := h.Store.ExpireClaims() if err != nil { h.Log.Printf("ws: expire claims: %v", err) } else if n > 0 { h.Log.Printf("ws: expired %d stale claim(s)", n) h.broadcastUI(outbound{Type: "claims_changed", Message: "expired"}) } now := time.Now().UTC() var stale []string h.mu.Lock() for id, a := range h.agents { if now.Sub(a.lastSeen) > agentStaleAfter { stale = append(stale, id) } } h.mu.Unlock() for _, id := range stale { h.Log.Printf("ws: agent %s heartbeat stale — releasing claims", id) h.dropAgent(id, true) } } func (h *Hub) authToken(r *http.Request) bool { if h.Token == "" { return true } tok := r.URL.Query().Get("token") if tok == "" { auth := r.Header.Get("Authorization") tok = strings.TrimPrefix(auth, "Bearer ") } if tok == "" { tok = r.Header.Get("X-Streamd-Token") } return tok == h.Token } func (h *Hub) handleAgentsAPI(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) return } writeJSON(w, map[string]any{"agents": h.snapshotAgents()}) } func (h *Hub) handleUI(w http.ResponseWriter, r *http.Request) { conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } h.mu.Lock() h.uis[conn] = struct{}{} h.mu.Unlock() defer func() { h.mu.Lock() delete(h.uis, conn) h.mu.Unlock() _ = conn.Close() }() _ = conn.SetReadDeadline(time.Now().Add(pongWait)) conn.SetPongHandler(func(string) error { _ = conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) // Push initial agent list. h.sendJSON(conn, outbound{Type: "agents", Agents: h.snapshotAgents()}) for { if _, _, err := conn.ReadMessage(); err != nil { return } } } func (h *Hub) handleAgent(w http.ResponseWriter, r *http.Request) { if !h.authToken(r) { http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized) return } agentID := strings.TrimSpace(r.URL.Query().Get("agent_id")) if agentID == "" { http.Error(w, `{"error":"agent_id required"}`, http.StatusBadRequest) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } ac := &agentConn{ conn: conn, agentID: agentID, lastSeen: time.Now().UTC(), send: make(chan []byte, 8), } h.mu.Lock() if old, ok := h.agents[agentID]; ok { close(old.send) _ = old.conn.Close() } h.agents[agentID] = ac h.mu.Unlock() h.Log.Printf("ws: agent connected id=%s", agentID) h.broadcastUI(outbound{Type: "agent_online", AgentID: agentID, Agents: h.snapshotAgents()}) go h.agentWriter(ac) defer h.dropAgent(agentID, true) _ = conn.SetReadDeadline(time.Now().Add(pongWait)) conn.SetPongHandler(func(string) error { _ = conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) // Greeting + first renew. h.renew(agentID) h.queue(ac, outbound{Type: "hello", AgentID: agentID, OK: true}) for { _, data, err := conn.ReadMessage() if err != nil { return } _ = conn.SetReadDeadline(time.Now().Add(pongWait)) var msg inbound if err := json.Unmarshal(data, &msg); err != nil { continue } switch strings.ToLower(msg.Type) { case "heartbeat", "ping", "hello": h.mu.Lock() if cur, ok := h.agents[agentID]; ok { cur.lastSeen = time.Now().UTC() } h.mu.Unlock() n, err := h.renew(agentID) if err != nil { h.queue(ac, outbound{Type: "heartbeat_ack", OK: false, Error: err.Error()}) continue } h.queue(ac, outbound{Type: "heartbeat_ack", OK: true, Message: "renewed", AgentID: agentID}) if n > 0 { h.broadcastUI(outbound{Type: "claims_changed", AgentID: agentID, Message: "renewed"}) } case "release_all": _, _ = h.Store.ReleaseAgentClaims(agentID) h.broadcastUI(outbound{Type: "claims_changed", AgentID: agentID, Message: "released"}) h.queue(ac, outbound{Type: "release_ack", OK: true}) } } } func (h *Hub) renew(agentID string) (int64, error) { return h.Store.RenewAgentClaims(agentID, claimRenewTTL) } func (h *Hub) dropAgent(agentID string, releaseClaims bool) { h.mu.Lock() ac, ok := h.agents[agentID] if ok { delete(h.agents, agentID) } h.mu.Unlock() if !ok { return } close(ac.send) _ = ac.conn.Close() if releaseClaims { n, err := h.Store.ReleaseAgentClaims(agentID) if err != nil { h.Log.Printf("ws: release claims for %s: %v", agentID, err) } else if n > 0 { h.Log.Printf("ws: released %d claim(s) for disconnected agent %s", n, agentID) } } h.Log.Printf("ws: agent disconnected id=%s", agentID) h.broadcastUI(outbound{Type: "agent_offline", AgentID: agentID, Agents: h.snapshotAgents()}) } func (h *Hub) agentWriter(ac *agentConn) { ticker := time.NewTicker(agentHeartbeatEvery) defer ticker.Stop() for { select { case msg, ok := <-ac.send: if !ok { return } _ = ac.conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := ac.conn.WriteMessage(websocket.TextMessage, msg); err != nil { return } case <-ticker.C: _ = ac.conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := ac.conn.WriteControl(websocket.PingMessage, []byte("ping"), time.Now().Add(writeWait)); err != nil { return } } } } func (h *Hub) queue(ac *agentConn, msg outbound) { b, err := json.Marshal(msg) if err != nil { return } select { case ac.send <- b: default: } } func (h *Hub) sendJSON(conn *websocket.Conn, msg outbound) { b, _ := json.Marshal(msg) _ = conn.SetWriteDeadline(time.Now().Add(writeWait)) _ = conn.WriteMessage(websocket.TextMessage, b) } func (h *Hub) broadcastUI(msg outbound) { b, err := json.Marshal(msg) if err != nil { return } h.mu.Lock() defer h.mu.Unlock() for conn := range h.uis { _ = conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := conn.WriteMessage(websocket.TextMessage, b); err != nil { _ = conn.Close() delete(h.uis, conn) } } } func (h *Hub) snapshotAgents() []any { h.mu.Lock() defer h.mu.Unlock() out := make([]any, 0, len(h.agents)) now := time.Now().UTC() for id, a := range h.agents { out = append(out, map[string]any{ "agent_id": id, "last_seen_s": now.Sub(a.lastSeen).Seconds(), "connected": true, }) } return out } func writeJSON(w http.ResponseWriter, v any) { w.Header().Set("Content-Type", "application/json") enc := json.NewEncoder(w) enc.SetIndent("", " ") _ = enc.Encode(v) }