rpc: fix crash if websocket conn has no origin

This commit is contained in:
Felix Lange 2019-01-18 11:12:08 +01:00
parent 6085cb2240
commit 962f4ed490
2 changed files with 25 additions and 9 deletions

View file

@ -142,6 +142,14 @@ type Conn interface {
SetWriteDeadline(time.Time) error SetWriteDeadline(time.Time) error
} }
// connWithRemoteAddr overrides the remote address of a connection.
type connWithRemoteAddr struct {
Conn
addr string
}
func (c connWithRemoteAddr) RemoteAddr() string { return c.addr }
// jsonCodec reads and writes JSON-RPC messages to the underlying connection. It also has // jsonCodec reads and writes JSON-RPC messages to the underlying connection. It also has
// support for parsing arguments and serializing (result) objects. // support for parsing arguments and serializing (result) objects.
type jsonCodec struct { type jsonCodec struct {
@ -166,12 +174,12 @@ func NewCodec(conn Conn, encode, decode func(v interface{}) error) ServerCodec {
} }
// Try to figure out the remote address. // Try to figure out the remote address.
type remoteNetAddr interface{ RemoteAddr() net.Addr }
type remoteStringAddr interface{ RemoteAddr() string } type remoteStringAddr interface{ RemoteAddr() string }
if netra, ok := conn.(remoteNetAddr); ok { type remoteNetAddr interface{ RemoteAddr() net.Addr }
codec.remoteAddr = netra.RemoteAddr().String() if ra, ok := conn.(remoteStringAddr); ok {
} else if sra, ok := conn.(remoteStringAddr); ok { codec.remoteAddr = ra.RemoteAddr()
codec.remoteAddr = sra.RemoteAddr() } else if ra, ok := conn.(remoteNetAddr); ok {
codec.remoteAddr = ra.RemoteAddr().String()
} }
return codec return codec
} }

View file

@ -76,7 +76,18 @@ func newWebsocketCodec(conn *websocket.Conn) ServerCodec {
decoder := func(v interface{}) error { decoder := func(v interface{}) error {
return websocketJSONCodec.Receive(conn, v) return websocketJSONCodec.Receive(conn, v)
} }
return NewCodec(conn, encoder, decoder) rpcconn := Conn(conn)
if conn.IsServerConn() {
// Override remote address with the actual socket address because
// package websocket crashes if there is no request origin.
addr := conn.Request().RemoteAddr
if wsaddr := conn.RemoteAddr().(*websocket.Addr); wsaddr.URL != nil {
// Add origin if present.
addr += "(" + wsaddr.URL.String() + ")"
}
rpcconn = connWithRemoteAddr{conn, addr}
}
return NewCodec(rpcconn, encoder, decoder)
} }
// NewWSServer creates a new websocket RPC server around an API provider. // NewWSServer creates a new websocket RPC server around an API provider.
@ -113,9 +124,6 @@ func wsHandshakeValidator(allowedOrigins []string) func(*websocket.Config, *http
log.Debug(fmt.Sprintf("Allowed origin(s) for WS RPC interface %v", origins.ToSlice())) log.Debug(fmt.Sprintf("Allowed origin(s) for WS RPC interface %v", origins.ToSlice()))
f := func(cfg *websocket.Config, req *http.Request) error { f := func(cfg *websocket.Config, req *http.Request) error {
// Set config origin to the peer address to make RemoteAddr work.
cfg.Origin = &url.URL{Scheme: "ws", Host: req.RemoteAddr}
// Verify origin against whitelist. // Verify origin against whitelist.
origin := strings.ToLower(req.Header.Get("Origin")) origin := strings.ToLower(req.Header.Get("Origin"))
if allowAllOrigins || origins.Contains(origin) { if allowAllOrigins || origins.Contains(origin) {