websocket.go 3.1 KB

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