package fleet import ( "encoding/json" "log" "net/http" "sync" "time" "forge-mesh/internal/api/types" "forge-mesh/internal/auth" "github.com/google/uuid" "github.com/gorilla/websocket" ) var upgrader = websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, } // Hub manages agent and deck WebSocket connections. type Hub struct { store *Store fleetSecret string tickets *auth.TicketStore mu sync.RWMutex agents map[string]*agentConn decks map[*deckConn]struct{} pendingCmds map[string][]types.FleetCommand } type agentConn struct { hostID string conn *websocket.Conn send chan []byte } type deckConn struct { conn *websocket.Conn send chan []byte } func NewHub(store *Store, fleetSecret string, tickets *auth.TicketStore) *Hub { return &Hub{ store: store, fleetSecret: fleetSecret, tickets: tickets, agents: make(map[string]*agentConn), decks: make(map[*deckConn]struct{}), pendingCmds: make(map[string][]types.FleetCommand), } } func (h *Hub) HandleAgentWS(w http.ResponseWriter, r *http.Request) { token := auth.ExtractBearer(r) if token == "" { token = r.URL.Query().Get("token") } if token == "" || !constantTimeEqual(token, h.fleetSecret) { http.Error(w, "unauthorized", http.StatusUnauthorized) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } ac := &agentConn{conn: conn, send: make(chan []byte, 16)} go h.writePump(ac, true) go h.readAgentPump(ac) } func (h *Hub) HandleDeckWS(w http.ResponseWriter, r *http.Request) { ticket := r.URL.Query().Get("ticket") if !h.tickets.Consume(ticket) { http.Error(w, "unauthorized", http.StatusUnauthorized) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } dc := &deckConn{conn: conn, send: make(chan []byte, 32)} h.mu.Lock() h.decks[dc] = struct{}{} h.mu.Unlock() go h.writePump(&agentConn{conn: conn, send: dc.send}, false) go h.readDeckPump(dc) } func (h *Hub) readAgentPump(ac *agentConn) { defer func() { h.unregisterAgent(ac) ac.conn.Close() }() ac.conn.SetReadLimit(1 << 20) _ = ac.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) ac.conn.SetPongHandler(func(string) error { return ac.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) }) for { _, data, err := ac.conn.ReadMessage() if err != nil { return } var msg types.WsMessage if err := json.Unmarshal(data, &msg); err != nil { continue } switch msg.Type { case "heartbeat": h.handleHeartbeat(ac, data) case "command_ack": h.broadcastToDecks(data) } } } func (h *Hub) handleHeartbeat(ac *agentConn, raw []byte) { var msg struct { types.WsMessage types.HeartbeatPayload } if err := json.Unmarshal(raw, &msg); err != nil { return } host, err := h.store.UpsertHeartbeat(types.HeartbeatPayload{ HostID: coalesce(msg.WsMessage.HostID, msg.HeartbeatPayload.HostID), Hostname: msg.Hostname, Arch: msg.Arch, Hashrate: msg.Hashrate, HashrateHps: coalesceFloat(msg.HashrateHps, msg.Hashrate), CurrentTier: msg.CurrentTier, TierType: msg.TierType, TierState: msg.TierState, Fingerprint: msg.Fingerprint, }) if err != nil { log.Printf("heartbeat store: %v", err) return } h.mu.Lock() if ac.hostID != "" && ac.hostID != host.ID { delete(h.agents, ac.hostID) } ac.hostID = host.ID h.agents[host.ID] = ac pending := h.pendingCmds[host.ID] delete(h.pendingCmds, host.ID) h.mu.Unlock() for _, cmd := range pending { h.sendCommand(ac, cmd) } card := ToFleetCard(host) update, _ := json.Marshal(map[string]any{ "type": "host_update", "host": card, "timestamp": time.Now().UTC().Format(time.RFC3339), }) h.broadcastToDecks(update) } func (h *Hub) readDeckPump(dc *deckConn) { defer func() { h.mu.Lock() delete(h.decks, dc) h.mu.Unlock() dc.conn.Close() }() dc.conn.SetReadLimit(1 << 18) for { if _, _, err := dc.conn.ReadMessage(); err != nil { return } } } func (h *Hub) writePump(ac *agentConn, ping bool) { ticker := time.NewTicker(30 * time.Second) defer func() { ticker.Stop() ac.conn.Close() }() for { select { case msg, ok := <-ac.send: _ = ac.conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if !ok { _ = ac.conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if err := ac.conn.WriteMessage(websocket.TextMessage, msg); err != nil { return } case <-ticker.C: if ping { _ = ac.conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := ac.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } } } func (h *Hub) unregisterAgent(ac *agentConn) { h.mu.Lock() defer h.mu.Unlock() if ac.hostID != "" { delete(h.agents, ac.hostID) _ = h.store.MarkOffline(ac.hostID) } } func (h *Hub) DispatchCommand(hostID, action string, args map[string]any) (*types.FleetCommand, error) { cmd := types.FleetCommand{ ID: uuid.NewString(), Action: action, Args: args, IssuedAt: time.Now().UTC(), } h.mu.Lock() ac, online := h.agents[hostID] if online { h.mu.Unlock() h.sendCommand(ac, cmd) return &cmd, nil } h.pendingCmds[hostID] = append(h.pendingCmds[hostID], cmd) h.mu.Unlock() return &cmd, nil } func (h *Hub) sendCommand(ac *agentConn, cmd types.FleetCommand) { payload, _ := json.Marshal(types.WsMessage{ Type: "command", HostID: ac.hostID, Command: &cmd, }) select { case ac.send <- payload: default: log.Printf("agent %s send buffer full", ac.hostID) } } // PushMiningProfile sends an updated mining profile to a connected agent. func (h *Hub) PushMiningProfile(hostID string, profile types.MiningProfile) bool { payload, err := json.Marshal(types.WsMessage{ Type: "mining_profile", HostID: hostID, Payload: map[string]any{"profile": profile}, }) if err != nil { return false } h.mu.RLock() ac, ok := h.agents[hostID] h.mu.RUnlock() if !ok { return false } select { case ac.send <- payload: return true default: return false } } func (h *Hub) PopPendingCommands(hostID string) []types.FleetCommand { h.mu.Lock() defer h.mu.Unlock() cmds := h.pendingCmds[hostID] delete(h.pendingCmds, hostID) return cmds } func (h *Hub) broadcastToDecks(data []byte) { h.mu.RLock() defer h.mu.RUnlock() for dc := range h.decks { select { case dc.send <- data: default: } } } func (h *Hub) HandleBeacon(store *Store) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var hb types.HeartbeatPayload if err := json.NewDecoder(r.Body).Decode(&hb); err != nil { http.Error(w, "bad request", http.StatusBadRequest) return } host, err := store.UpsertHeartbeat(hb) if err != nil { http.Error(w, "internal error", http.StatusInternalServerError) return } cmds := h.PopPendingCommands(host.ID) w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(types.BeaconResponse{OK: true, Commands: cmds}) } } func constantTimeEqual(a, b string) bool { if len(a) != len(b) { return false } var v byte for i := 0; i < len(a); i++ { v |= a[i] ^ b[i] } return v == 0 } func coalesce(values ...string) string { for _, v := range values { if v != "" { return v } } return "" } func coalesceFloat(values ...float64) float64 { for _, v := range values { if v > 0 { return v } } return 0 }