diff --git a/configs/config.yaml b/configs/config.yaml index 9617e6ab..a14103eb 100644 --- a/configs/config.yaml +++ b/configs/config.yaml @@ -32,6 +32,13 @@ coordinator: profilingEnabled: false metricEnabled: false urlPrefix: /coordinator + # a custom Origins for incoming Websocket connections: + # "" -- checks same origin policy + # "*" -- allows all + # "your address" -- checks for that address + origin: + userWs: + workerWs: # HTTP(S) server config server: address: :8000 diff --git a/pkg/config/coordinator/config.go b/pkg/config/coordinator/config.go index f326f193..89076169 100644 --- a/pkg/config/coordinator/config.go +++ b/pkg/config/coordinator/config.go @@ -15,8 +15,12 @@ type Config struct { DebugHost string Library games.Config Monitoring monitoring.Config - Server shared.Server - Analytics Analytics + Origin struct { + UserWs string + WorkerWs string + } + Server shared.Server + Analytics Analytics } Emulator emulator.Emulator Environment shared.Environment diff --git a/pkg/coordinator/handlers.go b/pkg/coordinator/handlers.go index df8b4e1e..16fed446 100644 --- a/pkg/coordinator/handlers.go +++ b/pkg/coordinator/handlers.go @@ -14,10 +14,10 @@ import ( "github.com/giongto35/cloud-game/v2/pkg/environment" "github.com/giongto35/cloud-game/v2/pkg/games" "github.com/giongto35/cloud-game/v2/pkg/ice" + "github.com/giongto35/cloud-game/v2/pkg/network/websocket" "github.com/giongto35/cloud-game/v2/pkg/service" "github.com/giongto35/cloud-game/v2/pkg/util" "github.com/gofrs/uuid" - "github.com/gorilla/websocket" ) type Server struct { @@ -32,15 +32,15 @@ type Server struct { workerClients map[string]*WorkerClient // browserClients are the map sessionID to browser Client browserClients map[string]*BrowserClient -} -var upgrader = websocket.Upgrader{} + userWsUpgrader, workerWsUpgrader websocket.Upgrader +} func NewServer(cfg coordinator.Config, library games.GameLibrary) *Server { // scan the lib right away library.Scan() - return &Server{ + s := &Server{ cfg: cfg, library: library, // Mapping roomID to server @@ -50,6 +50,12 @@ func NewServer(cfg coordinator.Config, library games.GameLibrary) *Server { // Mapping sessionID to browser browserClients: map[string]*BrowserClient{}, } + + // a custom Origin check + s.workerWsUpgrader = websocket.NewUpgrader(cfg.Coordinator.Origin.WorkerWs) + s.userWsUpgrader = websocket.NewUpgrader(cfg.Coordinator.Origin.UserWs) + + return s } // WSO handles all connections from a new worker to coordinator @@ -70,9 +76,7 @@ func (s *Server) WSO(w http.ResponseWriter, r *http.Request) { log.Printf("Warning! Unsecure connection. The worker may not work properly without HTTPS on its side!") } - // be aware of ReadBufferSize, WriteBufferSize (default 4096) - // https://pkg.go.dev/github.com/gorilla/websocket?tab=doc#Upgrader - c, err := upgrader.Upgrade(w, r, nil) + c, err := s.workerWsUpgrader.Upgrade(w, r, nil) if err != nil { log.Println("Coordinator: [!] WS upgrade:", err) return @@ -130,18 +134,17 @@ func (s *Server) WSO(w http.ResponseWriter, r *http.Request) { wc.Listen() } -// WSO handles all connections from user/frontend to coordinator +// WS handles all connections from user/frontend to coordinator func (s *Server) WS(w http.ResponseWriter, r *http.Request) { log.Println("Coordinator: A user is connecting...") + defer func() { if r := recover(); r != nil { log.Println("Warn: Something wrong. Recovered in ", r) } }() - // be aware of ReadBufferSize, WriteBufferSize (default 4096) - // https://pkg.go.dev/github.com/gorilla/websocket?tab=doc#Upgrader - c, err := upgrader.Upgrade(w, r, nil) + c, err := s.userWsUpgrader.Upgrade(w, r, nil) if err != nil { log.Println("Coordinator: [!] WS upgrade:", err) return diff --git a/pkg/network/websocket/websocket.go b/pkg/network/websocket/websocket.go index 18687162..7f3c3a66 100644 --- a/pkg/network/websocket/websocket.go +++ b/pkg/network/websocket/websocket.go @@ -2,11 +2,44 @@ package websocket import ( "crypto/tls" + "net/http" "net/url" "github.com/gorilla/websocket" ) +type Upgrader struct { + websocket.Upgrader + + origin string +} + +var DefaultUpgrader = Upgrader{ + Upgrader: websocket.Upgrader{ + ReadBufferSize: 4096, + WriteBufferSize: 4096, + EnableCompression: false, + }, +} + +func NewUpgrader(origin string) Upgrader { + u := DefaultUpgrader + switch { + case origin == "*": + u.CheckOrigin = func(r *http.Request) bool { return true } + case origin != "": + u.CheckOrigin = func(r *http.Request) bool { return r.Header.Get("Origin") == origin } + } + return u +} + +func (u *Upgrader) Upgrade(w http.ResponseWriter, r *http.Request, responseHeader http.Header) (*websocket.Conn, error) { + if u.origin != "" { + w.Header().Set("Access-Control-Allow-Origin", u.origin) + } + return u.Upgrader.Upgrade(w, r, responseHeader) +} + func Connect(address url.URL) (*websocket.Conn, error) { dialer := websocket.Dialer{} if address.Scheme == "wss" {