package httpserver import ( "net/http" "time" "github.com/gorilla/websocket" "github.com/tim/md-hub-secure/apps/server/internal/realtime" ) const ( wsWriteWait = 10 * time.Second wsPongWait = 60 * time.Second wsPingPeriod = 45 * time.Second ) var wsUpgrader = websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, CheckOrigin: func(r *http.Request) bool { origin := r.Header.Get("Origin") if origin == "" { return true } return origin == "http://"+r.Host || origin == "https://"+r.Host }, } func (s *Server) handleWebSocket(w http.ResponseWriter, r *http.Request) { conn, err := wsUpgrader.Upgrade(w, r, nil) if err != nil { s.logger.Warn("upgrade websocket", "error", err) return } client := s.hub.NewClient(16) s.hub.Register(client) go s.readWebSocket(conn, client) s.writeWebSocket(conn, client) } func (s *Server) readWebSocket(conn *websocket.Conn, client *realtime.Client) { defer func() { s.hub.Unregister(client) _ = conn.Close() }() conn.SetReadLimit(1024) _ = conn.SetReadDeadline(time.Now().Add(wsPongWait)) conn.SetPongHandler(func(string) error { return conn.SetReadDeadline(time.Now().Add(wsPongWait)) }) for { if _, _, err := conn.NextReader(); err != nil { return } } } func (s *Server) writeWebSocket(conn *websocket.Conn, client *realtime.Client) { ticker := time.NewTicker(wsPingPeriod) defer func() { ticker.Stop() s.hub.Unregister(client) _ = conn.Close() }() for { select { case message, ok := <-client.Send(): _ = conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) if !ok { _ = conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if err := conn.WriteMessage(websocket.TextMessage, message); err != nil { return } case <-ticker.C: _ = conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } }