MM-61130: Use a channelMember map at web_hub level (#28810)

Tests at very high scale indicates that the iteration
of all connections during websocket broadcast starts
to become a bottleneck.

To optimize this, we move the channelMember cache from
inside web_conn.go to the hubConnectionIndex.

This involves adding a new map keyed by the channelID
and containing all webConns where the user is a member
of that channel. Subsequently, a new method needed to
be added to invalidate the cache which previously
used to happen in web_conn.

And as a last step, we remove the cache from web_conn
to reduce SQL queries to the DB.

https://mattermost.atlassian.net/browse/MM-61130

```release-note
NONE
```
Этот коммит содержится в:
Agniva De Sarker
2024-11-08 09:57:54 +05:30
коммит произвёл GitHub
родитель 37d97e8024
Коммит bd8774bdce
11 изменённых файлов: 564 добавлений и 329 удалений

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

@@ -4,6 +4,7 @@
package platform
import (
"fmt"
"hash/maphash"
"runtime"
"runtime/debug"
@@ -14,6 +15,7 @@ import (
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/mlog"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
const (
@@ -45,6 +47,11 @@ type webConnSessionMessage struct {
isRegistered chan bool
}
type webConnRegisterMessage struct {
conn *WebConn
err chan error
}
type webConnCheckMessage struct {
userID string
connectionID string
@@ -65,7 +72,7 @@ type Hub struct {
connectionCount int64
platform *PlatformService
connectionIndex int
register chan *WebConn
register chan *webConnRegisterMessage
unregister chan *WebConn
broadcast chan *model.WebSocketEvent
stop chan struct{}
@@ -84,7 +91,7 @@ type Hub struct {
func newWebHub(ps *PlatformService) *Hub {
return &Hub{
platform: ps,
register: make(chan *WebConn),
register: make(chan *webConnRegisterMessage),
unregister: make(chan *WebConn),
broadcast: make(chan *model.WebSocketEvent, broadcastQueueSize),
stop: make(chan struct{}),
@@ -150,14 +157,15 @@ func (ps *PlatformService) GetHubForUserId(userID string) *Hub {
}
// HubRegister registers a connection to a hub.
func (ps *PlatformService) HubRegister(webConn *WebConn) {
func (ps *PlatformService) HubRegister(webConn *WebConn) error {
hub := ps.GetHubForUserId(webConn.UserId)
if hub != nil {
if metrics := ps.metricsIFace; metrics != nil {
metrics.IncrementWebSocketBroadcastUsersRegistered(strconv.Itoa(hub.connectionIndex), 1)
}
hub.Register(webConn)
return hub.Register(webConn)
}
return nil
}
// HubUnregister unregisters a connection from a hub.
@@ -262,11 +270,17 @@ func (ps *PlatformService) WebConnCountForUser(userID string) int {
}
// Register registers a connection to the hub.
func (h *Hub) Register(webConn *WebConn) {
func (h *Hub) Register(webConn *WebConn) error {
wr := &webConnRegisterMessage{
conn: webConn,
err: make(chan error),
}
select {
case h.register <- webConn:
case h.register <- wr:
return <-wr.err
case <-h.stop:
}
return nil
}
// Unregister unregisters a connection from the hub.
@@ -389,7 +403,10 @@ func (h *Hub) Start() {
ticker := time.NewTicker(inactiveConnReaperInterval)
defer ticker.Stop()
connIndex := newHubConnectionIndex(inactiveConnReaperInterval)
connIndex := newHubConnectionIndex(inactiveConnReaperInterval,
h.platform.Store,
h.platform.logger,
)
for {
select {
@@ -423,22 +440,27 @@ func (h *Hub) Start() {
req.result <- connIndex.ForUserActiveCount(req.userID)
case <-ticker.C:
connIndex.RemoveInactiveConnections()
case webConn := <-h.register:
case webConnReg := <-h.register:
// Mark the current one as active.
// There is no need to check if it was inactive or not,
// we will anyways need to make it active.
webConn.Active.Store(true)
webConnReg.conn.Active.Store(true)
connIndex.Add(webConn)
err := connIndex.Add(webConnReg.conn)
if err != nil {
webConnReg.err <- err
continue
}
atomic.StoreInt64(&h.connectionCount, int64(connIndex.AllActive()))
if webConn.IsAuthenticated() && webConn.reuseCount == 0 {
if webConnReg.conn.IsAuthenticated() && webConnReg.conn.reuseCount == 0 {
// The hello message should only be sent when the reuseCount is 0.
// i.e in server restart, or long timeout, or fresh connection case.
// In case of seq number not found in dead queue, it is handled by
// the webconn write pump.
webConn.send <- webConn.createHelloMessage()
webConnReg.conn.send <- webConnReg.conn.createHelloMessage()
}
webConnReg.err <- nil
case webConn := <-h.unregister:
// If already removed (via queue full), then removing again becomes a noop.
// But if not removed, mark inactive.
@@ -497,6 +519,13 @@ func (h *Hub) Start() {
for _, webConn := range connIndex.ForUser(userID) {
webConn.InvalidateCache()
}
err := connIndex.InvalidateCMCacheForUser(userID)
if err != nil {
h.platform.Log().Error("Error while invalidating channel member cache", mlog.String("user_id", userID), mlog.Err(err))
for _, webConn := range connIndex.ForUser(userID) {
closeAndRemoveConn(connIndex, webConn)
}
}
case activity := <-h.activity:
for _, webConn := range connIndex.ForUser(activity.userID) {
if !webConn.Active.Load() {
@@ -519,8 +548,7 @@ func (h *Hub) Start() {
mlog.String("user_id", directMsg.conn.UserId),
mlog.String("conn_id", directMsg.conn.GetConnectionID()))
}
close(directMsg.conn.send)
connIndex.Remove(directMsg.conn)
closeAndRemoveConn(connIndex, directMsg.conn)
}
case msg := <-h.broadcast:
if metrics := h.platform.metricsIFace; metrics != nil {
@@ -546,27 +574,29 @@ func (h *Hub) Start() {
mlog.String("user_id", webConn.UserId),
mlog.String("conn_id", webConn.GetConnectionID()))
}
close(webConn.send)
connIndex.Remove(webConn)
closeAndRemoveConn(connIndex, webConn)
}
}
}
var targetConns []*WebConn
if connID := msg.GetBroadcast().ConnectionId; connID != "" {
if webConn := connIndex.ForConnection(connID); webConn != nil {
broadcast(webConn)
continue
targetConns = append(targetConns, webConn)
}
} else if msg.GetBroadcast().UserId != "" {
candidates := connIndex.ForUser(msg.GetBroadcast().UserId)
for _, webConn := range candidates {
} else if userID := msg.GetBroadcast().UserId; userID != "" {
targetConns = connIndex.ForUser(userID)
} else if channelID := msg.GetBroadcast().ChannelId; channelID != "" {
targetConns = connIndex.ForChannel(channelID)
}
if targetConns != nil {
for _, webConn := range targetConns {
broadcast(webConn)
}
continue
}
candidates := connIndex.All()
for webConn := range candidates {
for webConn := range connIndex.All() {
broadcast(webConn)
}
case <-h.stop:
@@ -616,14 +646,24 @@ func areAllInactive(conns []*WebConn) bool {
return true
}
// closeAndRemoveConn closes the send channel which will close the
// websocket connection, and then it removes the webConn from the conn index.
func closeAndRemoveConn(connIndex *hubConnectionIndex, conn *WebConn) {
close(conn.send)
connIndex.Remove(conn)
}
// hubConnectionIndex provides fast addition, removal, and iteration of web connections.
// It requires 3 functionalities which need to be very fast:
// It requires 4 functionalities which need to be very fast:
// - check if a connection exists or not.
// - get all connections for a given userID.
// - get all connections for a given channelID.
// - get all connections.
type hubConnectionIndex struct {
// byUserId stores the list of connections for a given userID
byUserId map[string][]*WebConn
// byChannelID stores the list of connections for a given channelID.
byChannelID map[string][]*WebConn
// byConnection serves the dual purpose of storing the index of the webconn
// in the value of byUserId map, and also to get all connections.
byConnection map[*WebConn]int
@@ -631,21 +671,39 @@ type hubConnectionIndex struct {
// staleThreshold is the limit beyond which inactive connections
// will be deleted.
staleThreshold time.Duration
store store.Store
logger mlog.LoggerIFace
}
func newHubConnectionIndex(interval time.Duration) *hubConnectionIndex {
func newHubConnectionIndex(interval time.Duration,
store store.Store,
logger mlog.LoggerIFace,
) *hubConnectionIndex {
return &hubConnectionIndex{
byUserId: make(map[string][]*WebConn),
byChannelID: make(map[string][]*WebConn),
byConnection: make(map[*WebConn]int),
byConnectionId: make(map[string]*WebConn),
staleThreshold: interval,
store: store,
logger: logger,
}
}
func (i *hubConnectionIndex) Add(wc *WebConn) {
func (i *hubConnectionIndex) Add(wc *WebConn) error {
cm, err := i.store.Channel().GetAllChannelMembersForUser(request.EmptyContext(i.logger), wc.UserId, false, false)
if err != nil {
return fmt.Errorf("error getChannelMembersForUser: %v", err)
}
for chID := range cm {
i.byChannelID[chID] = append(i.byChannelID[chID], wc)
}
i.byUserId[wc.UserId] = append(i.byUserId[wc.UserId], wc)
i.byConnection[wc] = len(i.byUserId[wc.UserId]) - 1
i.byConnectionId[wc.GetConnectionID()] = wc
return nil
}
func (i *hubConnectionIndex) Remove(wc *WebConn) {
@@ -654,21 +712,67 @@ func (i *hubConnectionIndex) Remove(wc *WebConn) {
return
}
// Remove the wc from i.byUserId
// get the conn slice.
userConnections := i.byUserId[wc.UserId]
// get the last connection.
last := userConnections[len(userConnections)-1]
// set the slot that we are trying to remove to be the last connection.
// https://go.dev/wiki/SliceTricks#delete-without-preserving-order
userConnections[userConnIndex] = last
// remove the last connection pointer from slice.
userConnections[len(userConnections)-1] = nil
// remove the last connection from the slice.
i.byUserId[wc.UserId] = userConnections[:len(userConnections)-1]
// set the index of the connection that was moved to the new index.
i.byConnection[last] = userConnIndex
connectionID := wc.GetConnectionID()
// Remove webconns from i.byChannelID
// This has O(n) complexity. We are trading off speed while removing
// a connection, to improve broadcasting a message.
for chID, webConns := range i.byChannelID {
// https://go.dev/wiki/SliceTricks#filtering-without-allocating
filtered := webConns[:0]
for _, conn := range webConns {
if conn.GetConnectionID() != connectionID {
filtered = append(filtered, conn)
}
}
for i := len(filtered); i < len(webConns); i++ {
webConns[i] = nil
}
i.byChannelID[chID] = filtered
}
delete(i.byConnection, wc)
delete(i.byConnectionId, wc.GetConnectionID())
delete(i.byConnectionId, connectionID)
}
func (i *hubConnectionIndex) InvalidateCMCacheForUser(userID string) error {
// We make this query first to fail fast in case of an error.
cm, err := i.store.Channel().GetAllChannelMembersForUser(request.EmptyContext(i.logger), userID, false, false)
if err != nil {
return err
}
// Clear out all user entries which belong to channels.
for chID, webConns := range i.byChannelID {
// https://go.dev/wiki/SliceTricks#filtering-without-allocating
filtered := webConns[:0]
for _, conn := range webConns {
if conn.UserId != userID {
filtered = append(filtered, conn)
}
}
for i := len(filtered); i < len(webConns); i++ {
webConns[i] = nil
}
i.byChannelID[chID] = filtered
}
// re-populate the cache
for chID := range cm {
i.byChannelID[chID] = append(i.byChannelID[chID], i.ForUser(userID)...)
}
return nil
}
func (i *hubConnectionIndex) Has(wc *WebConn) bool {
@@ -691,6 +795,16 @@ func (i *hubConnectionIndex) ForUser(id string) []*WebConn {
return conns
}
// ForChannel returns all connections for a channelID.
func (i *hubConnectionIndex) ForChannel(channelID string) []*WebConn {
// Note: this is expensive because usually there will be
// more than 1 member for a channel, and broadcasting
// is a hot path, but worth it.
conns := make([]*WebConn, len(i.byChannelID[channelID]))
copy(conns, i.byChannelID[channelID])
return conns
}
// ForUserActiveCount returns the number of active connections for a userID
func (i *hubConnectionIndex) ForUserActiveCount(id string) int {
cnt := 0