websocket.go 2.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899
  1. package handlers
  2. import (
  3. "log"
  4. "net/http"
  5. "sync"
  6. "time"
  7. "github.com/gorilla/websocket"
  8. )
  9. type VisitorDot struct {
  10. UserID string `json:"user_id,omitempty"`
  11. Lat float64 `json:"lat"`
  12. Lng float64 `json:"lng"`
  13. }
  14. type WSHub struct {
  15. mu sync.RWMutex
  16. clients map[*websocket.Conn]VisitorDot
  17. allowedOrigins []string
  18. isProd bool
  19. }
  20. func NewWSHub(allowedOrigins []string, isProd bool) *WSHub {
  21. if isProd && len(allowedOrigins) == 0 {
  22. log.Println("WARN: WebSocket AllowedOrigins is empty in production — connections will be rejected")
  23. }
  24. return &WSHub{
  25. clients: make(map[*websocket.Conn]VisitorDot),
  26. allowedOrigins: allowedOrigins,
  27. isProd: isProd,
  28. }
  29. }
  30. func (h *WSHub) HandleWS(w http.ResponseWriter, r *http.Request) {
  31. upgrader := websocket.Upgrader{
  32. CheckOrigin: func(r *http.Request) bool {
  33. if len(h.allowedOrigins) == 0 {
  34. return !h.isProd
  35. }
  36. origin := r.Header.Get("Origin")
  37. for _, o := range h.allowedOrigins {
  38. if o == origin {
  39. return true
  40. }
  41. }
  42. return false
  43. },
  44. }
  45. conn, err := upgrader.Upgrade(w, r, nil)
  46. if err != nil {
  47. log.Printf("ws upgrade: %v", err)
  48. return
  49. }
  50. dot := VisitorDot{}
  51. h.mu.Lock()
  52. h.clients[conn] = dot
  53. h.mu.Unlock()
  54. for {
  55. var msg VisitorDot
  56. if err := conn.ReadJSON(&msg); err != nil {
  57. break
  58. }
  59. h.mu.Lock()
  60. h.clients[conn] = msg
  61. visitors := make([]VisitorDot, 0, len(h.clients))
  62. for _, v := range h.clients {
  63. visitors = append(visitors, v)
  64. }
  65. var failed []*websocket.Conn
  66. for c := range h.clients {
  67. if err := c.SetWriteDeadline(time.Now().Add(10 * time.Second)); err != nil {
  68. failed = append(failed, c)
  69. continue
  70. }
  71. if err := c.WriteJSON(map[string]interface{}{"visitors": visitors}); err != nil {
  72. failed = append(failed, c)
  73. }
  74. }
  75. for _, c := range failed {
  76. c.WriteMessage(websocket.CloseMessage, []byte{})
  77. c.Close()
  78. delete(h.clients, c)
  79. }
  80. h.mu.Unlock()
  81. }
  82. h.mu.Lock()
  83. delete(h.clients, conn)
  84. h.mu.Unlock()
  85. conn.Close()
  86. }