Add optional Origin handling for Websockets

This commit is contained in:
Sergey Stepanov 2021-12-15 17:58:49 +03:00
parent 535177bd46
commit ad822a624d
No known key found for this signature in database
GPG key ID: A56B4929BAA8556B
4 changed files with 60 additions and 13 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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" {