package websocket import ( "net/http" "time" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" "go.uber.org/zap" ) var upgrader = websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, CheckOrigin: func(r *http.Request) bool { return true }, } type Handler struct { hub *Hub log *zap.Logger } func NewHandler(hub *Hub, log *zap.Logger) *Handler { return &Handler{hub: hub, log: log} } func (h *Handler) ServeWS(c *gin.Context) { userID := c.GetUint("user_id") tenantID := c.GetString("tenant_id") if userID == 0 { c.AbortWithStatusJSON(401, gin.H{"code": -1, "message": "unauthorized"}) return } conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { h.log.Error("websocket upgrade", zap.Error(err)) return } client := NewClient(h.hub, userID, tenantID) h.hub.Register(client) go client.writePump(conn) go client.readPump(conn, h.hub) } func (c *Client) readPump(conn *websocket.Conn, hub *Hub) { defer func() { hub.Unregister(c) conn.Close() close(c.done) }() conn.SetReadLimit(maxMessageSize) conn.SetReadDeadline(deadline(pongWait)) conn.SetPongHandler(func(string) error { conn.SetReadDeadline(deadline(pongWait)) return nil }) for { _, _, err := conn.ReadMessage() if err != nil { break } } } func (c *Client) writePump(conn *websocket.Conn) { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() conn.Close() }() for { select { case msg, ok := <-c.send: conn.SetWriteDeadline(deadline(writeWait)) if !ok { conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if err := conn.WriteMessage(websocket.TextMessage, msg); err != nil { return } case <-ticker.C: conn.SetWriteDeadline(deadline(writeWait)) if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } } func deadline(d time.Duration) time.Time { return time.Now().Add(d) }