Files
validator/OmCTF-2025/services/block_game/backend/messaging/messaging.go

164 lines
3.2 KiB
Go

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"
}