Files
mostlymatter/app/web_conn_test.go
Jesse Hallam c1a27d8360 MM-39524: fix WebConn caching stale channel members (#18840)
When configured in a master/slave database environment, a read replica can sometimes return stale data to a `WebConn`, resulting in the user missing out on websocket events targetting that channel until the `WebConn` cache expires.

This is most easily reproducible by using Playbooks on community and starting a new run. The owner, or any automatically invited participants, typically find the websocket events dropped in that channel for up to 30 minutes, even after multiple page refreshes.

I've reproduced this locally, and while I've extended the unit tests, I note that they don't actually exercise this case given the need for a dedicated slave database during unit tests.

Fixes: https://mattermost.atlassian.net/browse/MM-39524
2021-10-28 11:08:09 -03:00

352 строки
11 KiB
Go

// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
// See LICENSE.txt for license information.
package app
import (
"bytes"
"net"
"net/http"
"net/http/httptest"
"testing"
"github.com/gorilla/websocket"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost-server/v6/model"
"github.com/mattermost/mattermost-server/v6/shared/i18n"
)
func TestWebConnShouldSendEvent(t *testing.T) {
th := Setup(t).InitBasic()
defer th.TearDown()
session, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser.Id, Roles: th.BasicUser.GetRawRoles(), TeamMembers: []*model.TeamMember{
{
UserId: th.BasicUser.Id,
TeamId: th.BasicTeam.Id,
Roles: model.TeamUserRoleId,
},
}})
require.Nil(t, err)
basicUserWc := &WebConn{
App: th.App,
UserId: th.BasicUser.Id,
T: i18n.T,
}
basicUserWc.SetSession(session)
basicUserWc.SetSessionToken(session.Token)
basicUserWc.SetSessionExpiresAt(session.ExpiresAt)
session2, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser2.Id, Roles: th.BasicUser2.GetRawRoles(), TeamMembers: []*model.TeamMember{
{
UserId: th.BasicUser2.Id,
TeamId: th.BasicTeam.Id,
Roles: model.TeamAdminRoleId,
},
}})
require.Nil(t, err)
basicUser2Wc := &WebConn{
App: th.App,
UserId: th.BasicUser2.Id,
T: i18n.T,
}
basicUser2Wc.SetSession(session2)
basicUser2Wc.SetSessionToken(session2.Token)
basicUser2Wc.SetSessionExpiresAt(session2.ExpiresAt)
session3, err := th.App.CreateSession(&model.Session{UserId: th.SystemAdminUser.Id, Roles: th.SystemAdminUser.GetRawRoles()})
require.Nil(t, err)
adminUserWc := &WebConn{
App: th.App,
UserId: th.SystemAdminUser.Id,
T: i18n.T,
}
adminUserWc.SetSession(session3)
adminUserWc.SetSessionToken(session3.Token)
adminUserWc.SetSessionExpiresAt(session3.ExpiresAt)
// By default, only BasicUser and BasicUser2 get added to the BasicTeam.
th.LinkUserToTeam(th.SystemAdminUser, th.BasicTeam)
// Create another channel with just BasicUser (implicitly) and SystemAdminUser to test channel broadcast
channel2 := th.CreateChannel(th.BasicTeam)
th.AddUserToChannel(th.SystemAdminUser, channel2)
cases := []struct {
Description string
Broadcast *model.WebsocketBroadcast
User1Expected bool
User2Expected bool
AdminExpected bool
}{
{"should send to all", &model.WebsocketBroadcast{}, true, true, true},
{"should only send to basic user", &model.WebsocketBroadcast{UserId: th.BasicUser.Id}, true, false, false},
{"should omit basic user 2", &model.WebsocketBroadcast{OmitUsers: map[string]bool{th.BasicUser2.Id: true}}, true, false, true},
{"should only send to admin", &model.WebsocketBroadcast{ContainsSensitiveData: true}, false, false, true},
{"should only send to non-admins", &model.WebsocketBroadcast{ContainsSanitizedData: true}, true, true, false},
{"should send to nobody", &model.WebsocketBroadcast{ContainsSensitiveData: true, ContainsSanitizedData: true}, false, false, false},
// needs more cases to get full coverage
}
event := model.NewWebSocketEvent("some_event", "", "", "", nil)
for _, c := range cases {
t.Run(c.Description, func(t *testing.T) {
event = event.SetBroadcast(c.Broadcast)
if c.User1Expected {
assert.True(t, basicUserWc.shouldSendEvent(event), "expected user 1")
} else {
assert.False(t, basicUserWc.shouldSendEvent(event), "did not expect user 1")
}
if c.User2Expected {
assert.True(t, basicUser2Wc.shouldSendEvent(event), "expected user 2")
} else {
assert.False(t, basicUser2Wc.shouldSendEvent(event), "did not expect user 2")
}
if c.AdminExpected {
assert.True(t, adminUserWc.shouldSendEvent(event), "expected admin")
} else {
assert.False(t, adminUserWc.shouldSendEvent(event), "did not expect admin")
}
})
}
t.Run("should send to basic user in basic channel", func(t *testing.T) {
event = event.SetBroadcast(&model.WebsocketBroadcast{ChannelId: th.BasicChannel.Id})
assert.True(t, basicUserWc.shouldSendEvent(event), "expected user 1")
assert.False(t, basicUser2Wc.shouldSendEvent(event), "did not expect user 2")
assert.False(t, adminUserWc.shouldSendEvent(event), "did not expect admin")
})
t.Run("should send to basic user and admin in channel2", func(t *testing.T) {
event = event.SetBroadcast(&model.WebsocketBroadcast{ChannelId: channel2.Id})
assert.True(t, basicUserWc.shouldSendEvent(event), "expected user 1")
assert.False(t, basicUser2Wc.shouldSendEvent(event), "did not expect user 2")
assert.True(t, adminUserWc.shouldSendEvent(event), "expected admin")
})
t.Run("channel member cache invalidated after user added to channel", func(t *testing.T) {
th.AddUserToChannel(th.BasicUser2, channel2)
basicUser2Wc.InvalidateCache()
event = event.SetBroadcast(&model.WebsocketBroadcast{ChannelId: channel2.Id})
assert.True(t, basicUserWc.shouldSendEvent(event), "expected user 1")
assert.True(t, basicUser2Wc.shouldSendEvent(event), "expected user 2")
assert.True(t, adminUserWc.shouldSendEvent(event), "expected admin")
})
event2 := model.NewWebSocketEvent(model.WebsocketEventUpdateTeam, th.BasicTeam.Id, "", "", nil)
assert.True(t, basicUserWc.shouldSendEvent(event2))
assert.True(t, basicUser2Wc.shouldSendEvent(event2))
event3 := model.NewWebSocketEvent(model.WebsocketEventUpdateTeam, "wrongId", "", "", nil)
assert.False(t, basicUserWc.shouldSendEvent(event3))
}
func TestWebConnAddDeadQueue(t *testing.T) {
th := Setup(t)
defer th.TearDown()
th.App.UpdateConfig(func(cfg *model.Config) { *cfg.ServiceSettings.EnableReliableWebSockets = true })
wc := th.App.NewWebConn(&WebConnConfig{
WebSocket: &websocket.Conn{},
})
for i := 0; i < 2; i++ {
msg := &model.WebSocketEvent{}
msg = msg.SetSequence(int64(i))
wc.addToDeadQueue(msg)
}
for i := 0; i < 2; i++ {
assert.Equal(t, int64(i), wc.deadQueue[i].GetSequence())
}
// Should push out the first two elements
for i := 0; i < deadQueueSize; i++ {
msg := &model.WebSocketEvent{}
msg = msg.SetSequence(int64(i + 2))
wc.addToDeadQueue(msg)
}
for i := 0; i < deadQueueSize; i++ {
assert.Equal(t, int64(i+2), wc.deadQueue[(i+2)%deadQueueSize].GetSequence())
}
}
func TestWebConnIsInDeadQueue(t *testing.T) {
th := Setup(t)
defer th.TearDown()
th.App.UpdateConfig(func(cfg *model.Config) {
*cfg.ServiceSettings.EnableReliableWebSockets = true
})
wc := th.App.NewWebConn(&WebConnConfig{
WebSocket: &websocket.Conn{},
})
var i int
for ; i < 2; i++ {
msg := &model.WebSocketEvent{}
msg = msg.SetSequence(int64(i))
wc.addToDeadQueue(msg)
}
wc.Sequence = int64(0)
ok, ind := wc.isInDeadQueue(wc.Sequence)
assert.True(t, ok)
assert.Equal(t, 0, ind)
assert.True(t, wc.hasMsgLoss())
wc.Sequence = int64(1)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.True(t, ok)
assert.Equal(t, 1, ind)
assert.True(t, wc.hasMsgLoss())
wc.Sequence = int64(2)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.False(t, ok)
assert.Equal(t, 0, ind)
assert.False(t, wc.hasMsgLoss())
for ; i < deadQueueSize+2; i++ {
msg := &model.WebSocketEvent{}
msg = msg.SetSequence(int64(i))
wc.addToDeadQueue(msg)
}
wc.Sequence = int64(129)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.True(t, ok)
assert.Equal(t, 1, ind)
wc.Sequence = int64(128)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.True(t, ok)
assert.Equal(t, 0, ind)
wc.Sequence = int64(2)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.True(t, ok)
assert.Equal(t, 2, ind)
assert.True(t, wc.hasMsgLoss())
wc.Sequence = int64(0)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.False(t, ok)
assert.Equal(t, 0, ind)
wc.Sequence = int64(130)
ok, ind = wc.isInDeadQueue(wc.Sequence)
assert.False(t, ok)
assert.Equal(t, 0, ind)
assert.False(t, wc.hasMsgLoss())
}
func TestWebConnDrainDeadQueue(t *testing.T) {
th := Setup(t)
defer th.TearDown()
th.App.UpdateConfig(func(cfg *model.Config) {
*cfg.ServiceSettings.EnableReliableWebSockets = true
})
var dialConn = func(t *testing.T, a *App, addr net.Addr) *WebConn {
d := websocket.Dialer{}
c, _, err := d.Dial("ws://"+addr.String()+"/ws", nil)
require.NoError(t, err)
cfg := &WebConnConfig{
WebSocket: c,
}
return a.NewWebConn(cfg)
}
t.Run("Empty Queue", func(t *testing.T) {
var handler = func(t *testing.T) http.HandlerFunc {
return func(w http.ResponseWriter, req *http.Request) {
upgrader := &websocket.Upgrader{}
conn, err := upgrader.Upgrade(w, req, nil)
cnt := 0
for err == nil {
_, _, err = conn.ReadMessage()
cnt++
}
assert.Equal(t, 1, cnt)
if _, ok := err.(*websocket.CloseError); !ok {
require.NoError(t, err)
}
}
}
s := httptest.NewServer(handler(t))
defer s.Close()
wc := dialConn(t, th.App, s.Listener.Addr())
defer wc.WebSocket.Close()
wc.clearDeadQueue()
err := wc.drainDeadQueue(0)
require.NoError(t, err)
})
var handler = func(t *testing.T, seqNum int64, limit int) http.HandlerFunc {
return func(w http.ResponseWriter, req *http.Request) {
upgrader := &websocket.Upgrader{}
conn, err := upgrader.Upgrade(w, req, nil)
var buf []byte
i := seqNum
for err == nil {
_, buf, err = conn.ReadMessage()
if err != nil && len(buf) > 0 {
ev, jsonErr := model.WebSocketEventFromJSON(bytes.NewReader(buf))
require.NoError(t, jsonErr)
require.LessOrEqual(t, int(i), limit)
assert.Equal(t, i, ev.GetSequence())
i++
}
}
if _, ok := err.(*websocket.CloseError); !ok {
require.NoError(t, err)
}
}
}
run := func(seqNum int64, limit int) {
s := httptest.NewServer(handler(t, seqNum, limit))
defer s.Close()
wc := dialConn(t, th.App, s.Listener.Addr())
defer wc.WebSocket.Close()
for i := 0; i < limit; i++ {
msg := model.NewWebSocketEvent("", "", "", "", map[string]bool{})
msg = msg.SetSequence(int64(i))
wc.addToDeadQueue(msg)
}
wc.Sequence = seqNum
ok, index := wc.isInDeadQueue(wc.Sequence)
require.True(t, ok)
err := wc.drainDeadQueue(index)
require.NoError(t, err)
}
t.Run("Half-full Queue", func(t *testing.T) {
t.Run("Middle", func(t *testing.T) { run(int64(2), 10) })
t.Run("Beginning", func(t *testing.T) { run(int64(0), 10) })
t.Run("End", func(t *testing.T) { run(int64(9), 10) })
t.Run("Full", func(t *testing.T) { run(int64(deadQueueSize-1), deadQueueSize) })
})
t.Run("Cycled Queue", func(t *testing.T) {
t.Run("First un-overwritten", func(t *testing.T) { run(int64(10), deadQueueSize+10) })
t.Run("End", func(t *testing.T) { run(int64(127), deadQueueSize+10) })
t.Run("Cycled End", func(t *testing.T) { run(int64(137), deadQueueSize+10) })
t.Run("Overwritten First", func(t *testing.T) { run(int64(128), deadQueueSize+10) })
})
}