package ws import ( "log" "net/http" "time" "github.com/gorilla/websocket" ) const ( writeWait = 10 * time.Second pongWait = 60 * time.Second pingPeriod = (pongWait * 9) / 10 maxMsgSize = 1024 * 1024 // 1MB ) var upgrader = websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, CheckOrigin: func(r *http.Request) bool { return true // Allow all origins in Phase 1 }, } // HandleDaemonWS handles WebSocket connections from daemons. func HandleDaemonWS(hub *Hub, apiKey string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { // Authenticate key := r.Header.Get("X-API-Key") if key != apiKey { http.Error(w, "unauthorized", http.StatusUnauthorized) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { log.Printf("WebSocket upgrade failed: %v", err) return } machineID := r.Header.Get("X-Machine-ID") if machineID == "" { machineID = "default" } client := &Client{ Type: ClientDaemon, ID: machineID, Conn: conn, Send: make(chan []byte, 256), Hub: hub, } hub.Register(client) go clientWritePump(client) go clientReadPump(client) } } // HandleDashboardWS handles WebSocket connections from dashboards. func HandleDashboardWS(hub *Hub, token string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { // Authenticate via query param or header t := r.URL.Query().Get("token") if t == "" { t = r.Header.Get("Authorization") } if t != token && t != "Bearer "+token { http.Error(w, "unauthorized", http.StatusUnauthorized) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { log.Printf("WebSocket upgrade failed: %v", err) return } client := &Client{ Type: ClientDashboard, Conn: conn, Send: make(chan []byte, 256), Hub: hub, } hub.Register(client) go clientWritePump(client) go clientReadPump(client) } } func clientReadPump(client *Client) { defer func() { client.Hub.Unregister(client) client.Conn.Close() }() client.Conn.SetReadLimit(maxMsgSize) client.Conn.SetReadDeadline(time.Now().Add(pongWait)) client.Conn.SetPongHandler(func(string) error { client.Conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) for { _, message, err := client.Conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) { log.Printf("WebSocket read error: %v", err) } return } client.Hub.HandleMessage(client, message) } } func clientWritePump(client *Client) { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() client.Conn.Close() }() for { select { case message, ok := <-client.Send: client.Conn.SetWriteDeadline(time.Now().Add(writeWait)) if !ok { client.Conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if err := client.Conn.WriteMessage(websocket.TextMessage, message); err != nil { return } case <-ticker.C: client.Conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := client.Conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } }