51 lines
1.0 KiB
Go
51 lines
1.0 KiB
Go
|
|
package handler
|
||
|
|
|
||
|
|
import (
|
||
|
|
"log"
|
||
|
|
"net/http"
|
||
|
|
|
||
|
|
"github.com/gin-gonic/gin"
|
||
|
|
gorillaWs "github.com/gorilla/websocket"
|
||
|
|
ws "github.com/muyuqingfeng/iloom/shared/pkg/websocket"
|
||
|
|
)
|
||
|
|
|
||
|
|
var upgrader = gorillaWs.Upgrader{
|
||
|
|
ReadBufferSize: 1024,
|
||
|
|
WriteBufferSize: 1024,
|
||
|
|
CheckOrigin: func(r *http.Request) bool {
|
||
|
|
return true
|
||
|
|
},
|
||
|
|
}
|
||
|
|
|
||
|
|
func HandleWebSocket(hub *ws.Hub) gin.HandlerFunc {
|
||
|
|
return func(c *gin.Context) {
|
||
|
|
companyID := c.GetHeader("X-Company-ID")
|
||
|
|
userID := c.GetHeader("X-User-ID")
|
||
|
|
|
||
|
|
// Also try query params for WebSocket connections
|
||
|
|
if companyID == "" {
|
||
|
|
companyID = c.Query("company_id")
|
||
|
|
}
|
||
|
|
if userID == "" {
|
||
|
|
userID = c.Query("user_id")
|
||
|
|
}
|
||
|
|
|
||
|
|
if companyID == "" {
|
||
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": "missing company_id"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||
|
|
if err != nil {
|
||
|
|
log.Printf("websocket upgrade error: %v", err)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
client := ws.NewClient(hub, conn, companyID, userID)
|
||
|
|
client.Register()
|
||
|
|
|
||
|
|
go client.WritePump()
|
||
|
|
go client.ReadPump()
|
||
|
|
}
|
||
|
|
}
|