353 lines
7.3 KiB
Go
353 lines
7.3 KiB
Go
package fleet
|
|
|
|
import (
|
|
"crypto/subtle"
|
|
"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 {
|
|
return subtle.ConstantTimeCompare([]byte(a), []byte(b)) == 1
|
|
}
|
|
|
|
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
|
|
}
|