diff --git a/app/echo.go b/app/echo.go index 3bfddf1..43bd1e9 100644 --- a/app/echo.go +++ b/app/echo.go @@ -1,11 +1,24 @@ package main import ( + "fmt" + "time" "net/http" "github.com/gorilla/websocket" ) +var upgrader = websocket.Upgrader{ + ReadBufferSize: 1024, + WriteBufferSize: 1024, + CheckOrigin: func(r *http.Request) bool { + if r.Header.Get("origin") == "https://statbate.com" { + return true + } + return false + }, +} + var ( wsClients = make(map[*websocket.Conn]struct{}) @@ -13,37 +26,68 @@ var ( Count chan int Send chan []byte Add chan *websocket.Conn + Del chan *websocket.Conn }{ Count: make(chan int, 100), Send: make(chan []byte, 100), Add: make(chan *websocket.Conn, 100), + Del: make(chan *websocket.Conn, 100), } ) func broadcast() { + ticker := time.NewTicker(30 * time.Second) for { select { case conn := <-ws.Add: wsClients[conn] = struct{}{} + case conn := <-ws.Del: + delete(wsClients, conn) + case <-ws.Count: ws.Count <- len(wsClients) case message := <-ws.Send: - for conn := range wsClients { - if err := conn.WriteMessage(1, message); err != nil { - conn.Close() - delete(wsClients, conn) - } - } + sendMessage(message) + + case <-ticker.C: + sendMessage([]byte("ping")) + } + } +} + +func sendMessage(message []byte) { + for conn := range wsClients { + if err := conn.WriteMessage(1, message); err != nil { + conn.Close() + delete(wsClients, conn) } } } func wsHandler(w http.ResponseWriter, r *http.Request) { - conn, err := websocket.Upgrade(w, r, w.Header(), 1024, 1024) + conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } - ws.Add <- conn + go readWS(conn) +} + +func readWS(conn *websocket.Conn) { + defer conn.Close() + + ws.Add <- conn + + defer func() { + ws.Del <- conn + }() + + for { + _, _, err := conn.ReadMessage() + if err != nil { + fmt.Println("readWS", err.Error()) + return + } + } }