websocket.go 2.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. package handlers
  2. import (
  3. "encoding/json"
  4. "log"
  5. "net/http"
  6. "github.com/gorilla/websocket"
  7. )
  8. type VisitorDot struct {
  9. UserID string `json:"user_id,omitempty"`
  10. Lat float64 `json:"lat"`
  11. Lng float64 `json:"lng"`
  12. }
  13. type visitorsPayload struct {
  14. Visitors []VisitorDot `json:"visitors"`
  15. }
  16. type WSHub struct {
  17. clients map[*Client]bool
  18. register chan *Client
  19. unregister chan *Client
  20. broadcast chan []byte
  21. allowedOrigins []string
  22. isProd bool
  23. }
  24. func NewWSHub(allowedOrigins []string, isProd bool) *WSHub {
  25. return &WSHub{
  26. clients: make(map[*Client]bool),
  27. register: make(chan *Client),
  28. unregister: make(chan *Client),
  29. broadcast: make(chan []byte, 256),
  30. allowedOrigins: allowedOrigins,
  31. isProd: isProd,
  32. }
  33. }
  34. func (h *WSHub) Run() {
  35. for {
  36. select {
  37. case client := <-h.register:
  38. h.clients[client] = true
  39. h.broadcastVisitors()
  40. case client := <-h.unregister:
  41. if _, ok := h.clients[client]; ok {
  42. delete(h.clients, client)
  43. close(client.send)
  44. h.broadcastVisitors()
  45. }
  46. case message := <-h.broadcast:
  47. for client := range h.clients {
  48. select {
  49. case client.send <- message:
  50. default:
  51. close(client.send)
  52. delete(h.clients, client)
  53. }
  54. }
  55. }
  56. }
  57. }
  58. func (h *WSHub) broadcastVisitors() {
  59. visitors := make([]VisitorDot, 0, len(h.clients))
  60. for c := range h.clients {
  61. visitors = append(visitors, VisitorDot{
  62. UserID: c.userID,
  63. Lat: c.lat,
  64. Lng: c.lng,
  65. })
  66. }
  67. data, err := json.Marshal(visitorsPayload{Visitors: visitors})
  68. if err != nil {
  69. return
  70. }
  71. msg, err := json.Marshal(wsMessage{Type: "visitors", Data: data})
  72. if err != nil {
  73. return
  74. }
  75. for client := range h.clients {
  76. select {
  77. case client.send <- msg:
  78. default:
  79. close(client.send)
  80. delete(h.clients, client)
  81. }
  82. }
  83. }
  84. func (h *WSHub) HandleWS(w http.ResponseWriter, r *http.Request) {
  85. upgrader := websocket.Upgrader{
  86. CheckOrigin: func(r *http.Request) bool {
  87. if len(h.allowedOrigins) == 0 {
  88. return !h.isProd
  89. }
  90. origin := r.Header.Get("Origin")
  91. for _, o := range h.allowedOrigins {
  92. if o == origin {
  93. return true
  94. }
  95. }
  96. return false
  97. },
  98. }
  99. conn, err := upgrader.Upgrade(w, r, nil)
  100. if err != nil {
  101. log.Printf("ws upgrade: %v", err)
  102. return
  103. }
  104. client := &Client{
  105. hub: h,
  106. conn: conn,
  107. send: make(chan []byte, 256),
  108. }
  109. h.register <- client
  110. go client.writePump()
  111. go client.readPump()
  112. }