[MM-19914] Fix data races in WebSocketEvent (#13039)

* Make WebSocketEvent type immutable

* Update code to use updated immutable WebSocketEvent type

* Export WebSocketEvent fields and mark them as deprecated
Этот коммит содержится в:
Claudio Costa
2019-12-24 09:32:11 +01:00
коммит произвёл GitHub
родитель 9ab7bee0a6
Коммит 80dd2915db
16 изменённых файлов: 216 добавлений и 94 удалений

Просмотреть файл

@@ -4,6 +4,7 @@
package model
import (
"bytes"
"encoding/json"
"net/http"
"time"
@@ -121,9 +122,9 @@ func (wsc *WebSocketClient) Listen() {
return
}
var event WebSocketEvent
if err := json.Unmarshal(rawMsg, &event); err == nil && event.IsValid() {
wsc.EventChannel <- &event
event := WebSocketEventFromJson(bytes.NewReader(rawMsg))
if event.IsValid() {
wsc.EventChannel <- event
continue
}

Просмотреть файл

@@ -76,30 +76,41 @@ type precomputedWebSocketEventJSON struct {
Broadcast json.RawMessage
}
type WebSocketEvent struct {
// webSocketEventJSON mirrors WebSocketEvent to make some of its unexported fields serializable
type webSocketEventJSON struct {
Event string `json:"event"`
Data map[string]interface{} `json:"data"`
Broadcast *WebsocketBroadcast `json:"broadcast"`
Sequence int64 `json:"seq"`
}
// **NOTE**: Direct access to WebSocketEvent fields is deprecated. They will be
// made unexported in next major version release. Provided getter functions should be used instead.
type WebSocketEvent struct {
Event string // Deprecated: use EventType()
Data map[string]interface{} // Deprecated: use GetData()
Broadcast *WebsocketBroadcast // Deprecated: use GetBroadcast()
Sequence int64 // Deprecated: use GetSequence()
precomputedJSON *precomputedWebSocketEventJSON
}
// PrecomputeJSON precomputes and stores the serialized JSON for all fields other than Sequence.
// This makes ToJson much more efficient when sending the same event to multiple connections.
func (m *WebSocketEvent) PrecomputeJSON() {
event, _ := json.Marshal(m.Event)
data, _ := json.Marshal(m.Data)
broadcast, _ := json.Marshal(m.Broadcast)
m.precomputedJSON = &precomputedWebSocketEventJSON{
func (ev *WebSocketEvent) PrecomputeJSON() *WebSocketEvent {
copy := ev.Copy()
event, _ := json.Marshal(copy.Event)
data, _ := json.Marshal(copy.Data)
broadcast, _ := json.Marshal(copy.Broadcast)
copy.precomputedJSON = &precomputedWebSocketEventJSON{
Event: json.RawMessage(event),
Data: json.RawMessage(data),
Broadcast: json.RawMessage(broadcast),
}
return copy
}
func (m *WebSocketEvent) Add(key string, value interface{}) {
m.Data[key] = value
func (ev *WebSocketEvent) Add(key string, value interface{}) {
ev.Data[key] = value
}
func NewWebSocketEvent(event, teamId, channelId, userId string, omitUsers map[string]bool) *WebSocketEvent {
@@ -107,26 +118,85 @@ func NewWebSocketEvent(event, teamId, channelId, userId string, omitUsers map[st
Broadcast: &WebsocketBroadcast{TeamId: teamId, ChannelId: channelId, UserId: userId, OmitUsers: omitUsers}}
}
func (o *WebSocketEvent) IsValid() bool {
return o.Event != ""
}
func (o *WebSocketEvent) EventType() string {
return o.Event
}
func (o *WebSocketEvent) ToJson() string {
if o.precomputedJSON != nil {
return fmt.Sprintf(`{"event": %s, "data": %s, "broadcast": %s, "seq": %d}`, o.precomputedJSON.Event, o.precomputedJSON.Data, o.precomputedJSON.Broadcast, o.Sequence)
func (ev *WebSocketEvent) Copy() *WebSocketEvent {
copy := &WebSocketEvent{
Event: ev.Event,
Data: ev.Data,
Broadcast: ev.Broadcast,
Sequence: ev.Sequence,
precomputedJSON: ev.precomputedJSON,
}
b, _ := json.Marshal(o)
return copy
}
func (ev *WebSocketEvent) GetData() map[string]interface{} {
return ev.Data
}
func (ev *WebSocketEvent) GetBroadcast() *WebsocketBroadcast {
return ev.Broadcast
}
func (ev *WebSocketEvent) GetSequence() int64 {
return ev.Sequence
}
func (ev *WebSocketEvent) SetEvent(event string) *WebSocketEvent {
copy := ev.Copy()
copy.Event = event
return copy
}
func (ev *WebSocketEvent) SetData(data map[string]interface{}) *WebSocketEvent {
copy := ev.Copy()
copy.Data = data
return copy
}
func (ev *WebSocketEvent) SetBroadcast(broadcast *WebsocketBroadcast) *WebSocketEvent {
copy := ev.Copy()
copy.Broadcast = broadcast
return copy
}
func (ev *WebSocketEvent) SetSequence(seq int64) *WebSocketEvent {
copy := ev.Copy()
copy.Sequence = seq
return copy
}
func (ev *WebSocketEvent) IsValid() bool {
return ev.Event != ""
}
func (ev *WebSocketEvent) EventType() string {
return ev.Event
}
func (ev *WebSocketEvent) ToJson() string {
if ev.precomputedJSON != nil {
return fmt.Sprintf(`{"event": %s, "data": %s, "broadcast": %s, "seq": %d}`, ev.precomputedJSON.Event, ev.precomputedJSON.Data, ev.precomputedJSON.Broadcast, ev.Sequence)
}
b, _ := json.Marshal(webSocketEventJSON{
ev.Event,
ev.Data,
ev.Broadcast,
ev.Sequence,
})
return string(b)
}
func WebSocketEventFromJson(data io.Reader) *WebSocketEvent {
var o *WebSocketEvent
json.NewDecoder(data).Decode(&o)
return o
var ev WebSocketEvent
var o webSocketEventJSON
if err := json.NewDecoder(data).Decode(&o); err != nil {
return nil
}
ev.Event = o.Event
ev.Data = o.Data
ev.Broadcast = o.Broadcast
ev.Sequence = o.Sequence
return &ev
}
type WebSocketResponse struct {

Просмотреть файл

@@ -17,14 +17,68 @@ func TestWebSocketEvent(t *testing.T) {
json := m.ToJson()
result := WebSocketEventFromJson(strings.NewReader(json))
badresult := WebSocketEventFromJson(strings.NewReader("junk"))
require.Nil(t, badresult, "should not have parsed")
require.True(t, m.IsValid(), "should be valid")
require.Equal(t, m.GetBroadcast().TeamId, result.GetBroadcast().TeamId, "Ids do not match")
require.Equal(t, m.GetData()["RootId"], result.GetData()["RootId"], "Ids do not match")
}
require.Equal(t, m.Broadcast.TeamId, result.Broadcast.TeamId, "Ids do not match")
func TestWebSocketEventImmutable(t *testing.T) {
m := NewWebSocketEvent("some_event", NewId(), NewId(), NewId(), nil)
require.Equal(t, m.Data["RootId"], result.Data["RootId"], "Ids do not match")
new := m.SetEvent("new_event")
if new == m {
require.Fail(t, "pointers should not be the same")
}
require.NotEqual(t, m.Event, new.Event)
require.Equal(t, new.Event, "new_event")
require.Equal(t, new.Event, new.EventType())
new = m.SetSequence(45)
if new == m {
require.Fail(t, "pointers should not be the same")
}
require.NotEqual(t, m.Sequence, new.Sequence)
require.Equal(t, new.Sequence, int64(45))
require.Equal(t, new.Sequence, new.GetSequence())
broadcast := &WebsocketBroadcast{}
new = m.SetBroadcast(broadcast)
if new == m {
require.Fail(t, "pointers should not be the same")
}
require.NotEqual(t, m.Broadcast, new.Broadcast)
require.Equal(t, new.Broadcast, broadcast)
require.Equal(t, new.Broadcast, new.GetBroadcast())
data := map[string]interface{}{
"key": "val",
"key2": "val2",
}
new = m.SetData(data)
if new == m {
require.Fail(t, "pointers should not be the same")
}
require.NotEqual(t, m, new)
require.Equal(t, new.Data, data)
require.Equal(t, new.Data, new.GetData())
copy := m.Copy()
if copy == m {
require.Fail(t, "pointers should not be the same")
}
require.Equal(t, m, copy)
}
func TestWebSocketEventFromJson(t *testing.T) {
ev := WebSocketEventFromJson(strings.NewReader("junk"))
require.Nil(t, ev, "should not have parsed")
data := `{"event": "test", "data": {"key": "val"}, "seq": 45, "broadcast": {"user_id": "userid"}}`
ev = WebSocketEventFromJson(strings.NewReader(data))
require.NotNil(t, ev, "should have parsed")
require.Equal(t, ev.Event, "test")
require.Equal(t, ev.Sequence, int64(45))
require.Equal(t, ev.Data, map[string]interface{}{"key": "val"})
require.Equal(t, ev.Broadcast, &WebsocketBroadcast{UserId: "userid"})
}
func TestWebSocketResponse(t *testing.T) {
@@ -46,7 +100,7 @@ func TestWebSocketResponse(t *testing.T) {
func TestWebSocketEvent_PrecomputeJSON(t *testing.T) {
event := NewWebSocketEvent(WEBSOCKET_EVENT_POSTED, "foo", "bar", "baz", nil)
event.Sequence = 7
event = event.SetSequence(7)
before := event.ToJson()
event.PrecomputeJSON()
@@ -60,7 +114,7 @@ var stringSink string
func BenchmarkWebSocketEvent_ToJson(b *testing.B) {
event := NewWebSocketEvent(WEBSOCKET_EVENT_POSTED, "foo", "bar", "baz", nil)
for i := 0; i < 100; i++ {
event.Data[NewId()] = NewId()
event.GetData()[NewId()] = NewId()
}
b.Run("SerializedNTimes", func(b *testing.B) {