| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165 |
- 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()
- }
|