From 6d361db63869f7313b27f1067a457bd5b2f5685e Mon Sep 17 00:00:00 2001 From: Claudio Costa Date: Tue, 7 Dec 2021 15:24:18 +0100 Subject: [PATCH] [MM-40485] Enable receiving binary websocket messages (#19128) * Enable receiving binary websocket messages * Improve error message * Prefer anonymous declaration * Simplify * Improve test * Use MessagePack to clone WebSocketRequest struct * Use short form * Fix test --- api4/websocket_test.go | 41 ++++++++++++++++++++++++++ app/web_conn.go | 19 ++++++++++-- model/websocket_client.go | 26 ++++++++++++++++ model/websocket_client_test.go | 54 ++++++++++++++++++++++++++++++++++ model/websocket_request.go | 20 ++++++------- 5 files changed, 148 insertions(+), 12 deletions(-) diff --git a/api4/websocket_test.go b/api4/websocket_test.go index 802b34b8c6..7ff693330b 100644 --- a/api4/websocket_test.go +++ b/api4/websocket_test.go @@ -234,6 +234,47 @@ func TestWebSocketReconnectRace(t *testing.T) { wg.Wait() } +func TestWebSocketSendBinary(t *testing.T) { + th := Setup(t).InitBasic() + defer th.TearDown() + + client := th.CreateClient() + th.LoginBasicWithClient(client) + WebSocketClient, err := th.CreateWebSocketClientWithClient(client) + require.NoError(t, err) + defer WebSocketClient.Close() + WebSocketClient.Listen() + resp := <-WebSocketClient.ResponseChannel + require.Equal(t, resp.Status, model.StatusOk) + + client2 := th.CreateClient() + th.LoginBasic2WithClient(client2) + WebSocketClient2, err := th.CreateWebSocketClientWithClient(client2) + require.NoError(t, err) + defer WebSocketClient2.Close() + + time.Sleep(1000 * time.Millisecond) + + WebSocketClient.SendBinaryMessage("get_statuses", nil) + resp = <-WebSocketClient.ResponseChannel + require.Nil(t, resp.Error, resp.Error) + require.Equal(t, resp.SeqReply, WebSocketClient.Sequence-1) + + status, ok := resp.Data[th.BasicUser.Id] + require.True(t, ok) + require.Equal(t, model.StatusOnline, status) + status, ok = resp.Data[th.BasicUser2.Id] + require.True(t, ok) + require.Equal(t, model.StatusOnline, status) + + WebSocketClient.SendBinaryMessage("get_statuses_by_ids", map[string]interface{}{ + "user_ids": []string{th.BasicUser2.Id}, + }) + status, ok = resp.Data[th.BasicUser2.Id] + require.True(t, ok) + require.Equal(t, model.StatusOnline, status) +} + func TestWebSocketStatuses(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() diff --git a/app/web_conn.go b/app/web_conn.go index 9a4f329345..f71532174c 100644 --- a/app/web_conn.go +++ b/app/web_conn.go @@ -17,6 +17,7 @@ import ( "time" "github.com/gorilla/websocket" + "github.com/vmihailenco/msgpack/v5" "github.com/mattermost/mattermost-server/v6/model" "github.com/mattermost/mattermost-server/v6/plugin" @@ -343,9 +344,23 @@ func (wc *WebConn) readPump() { }) for { + msgType, rd, err := wc.WebSocket.NextReader() + if err != nil { + wc.logSocketErr("websocket.NextReader", err) + return + } + + var decoder interface { + Decode(v interface{}) error + } + if msgType == websocket.TextMessage { + decoder = json.NewDecoder(rd) + } else { + decoder = msgpack.NewDecoder(rd) + } var req model.WebSocketRequest - if err := wc.WebSocket.ReadJSON(&req); err != nil { - wc.logSocketErr("websocket.read", err) + if err = decoder.Decode(&req); err != nil { + wc.logSocketErr("websocket.Decode", err) return } diff --git a/model/websocket_client.go b/model/websocket_client.go index df35695a8f..b80d477751 100644 --- a/model/websocket_client.go +++ b/model/websocket_client.go @@ -14,6 +14,7 @@ import ( "github.com/mattermost/mattermost-server/v6/shared/mlog" "github.com/gorilla/websocket" + "github.com/vmihailenco/msgpack/v5" ) const ( @@ -26,6 +27,7 @@ type msgType int const ( msgTypeJSON msgType = iota + 1 msgTypePong + msgTypeBinary ) type writeMessage struct { @@ -182,6 +184,10 @@ func (wsc *WebSocketClient) writer() { switch msg.msgType { case msgTypeJSON: wsc.Conn.WriteJSON(msg.data) + case msgTypeBinary: + if data, ok := msg.data.([]byte); ok { + wsc.Conn.WriteMessage(websocket.BinaryMessage, data) + } case msgTypePong: wsc.Conn.WriteMessage(websocket.PongMessage, []byte{}) } @@ -275,6 +281,26 @@ func (wsc *WebSocketClient) SendMessage(action string, data map[string]interface } } +func (wsc *WebSocketClient) SendBinaryMessage(action string, data map[string]interface{}) error { + req := &WebSocketRequest{} + req.Seq = wsc.Sequence + req.Action = action + req.Data = data + + binaryData, err := msgpack.Marshal(req) + if err != nil { + return fmt.Errorf("failed to marshal request to msgpack: %w", err) + } + + wsc.Sequence++ + wsc.writeChan <- writeMessage{ + msgType: msgTypeBinary, + data: binaryData, + } + + return nil +} + // UserTyping will push a user_typing event out to all connected users // who are in the specified channel func (wsc *WebSocketClient) UserTyping(channelId, parentId string) { diff --git a/model/websocket_client_test.go b/model/websocket_client_test.go index 28d5ffb7ee..cbf1128ae9 100644 --- a/model/websocket_client_test.go +++ b/model/websocket_client_test.go @@ -13,6 +13,7 @@ import ( "github.com/gorilla/websocket" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/vmihailenco/msgpack/v5" ) func dummyWebsocketHandler(t *testing.T) http.HandlerFunc { @@ -149,3 +150,56 @@ func TestWebSocketClose(t *testing.T) { checkWriteChan(cli.writeChan) }) } + +func binaryWebsocketHandler(t *testing.T, clientData map[string]interface{}, doneCh chan struct{}) http.HandlerFunc { + return func(w http.ResponseWriter, req *http.Request) { + defer close(doneCh) + upgrader := &websocket.Upgrader{ + ReadBufferSize: 1024, + WriteBufferSize: 1024, + } + conn, err := upgrader.Upgrade(w, req, nil) + require.NoError(t, err) + defer conn.Close() + + for { + msgType, buf, err := conn.ReadMessage() + require.NoError(t, err) + if msgType == websocket.BinaryMessage { + require.Equal(t, msgType, websocket.BinaryMessage) + wsReq := &WebSocketRequest{} + err = msgpack.Unmarshal(buf, wsReq) + require.NoError(t, err) + require.Equal(t, clientData, wsReq.Data) + break + } + } + } +} + +func TestWebSocketSendBinaryMessage(t *testing.T) { + clientData := map[string]interface{}{ + "data": []byte("some data to send as binary"), + } + + doneCh := make(chan struct{}) + s := httptest.NewServer(binaryWebsocketHandler(t, clientData, doneCh)) + defer s.Close() + + url := strings.Replace(s.URL, "http://", "ws://", 1) + cli, err := NewWebSocketClient4(url, "authToken") + require.NoError(t, err) + cli.Listen() + defer cli.Close() + + err = cli.SendBinaryMessage("binaryAction", map[string]interface{}{ + "unmarshable": func() {}, + }) + require.Error(t, err) + + err = cli.SendBinaryMessage("binaryAction", clientData) + require.NoError(t, err) + + // This is to make sure the message is handled prior to exiting. + <-doneCh +} diff --git a/model/websocket_request.go b/model/websocket_request.go index 9e86397833..a7750bcea8 100644 --- a/model/websocket_request.go +++ b/model/websocket_request.go @@ -4,31 +4,31 @@ package model import ( - "encoding/json" - "github.com/mattermost/mattermost-server/v6/shared/i18n" + + "github.com/vmihailenco/msgpack/v5" ) // WebSocketRequest represents a request made to the server through a websocket. type WebSocketRequest struct { // Client-provided fields - Seq int64 `json:"seq"` // A counter which is incremented for every request made. - Action string `json:"action"` // The action to perform for a request. For example: get_statuses, user_typing. - Data map[string]interface{} `json:"data"` // The metadata for an action. + Seq int64 `json:"seq" msgpack:"seq"` // A counter which is incremented for every request made. + Action string `json:"action" msgpack:"action"` // The action to perform for a request. For example: get_statuses, user_typing. + Data map[string]interface{} `json:"data" msgpack:"data"` // The metadata for an action. // Server-provided fields - Session Session `json:"-"` - T i18n.TranslateFunc `json:"-"` - Locale string `json:"-"` + Session Session `json:"-" msgpack:"-"` + T i18n.TranslateFunc `json:"-" msgpack:"-"` + Locale string `json:"-" msgpack:"-"` } func (o *WebSocketRequest) Clone() (*WebSocketRequest, error) { - buf, err := json.Marshal(o) + buf, err := msgpack.Marshal(o) if err != nil { return nil, err } var ret WebSocketRequest - err = json.Unmarshal(buf, &ret) + err = msgpack.Unmarshal(buf, &ret) if err != nil { return nil, err }