diff --git a/app/web_conn.go b/app/web_conn.go index c1d4b918d1..4bd511babc 100644 --- a/app/web_conn.go +++ b/app/web_conn.go @@ -11,6 +11,7 @@ import ( "net" "net/http" "strconv" + "strings" "sync" "sync/atomic" "time" @@ -18,6 +19,7 @@ import ( "github.com/gorilla/websocket" "github.com/mattermost/mattermost-server/v6/model" + "github.com/mattermost/mattermost-server/v6/plugin" "github.com/mattermost/mattermost-server/v6/shared/i18n" "github.com/mattermost/mattermost-server/v6/shared/mlog" ) @@ -40,6 +42,14 @@ const ( reconnectLossless = "lossless" ) +const websocketMessagePluginPrefix = "custom_" + +type pluginWSPostedHook struct { + connectionID string + userID string + req *model.WebSocketRequest +} + type WebConnConfig struct { WebSocket *websocket.Conn Session model.Session @@ -90,6 +100,7 @@ type WebConn struct { connectionID atomic.Value endWritePump chan struct{} pumpFinished chan struct{} + pluginPosted chan pluginWSPostedHook } // CheckConnResult indicates whether a connectionID was present in the hub or not. @@ -179,6 +190,7 @@ func (a *App) NewWebConn(cfg *WebConnConfig) *WebConn { active: cfg.Active, endWritePump: make(chan struct{}), pumpFinished: make(chan struct{}), + pluginPosted: make(chan pluginWSPostedHook, 10), } wc.SetSession(&cfg.Session) @@ -186,9 +198,31 @@ func (a *App) NewWebConn(cfg *WebConnConfig) *WebConn { wc.SetSessionExpiresAt(cfg.Session.ExpiresAt) wc.SetConnectionID(cfg.ConnectionID) + if pluginsEnvironment := wc.App.GetPluginsEnvironment(); pluginsEnvironment != nil { + wc.App.Srv().Go(func() { + pluginsEnvironment.RunMultiPluginHook(func(hooks plugin.Hooks) bool { + hooks.OnWebSocketConnect(wc.GetConnectionID(), wc.UserId) + return true + }, plugin.OnWebSocketConnectID) + }) + } + return wc } +func (wc *WebConn) pluginPostedConsumer(wg *sync.WaitGroup) { + defer wg.Done() + + for msg := range wc.pluginPosted { + if pluginsEnvironment := wc.App.GetPluginsEnvironment(); pluginsEnvironment != nil { + pluginsEnvironment.RunMultiPluginHook(func(hooks plugin.Hooks) bool { + hooks.WebSocketMessageHasBeenPosted(msg.connectionID, msg.userID, msg.req) + return true + }, plugin.WebSocketMessageHasBeenPostedID) + } + } +} + // Close closes the WebConn. func (wc *WebConn) Close() { wc.WebSocket.Close() @@ -261,11 +295,25 @@ func (wc *WebConn) Pump() { defer wg.Done() wc.writePump() }() + + wg.Add(1) + go wc.pluginPostedConsumer(&wg) + wc.readPump() close(wc.endWritePump) + close(wc.pluginPosted) wg.Wait() wc.App.HubUnregister(wc) close(wc.pumpFinished) + + if pluginsEnvironment := wc.App.GetPluginsEnvironment(); pluginsEnvironment != nil { + wc.App.Srv().Go(func() { + pluginsEnvironment.RunMultiPluginHook(func(hooks plugin.Hooks) bool { + hooks.OnWebSocketDisconnect(wc.GetConnectionID(), wc.UserId) + return true + }, plugin.OnWebSocketDisconnectID) + }) + } } func (wc *WebConn) readPump() { @@ -290,7 +338,20 @@ func (wc *WebConn) readPump() { wc.logSocketErr("websocket.read", err) return } - wc.App.Srv().WebSocketRouter.ServeWebSocket(wc, &req) + + // Messages which actions are prefixed with the plugin prefix + // should only be dispatched to the plugins + if !strings.HasPrefix(req.Action, websocketMessagePluginPrefix) { + wc.App.Srv().WebSocketRouter.ServeWebSocket(wc, &req) + } + + clonedReq, err := req.Clone() + if err != nil { + wc.logSocketErr("websocket.cloneRequest", err) + continue + } + + wc.pluginPosted <- pluginWSPostedHook{wc.GetConnectionID(), wc.UserId, clonedReq} } } diff --git a/model/websocket_request.go b/model/websocket_request.go index 2468e21e2b..cc5d417f01 100644 --- a/model/websocket_request.go +++ b/model/websocket_request.go @@ -28,6 +28,19 @@ func (o *WebSocketRequest) ToJson() string { return string(b) } +func (o *WebSocketRequest) Clone() (*WebSocketRequest, error) { + buf, err := json.Marshal(o) + if err != nil { + return nil, err + } + var ret WebSocketRequest + err = json.Unmarshal(buf, &ret) + if err != nil { + return nil, err + } + return &ret, nil +} + func WebSocketRequestFromJson(data io.Reader) *WebSocketRequest { var o *WebSocketRequest json.NewDecoder(data).Decode(&o) diff --git a/plugin/client_rpc_generated.go b/plugin/client_rpc_generated.go index e6ba2e0d29..8901d59aed 100644 --- a/plugin/client_rpc_generated.go +++ b/plugin/client_rpc_generated.go @@ -566,6 +566,109 @@ func (s *hooksRPCServer) OnPluginClusterEvent(args *Z_OnPluginClusterEventArgs, return nil } +func init() { + hookNameToId["OnWebSocketConnect"] = OnWebSocketConnectID +} + +type Z_OnWebSocketConnectArgs struct { + A string + B string +} + +type Z_OnWebSocketConnectReturns struct { +} + +func (g *hooksRPCClient) OnWebSocketConnect(webConnID, userID string) { + _args := &Z_OnWebSocketConnectArgs{webConnID, userID} + _returns := &Z_OnWebSocketConnectReturns{} + if g.implemented[OnWebSocketConnectID] { + if err := g.client.Call("Plugin.OnWebSocketConnect", _args, _returns); err != nil { + g.log.Error("RPC call OnWebSocketConnect to plugin failed.", mlog.Err(err)) + } + } + +} + +func (s *hooksRPCServer) OnWebSocketConnect(args *Z_OnWebSocketConnectArgs, returns *Z_OnWebSocketConnectReturns) error { + if hook, ok := s.impl.(interface { + OnWebSocketConnect(webConnID, userID string) + }); ok { + hook.OnWebSocketConnect(args.A, args.B) + } else { + return encodableError(fmt.Errorf("Hook OnWebSocketConnect called but not implemented.")) + } + return nil +} + +func init() { + hookNameToId["OnWebSocketDisconnect"] = OnWebSocketDisconnectID +} + +type Z_OnWebSocketDisconnectArgs struct { + A string + B string +} + +type Z_OnWebSocketDisconnectReturns struct { +} + +func (g *hooksRPCClient) OnWebSocketDisconnect(webConnID, userID string) { + _args := &Z_OnWebSocketDisconnectArgs{webConnID, userID} + _returns := &Z_OnWebSocketDisconnectReturns{} + if g.implemented[OnWebSocketDisconnectID] { + if err := g.client.Call("Plugin.OnWebSocketDisconnect", _args, _returns); err != nil { + g.log.Error("RPC call OnWebSocketDisconnect to plugin failed.", mlog.Err(err)) + } + } + +} + +func (s *hooksRPCServer) OnWebSocketDisconnect(args *Z_OnWebSocketDisconnectArgs, returns *Z_OnWebSocketDisconnectReturns) error { + if hook, ok := s.impl.(interface { + OnWebSocketDisconnect(webConnID, userID string) + }); ok { + hook.OnWebSocketDisconnect(args.A, args.B) + } else { + return encodableError(fmt.Errorf("Hook OnWebSocketDisconnect called but not implemented.")) + } + return nil +} + +func init() { + hookNameToId["WebSocketMessageHasBeenPosted"] = WebSocketMessageHasBeenPostedID +} + +type Z_WebSocketMessageHasBeenPostedArgs struct { + A string + B string + C *model.WebSocketRequest +} + +type Z_WebSocketMessageHasBeenPostedReturns struct { +} + +func (g *hooksRPCClient) WebSocketMessageHasBeenPosted(webConnID, userID string, req *model.WebSocketRequest) { + _args := &Z_WebSocketMessageHasBeenPostedArgs{webConnID, userID, req} + _returns := &Z_WebSocketMessageHasBeenPostedReturns{} + if g.implemented[WebSocketMessageHasBeenPostedID] { + if err := g.client.Call("Plugin.WebSocketMessageHasBeenPosted", _args, _returns); err != nil { + g.log.Error("RPC call WebSocketMessageHasBeenPosted to plugin failed.", mlog.Err(err)) + } + } + +} + +func (s *hooksRPCServer) WebSocketMessageHasBeenPosted(args *Z_WebSocketMessageHasBeenPostedArgs, returns *Z_WebSocketMessageHasBeenPostedReturns) error { + if hook, ok := s.impl.(interface { + WebSocketMessageHasBeenPosted(webConnID, userID string, req *model.WebSocketRequest) + }); ok { + hook.WebSocketMessageHasBeenPosted(args.A, args.B, args.C) + } else { + return encodableError(fmt.Errorf("Hook WebSocketMessageHasBeenPosted called but not implemented.")) + } + return nil +} + type Z_RegisterCommandArgs struct { A *model.Command } diff --git a/plugin/hooks.go b/plugin/hooks.go index a88f9e90f9..1a3fdd4ff7 100644 --- a/plugin/hooks.go +++ b/plugin/hooks.go @@ -15,28 +15,31 @@ import ( // Feel free to add more, but do not change existing assignments. Follow the naming convention of // ID as the autogenerated glue code depends on that. const ( - OnActivateID = 0 - OnDeactivateID = 1 - ServeHTTPID = 2 - OnConfigurationChangeID = 3 - ExecuteCommandID = 4 - MessageWillBePostedID = 5 - MessageWillBeUpdatedID = 6 - MessageHasBeenPostedID = 7 - MessageHasBeenUpdatedID = 8 - UserHasJoinedChannelID = 9 - UserHasLeftChannelID = 10 - UserHasJoinedTeamID = 11 - UserHasLeftTeamID = 12 - ChannelHasBeenCreatedID = 13 - FileWillBeUploadedID = 14 - UserWillLogInID = 15 - UserHasLoggedInID = 16 - UserHasBeenCreatedID = 17 - ReactionHasBeenAddedID = 18 - ReactionHasBeenRemovedID = 19 - OnPluginClusterEventID = 20 - TotalHooksID = iota + OnActivateID = 0 + OnDeactivateID = 1 + ServeHTTPID = 2 + OnConfigurationChangeID = 3 + ExecuteCommandID = 4 + MessageWillBePostedID = 5 + MessageWillBeUpdatedID = 6 + MessageHasBeenPostedID = 7 + MessageHasBeenUpdatedID = 8 + UserHasJoinedChannelID = 9 + UserHasLeftChannelID = 10 + UserHasJoinedTeamID = 11 + UserHasLeftTeamID = 12 + ChannelHasBeenCreatedID = 13 + FileWillBeUploadedID = 14 + UserWillLogInID = 15 + UserHasLoggedInID = 16 + UserHasBeenCreatedID = 17 + ReactionHasBeenAddedID = 18 + ReactionHasBeenRemovedID = 19 + OnPluginClusterEventID = 20 + OnWebSocketConnectID = 21 + OnWebSocketDisconnectID = 22 + WebSocketMessageHasBeenPostedID = 23 + TotalHooksID = iota ) const ( @@ -219,4 +222,25 @@ type Hooks interface { // // Minimum server version: 5.36 OnPluginClusterEvent(c *Context, ev model.PluginClusterEvent) + + // OnWebSocketConnect is invoked when a new websocket connection is opened. + // + // This is used to track which users have connections opened with the Mattermost + // websocket. + // + // Minimum server version: 6.0 + OnWebSocketConnect(webConnID, userID string) + + // OnWebSocketDisconnect is invoked when a websocket connection is closed. + // + // This is used to track which users have connections opened with the Mattermost + // websocket. + // + // Minimum server version: 6.0 + OnWebSocketDisconnect(webConnID, userID string) + + // WebSocketMessageHasBeenPosted is invoked when a websocket message is received. + // + // Minimum server version: 6.0 + WebSocketMessageHasBeenPosted(webConnID, userID string, req *model.WebSocketRequest) } diff --git a/plugin/hooks_timer_layer_generated.go b/plugin/hooks_timer_layer_generated.go index 9bac911ae7..1d757d65c9 100644 --- a/plugin/hooks_timer_layer_generated.go +++ b/plugin/hooks_timer_layer_generated.go @@ -168,3 +168,21 @@ func (hooks *hooksTimerLayer) OnPluginClusterEvent(c *Context, ev model.PluginCl hooks.hooksImpl.OnPluginClusterEvent(c, ev) hooks.recordTime(startTime, "OnPluginClusterEvent", true) } + +func (hooks *hooksTimerLayer) OnWebSocketConnect(webConnID, userID string) { + startTime := timePkg.Now() + hooks.hooksImpl.OnWebSocketConnect(webConnID, userID) + hooks.recordTime(startTime, "OnWebSocketConnect", true) +} + +func (hooks *hooksTimerLayer) OnWebSocketDisconnect(webConnID, userID string) { + startTime := timePkg.Now() + hooks.hooksImpl.OnWebSocketDisconnect(webConnID, userID) + hooks.recordTime(startTime, "OnWebSocketDisconnect", true) +} + +func (hooks *hooksTimerLayer) WebSocketMessageHasBeenPosted(webConnID, userID string, req *model.WebSocketRequest) { + startTime := timePkg.Now() + hooks.hooksImpl.WebSocketMessageHasBeenPosted(webConnID, userID, req) + hooks.recordTime(startTime, "WebSocketMessageHasBeenPosted", true) +}