package handlers import ( "encoding/json" "log/slog" "net/http" "sync" "github.com/gorilla/websocket" "golang.org/x/time/rate" "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/middleware" ) type VisitorDot struct { UserID string `json:"user_id,omitempty"` Lat float64 `json:"lat"` Lng float64 `json:"lng"` } type visitorsPayload struct { Visitors []VisitorDot `json:"visitors"` } const maxWSConnections = 1000 type WSHub struct { mu sync.Mutex clients map[*Client]bool register chan *Client unregister chan *Client broadcast chan []byte allowedOrigins []string isProd bool connLimiter *rate.Limiter } func NewWSHub(allowedOrigins []string, isProd bool) *WSHub { return &WSHub{ clients: make(map[*Client]bool), register: make(chan *Client), unregister: make(chan *Client), broadcast: make(chan []byte, 256), allowedOrigins: allowedOrigins, isProd: isProd, connLimiter: rate.NewLimiter(maxWSConnections, maxWSConnections), } } func (h *WSHub) Run() { for { select { case client := <-h.register: h.mu.Lock() if len(h.clients) >= maxWSConnections { h.mu.Unlock() slog.Warn("ws max connections reached, rejecting") client.conn.Close() continue } h.clients[client] = true h.mu.Unlock() h.broadcastVisitors() case client := <-h.unregister: h.mu.Lock() if _, ok := h.clients[client]; ok { delete(h.clients, client) h.mu.Unlock() close(client.send) h.broadcastVisitors() } else { h.mu.Unlock() } case message := <-h.broadcast: h.mu.Lock() for client := range h.clients { select { case client.send <- message: default: close(client.send) delete(h.clients, client) } } h.mu.Unlock() } } } func (h *WSHub) broadcastVisitors() { h.mu.Lock() visitors := make([]VisitorDot, 0, len(h.clients)) clients := make([]*Client, 0, len(h.clients)) for c := range h.clients { visitors = append(visitors, VisitorDot{ UserID: c.userID, Lat: c.lat, Lng: c.lng, }) clients = append(clients, c) } h.mu.Unlock() data, err := json.Marshal(visitorsPayload{Visitors: visitors}) if err != nil { return } msg, err := json.Marshal(wsMessage{Type: "visitors", Data: data}) if err != nil { return } for _, client := range clients { select { case client.send <- msg: default: h.mu.Lock() close(client.send) delete(h.clients, client) h.mu.Unlock() } } } func (h *WSHub) HandleWS(w http.ResponseWriter, r *http.Request) { upgrader := websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { if len(h.allowedOrigins) == 0 { return !h.isProd } origin := r.Header.Get("Origin") for _, o := range h.allowedOrigins { if o == origin { return true } } return false }, } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { slog.Error("ws upgrade failed", "error", err) return } client := newClient(h, conn) userID := middleware.GetUserID(r.Context()) if userID != "" { client.userID = userID } if !h.connLimiter.Allow() { slog.Warn("ws connection rate limit exceeded") conn.Close() return } h.register <- client go client.writePump() go client.readPump() }