websocket.go 3.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171
  1. package handlers
  2. import (
  3. "encoding/json"
  4. "log/slog"
  5. "net/http"
  6. "sync"
  7. "github.com/gorilla/websocket"
  8. "golang.org/x/time/rate"
  9. "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/middleware"
  10. "gogs.fxtmmsk.ru/foxtime/photoplaces/backend/internal/services"
  11. )
  12. type VisitorDot struct {
  13. UserID string `json:"user_id,omitempty"`
  14. Lat float64 `json:"lat"`
  15. Lng float64 `json:"lng"`
  16. }
  17. type visitorsPayload struct {
  18. Visitors []VisitorDot `json:"visitors"`
  19. }
  20. const maxWSConnections = 1000
  21. type WSHub struct {
  22. mu sync.Mutex
  23. clients map[*Client]bool
  24. register chan *Client
  25. unregister chan *Client
  26. broadcast chan []byte
  27. allowedOrigins []string
  28. isProd bool
  29. authSvc *services.AuthService
  30. connLimiter *rate.Limiter
  31. }
  32. func NewWSHub(allowedOrigins []string, isProd bool, authSvc *services.AuthService) *WSHub {
  33. return &WSHub{
  34. clients: make(map[*Client]bool),
  35. register: make(chan *Client),
  36. unregister: make(chan *Client),
  37. broadcast: make(chan []byte, 256),
  38. allowedOrigins: allowedOrigins,
  39. isProd: isProd,
  40. authSvc: authSvc,
  41. connLimiter: rate.NewLimiter(maxWSConnections, maxWSConnections),
  42. }
  43. }
  44. func (h *WSHub) Run() {
  45. for {
  46. select {
  47. case client := <-h.register:
  48. h.mu.Lock()
  49. if len(h.clients) >= maxWSConnections {
  50. h.mu.Unlock()
  51. slog.Warn("ws max connections reached, rejecting")
  52. client.conn.Close()
  53. continue
  54. }
  55. h.clients[client] = true
  56. h.mu.Unlock()
  57. h.broadcastVisitors()
  58. case client := <-h.unregister:
  59. h.mu.Lock()
  60. if _, ok := h.clients[client]; ok {
  61. delete(h.clients, client)
  62. h.mu.Unlock()
  63. close(client.send)
  64. h.broadcastVisitors()
  65. } else {
  66. h.mu.Unlock()
  67. }
  68. case message := <-h.broadcast:
  69. h.mu.Lock()
  70. for client := range h.clients {
  71. select {
  72. case client.send <- message:
  73. default:
  74. close(client.send)
  75. delete(h.clients, client)
  76. }
  77. }
  78. h.mu.Unlock()
  79. }
  80. }
  81. }
  82. func (h *WSHub) broadcastVisitors() {
  83. h.mu.Lock()
  84. visitors := make([]VisitorDot, 0, len(h.clients))
  85. for c := range h.clients {
  86. visitors = append(visitors, VisitorDot{
  87. UserID: c.userID,
  88. Lat: c.lat,
  89. Lng: c.lng,
  90. })
  91. }
  92. clients := make([]*Client, 0, len(h.clients))
  93. for c := range h.clients {
  94. clients = append(clients, c)
  95. }
  96. h.mu.Unlock()
  97. data, err := json.Marshal(visitorsPayload{Visitors: visitors})
  98. if err != nil {
  99. return
  100. }
  101. msg, err := json.Marshal(wsMessage{Type: "visitors", Data: data})
  102. if err != nil {
  103. return
  104. }
  105. for _, client := range clients {
  106. select {
  107. case client.send <- msg:
  108. default:
  109. h.mu.Lock()
  110. close(client.send)
  111. delete(h.clients, client)
  112. h.mu.Unlock()
  113. }
  114. }
  115. }
  116. func (h *WSHub) HandleWS(w http.ResponseWriter, r *http.Request) {
  117. upgrader := websocket.Upgrader{
  118. CheckOrigin: func(r *http.Request) bool {
  119. if len(h.allowedOrigins) == 0 {
  120. return !h.isProd
  121. }
  122. origin := r.Header.Get("Origin")
  123. for _, o := range h.allowedOrigins {
  124. if o == origin {
  125. return true
  126. }
  127. }
  128. return false
  129. },
  130. }
  131. conn, err := upgrader.Upgrade(w, r, nil)
  132. if err != nil {
  133. slog.Error("ws upgrade failed", "error", err)
  134. return
  135. }
  136. client := newClient(h, conn)
  137. userID := middleware.GetUserID(r.Context())
  138. if userID != "" {
  139. client.userID = userID
  140. }
  141. if !h.connLimiter.Allow() {
  142. slog.Warn("ws connection rate limit exceeded")
  143. conn.Close()
  144. return
  145. }
  146. h.register <- client
  147. go client.writePump()
  148. go client.readPump()
  149. }