package messaging import ( "encoding/json" "errors" "fmt" "log" "net/http" "sync" "time" "github.com/gorilla/websocket" "omctf.ru/block-game-backend/utils" ) type Message struct { Type string `json:"type"` Option json.RawMessage `json:"option"` } // Simpler wrapper for gorilla/websocket.Conn type Conn struct { RawConn *websocket.Conn Recv chan Message send chan Message closeOnce sync.Once } var upgrader = websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, CheckOrigin: func(r *http.Request) bool { // return r.Header.Get("Origin") == "localhost:8080" return true // Temporarily allow all origins }, } func InitWebsocket(w http.ResponseWriter, r *http.Request) (*Conn, error) { rawConn, err := upgrader.Upgrade(w, r, nil) if err != nil { return nil, err } conn := &Conn{ RawConn: rawConn, Recv: make(chan Message, 5), send: make(chan Message, 5), closeOnce: sync.Once{}, } return conn, nil } func (c *Conn) Close() { c.closeOnce.Do(func() { close(c.Recv) close(c.send) c.RawConn.Close() }) } const ( pingPeriod = time.Second * 10 ) func (c *Conn) LoopReceive() { c.RawConn.SetReadDeadline(time.Now().Add(pingPeriod * 2)) c.RawConn.SetPongHandler(func(string) error { c.RawConn.SetReadDeadline(time.Now().Add(pingPeriod * 2)) return nil }) defer c.Close() for { var msg Message err := c.RawConn.ReadJSON(&msg) if errors.Is(err, &json.SyntaxError{}) || errors.Is(err, &json.UnmarshalTypeError{}) { c.Error(ErrorBadRequestf("invalid json: %w", err)) continue } else if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) { return } else if err != nil { log.Printf("error reading message: %s", err) c.Error(fmt.Errorf("internal server error")) return } c.Recv <- msg } } func (c *Conn) LoopSend() { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() c.Close() }() for { select { case _, ok := <-ticker.C: if !ok { return } c.RawConn.SetWriteDeadline(time.Now().Add(pingPeriod)) if err := c.RawConn.WriteMessage(websocket.PingMessage, nil); err != nil { return } case msg, ok := <-c.send: if !ok { return } c.RawConn.SetWriteDeadline(time.Now().Add(pingPeriod)) err := c.RawConn.WriteJSON(msg) if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) { return } else if err != nil { log.Printf("error sending message: %v", err) if msg.Type != "error" { c.Error(fmt.Errorf("internal server error")) } return } } } } func (c *Conn) Start() { go c.LoopReceive() go c.LoopSend() } func (c *Conn) Send(msg Message) { defer func() { recover() }() c.send <- msg } func (c *Conn) Error(err error) { c.Send(Message{"error", utils.MustMarshal(err.Error())}) } type ErrorBadRequest struct { Reason error } func ErrorBadRequestf(s string, args ...any) error { return ErrorBadRequest{Reason: fmt.Errorf(s, args...)} } func (e ErrorBadRequest) Error() string { return fmt.Sprintf("bad request: %s", e.Reason.Error()) } func (e ErrorBadRequest) Unwrap() error { return e.Reason } type ErrorConnectionClosed struct{} func (e ErrorConnectionClosed) Error() string { return "connection is closed" }