package handlers import ( "encoding/json" "log/slog" "time" "github.com/gorilla/websocket" "golang.org/x/time/rate" ) const ( writeWait = 10 * time.Second pongWait = 60 * time.Second pingPeriod = (pongWait * 9) / 10 maxMessageSize = 4096 visitorUpdateLimit = 10 visitorUpdateBurst = 1 ) type wsMessage struct { Type string `json:"type"` Data json.RawMessage `json:"data"` } type Client struct { hub *WSHub conn *websocket.Conn send chan []byte userID string lat float64 lng float64 limiter *rate.Limiter } func newClient(hub *WSHub, conn *websocket.Conn) *Client { return &Client{ hub: hub, conn: conn, send: make(chan []byte, 256), limiter: rate.NewLimiter(visitorUpdateLimit, visitorUpdateBurst), } } func (c *Client) readPump() { defer func() { c.hub.unregister <- c c.conn.Close() }() c.conn.SetReadLimit(maxMessageSize) c.conn.SetReadDeadline(time.Now().Add(pongWait)) c.conn.SetPongHandler(func(string) error { c.conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) for { var msg wsMessage if err := c.conn.ReadJSON(&msg); err != nil { break } switch msg.Type { case "visitor_update": if !c.limiter.Allow() { slog.Warn("ws visitor_update rate limit exceeded", "user_id", c.userID) continue } var dot VisitorDot if err := json.Unmarshal(msg.Data, &dot); err != nil { continue } c.userID = dot.UserID c.lat = dot.Lat c.lng = dot.Lng c.hub.broadcastVisitors() } } } func (c *Client) writePump() { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() c.conn.Close() }() for { select { case message, ok := <-c.send: c.conn.SetWriteDeadline(time.Now().Add(writeWait)) if !ok { c.conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if err := c.conn.WriteMessage(websocket.TextMessage, message); err != nil { return } case <-ticker.C: c.conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := c.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } }