package relay import ( "encoding/json" "log" "net" "net/http" "strings" "time" "github.com/duffy/usb-server/internal/protocol" "github.com/gorilla/websocket" ) const ( // readTimeout is how long a client may stay silent before we drop it. // It must exceed pingInterval so that keepalive pongs refresh it. readTimeout = 60 * time.Second // pingInterval is how often the relay pings each client. pingInterval = 20 * time.Second // writeTimeout bounds a single frame write. Without it, a peer that has // stopped reading would pin its write pump forever. writeTimeout = 20 * time.Second // maxMessageSize caps an inbound frame. Tunnel frames are at most 64 KB // of USB payload plus the tunnel header; 1 MB leaves ample headroom. maxMessageSize = 1024 * 1024 ) var upgrader = websocket.Upgrader{ ReadBufferSize: 64 * 1024, WriteBufferSize: 64 * 1024, CheckOrigin: func(r *http.Request) bool { return true // relay accepts all origins }, } // Server is the WebSocket relay server type Server struct { hub *Hub addr string diag *diagStore } // NewServer creates a new relay server func NewServer(addr string) *Server { return &Server{ hub: NewHub(), addr: addr, diag: newDiagStore(), } } // Run starts the relay server func (s *Server) Run() error { mux := http.NewServeMux() mux.HandleFunc("/ws", s.handleWebSocket) mux.HandleFunc("/health", s.handleHealth) mux.HandleFunc("/diag/", s.handleDiag) // Timeouts bound how long a stuck client can hold a connection. The // WebSocket route needs no write timeout — those connections are // long-lived by design — so it is left to the per-message deadlines the // write pump sets. server := &http.Server{ Addr: s.addr, Handler: mux, ReadHeaderTimeout: 15 * time.Second, IdleTimeout: 120 * time.Second, } log.Printf("[relay] starting on %s", s.addr) return server.ListenAndServe() } func (s *Server) handleHealth(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) json.NewEncoder(w).Encode(map[string]string{"status": "ok"}) } func (s *Server) handleWebSocket(w http.ResponseWriter, r *http.Request) { conn, err := upgrader.Upgrade(w, r, nil) if err != nil { log.Printf("[relay] upgrade error: %v", err) return } defer conn.Close() // Set read limits and deadlines conn.SetReadLimit(maxMessageSize) conn.SetReadDeadline(time.Now().Add(readTimeout)) conn.SetPongHandler(func(string) error { conn.SetReadDeadline(time.Now().Add(readTimeout)) return nil }) // Wait for registration message _, msgData, err := conn.ReadMessage() if err != nil { log.Printf("[relay] read error during registration: %v", err) return } var reg protocol.Register if err := json.Unmarshal(msgData, ®); err != nil || reg.Type != protocol.MsgRegister { log.Printf("[relay] invalid registration message") conn.WriteJSON(&protocol.ErrorMsg{Type: protocol.MsgError, Message: "invalid registration"}) return } if reg.Hash == "" || reg.ClientID == "" || !protocol.ValidMode(reg.Mode) { conn.WriteJSON(&protocol.ErrorMsg{Type: protocol.MsgError, Message: "missing or invalid registration fields"}) return } client := newClient(reg.ClientID, reg.Hash, reg.Mode, reg.Name, conn) client.DirectPort = reg.DirectPort client.PublicIP = clientIP(r) s.hub.Register(client) defer s.hub.Unregister(client) // The write pump owns the socket's write side: every frame for this // client, plus keepalive pings, goes through it. Nothing else may write, // which is what keeps one unresponsive peer from blocking the hub. go s.writePump(client) // Read loop for { msgType, data, err := conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) { log.Printf("[relay] read error from %s: %v", protocol.ShortID(client.ID), err) } break } conn.SetReadDeadline(time.Now().Add(readTimeout)) switch msgType { case websocket.TextMessage: s.hub.HandleTextMessage(client, data) case websocket.BinaryMessage: s.hub.HandleBinaryMessage(client, data) } } client.kill() } // clientIP determines the address a client connects from, which is passed on // to its peers so they can reach it directly. // // X-Forwarded-For is honoured because relays are commonly deployed behind a // reverse proxy, where RemoteAddr would otherwise be the proxy itself. Only // the first entry is used: later ones are supplied by upstream hops and are // not trustworthy. A wrong value here costs a failed direct attempt and a // fallback to relaying, never a security property — the peer still has to // prove group membership in the handshake. func clientIP(r *http.Request) string { if fwd := r.Header.Get("X-Forwarded-For"); fwd != "" { first := strings.TrimSpace(strings.Split(fwd, ",")[0]) if ip := net.ParseIP(first); ip != nil { return ip.String() } } if real := strings.TrimSpace(r.Header.Get("X-Real-IP")); real != "" { if ip := net.ParseIP(real); ip != nil { return ip.String() } } host, _, err := net.SplitHostPort(r.RemoteAddr) if err != nil { return "" } if ip := net.ParseIP(host); ip != nil { return ip.String() } return "" } // writePump serialises all writes to one client's socket. func (s *Server) writePump(client *Client) { ticker := time.NewTicker(pingInterval) defer ticker.Stop() defer client.Conn.Close() // unblocks the read loop when we give up for { select { case msg := <-client.Send: client.Conn.SetWriteDeadline(time.Now().Add(writeTimeout)) if err := client.Conn.WriteMessage(msg.typ, msg.data); err != nil { log.Printf("[relay] write error to %s: %v", protocol.ShortID(client.ID), err) client.kill() return } case <-ticker.C: client.Conn.SetWriteDeadline(time.Now().Add(writeTimeout)) if err := client.Conn.WriteMessage(websocket.PingMessage, nil); err != nil { client.kill() return } case <-client.dead: return } } }