diff --git a/rpc/json.go b/rpc/json.go index f2df2d4665..dc863ccb78 100644 --- a/rpc/json.go +++ b/rpc/json.go @@ -142,6 +142,14 @@ type Conn interface { 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 // support for parsing arguments and serializing (result) objects. 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. - type remoteNetAddr interface{ RemoteAddr() net.Addr } type remoteStringAddr interface{ RemoteAddr() string } - if netra, ok := conn.(remoteNetAddr); ok { - codec.remoteAddr = netra.RemoteAddr().String() - } else if sra, ok := conn.(remoteStringAddr); ok { - codec.remoteAddr = sra.RemoteAddr() + type remoteNetAddr interface{ RemoteAddr() net.Addr } + if ra, ok := conn.(remoteStringAddr); ok { + codec.remoteAddr = ra.RemoteAddr() + } else if ra, ok := conn.(remoteNetAddr); ok { + codec.remoteAddr = ra.RemoteAddr().String() } return codec } diff --git a/rpc/websocket.go b/rpc/websocket.go index 32844e81fc..b8e067a5f2 100644 --- a/rpc/websocket.go +++ b/rpc/websocket.go @@ -76,7 +76,18 @@ func newWebsocketCodec(conn *websocket.Conn) ServerCodec { decoder := func(v interface{}) error { 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. @@ -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())) 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. origin := strings.ToLower(req.Header.Get("Origin")) if allowAllOrigins || origins.Contains(origin) {