websocket.go 3.2 KB

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