Răsfoiți Sursa

update chat

Mike 4 luni în urmă
părinte
comite
e64f35f883

+ 1 - 1
cmd/server/factory.go

@@ -491,7 +491,7 @@ func TOC(deps Container) *toc.Server {
 			SessionRetriever:  deps.inMemorySessionManager,
 			RandIntn:          rand.Intn,
 		},
-		toc.NewIPRateLimiter(rate.Every(1*time.Minute), 10, 1*time.Minute),
+		toc.NewIPRateLimiter(rate.Every(1*time.Minute), 100000, 1*time.Minute),
 		deps.icbmSvc.RestoreWarningLevel,
 		deps.icbmSvc.UpdateWarnLevel,
 	)

+ 8 - 5
foodgroup/auth.go

@@ -81,8 +81,8 @@ type AuthService struct {
 // This method does not verify that the user and chat room exist because it
 // implicitly trusts the contents of the token signed by
 // {{OServiceService.ServiceRequest}}.
-func (s AuthService) RegisterChatSession(ctx context.Context, serverCookie state.ServerCookie) (*state.SessionInstance, error) {
-	sess, err := s.chatSessionRegistry.AddSession(ctx, serverCookie.ChatCookie, serverCookie.ScreenName)
+func (s AuthService) RegisterChatSession(ctx context.Context, serverCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error) {
+	sess, err := s.chatSessionRegistry.AddSession(ctx, serverCookie.ChatCookie, serverCookie.ScreenName, cfg)
 	if err != nil {
 		return nil, fmt.Errorf("AddSession: %w", err)
 	}
@@ -207,9 +207,12 @@ func (s AuthService) Signout(ctx context.Context, session *state.Session) {
 
 // SignoutChat removes user from chat room and notifies remaining participants
 // of their departure.
-func (s AuthService) SignoutChat(ctx context.Context, instance *state.SessionInstance) {
-	alertUserLeft(ctx, instance, s.chatMessageRelayer)
-	s.chatSessionRegistry.RemoveSession(instance)
+func (s AuthService) SignoutChat(ctx context.Context, sess *state.Session) {
+	instances := sess.Instances()
+	for _, instance := range instances {
+		alertUserLeft(ctx, instance, s.chatMessageRelayer)
+	}
+	s.chatSessionRegistry.RemoveSession(sess)
 }
 
 // BUCPChallenge processes a BUCP authentication challenge request. It

+ 4 - 4
foodgroup/auth_test.go

@@ -1745,7 +1745,7 @@ func TestAuthService_RegisterChatSession_HappyPath(t *testing.T) {
 
 	chatSessionRegistry := newMockChatSessionRegistry(t)
 	chatSessionRegistry.EXPECT().
-		AddSession(mock.Anything, serverCookie.ChatCookie, instance.DisplayScreenName()).
+		AddSession(mock.Anything, serverCookie.ChatCookie, instance.DisplayScreenName(), mock.Anything).
 		Return(instance, nil)
 
 	chatCookieBuf := &bytes.Buffer{}
@@ -1753,7 +1753,7 @@ func TestAuthService_RegisterChatSession_HappyPath(t *testing.T) {
 
 	svc := NewAuthService(config.Config{}, nil, nil, chatSessionRegistry, nil, nil, nil, nil, nil, wire.DefaultRateLimitClasses(), nil, slog.Default())
 
-	have, err := svc.RegisterChatSession(context.Background(), serverCookie)
+	have, err := svc.RegisterChatSession(context.Background(), serverCookie, nil)
 	assert.NoError(t, err)
 	assert.Equal(t, instance, have)
 }
@@ -2107,11 +2107,11 @@ func TestAuthService_SignoutChat(t *testing.T) {
 			sessionManager := newMockChatSessionRegistry(t)
 			for _, params := range tt.mockParams.removeSessionParams {
 				sessionManager.EXPECT().
-					RemoveSession(matchSession(params.screenName))
+					RemoveSession(matchUserSession(params.screenName))
 			}
 
 			svc := NewAuthService(config.Config{}, nil, nil, sessionManager, nil, nil, chatMessageRelayer, nil, nil, wire.DefaultRateLimitClasses(), nil, slog.Default())
-			svc.SignoutChat(context.Background(), tt.instance)
+			svc.SignoutChat(context.Background(), tt.instance.Session())
 		})
 	}
 }

+ 40 - 21
foodgroup/mock_chat_session_registry_test.go

@@ -39,8 +39,16 @@ func (_m *mockChatSessionRegistry) EXPECT() *mockChatSessionRegistry_Expecter {
 }
 
 // AddSession provides a mock function for the type mockChatSessionRegistry
-func (_mock *mockChatSessionRegistry) AddSession(ctx context.Context, chatCookie string, screenName state.DisplayScreenName) (*state.SessionInstance, error) {
-	ret := _mock.Called(ctx, chatCookie, screenName)
+func (_mock *mockChatSessionRegistry) AddSession(ctx context.Context, chatCookie string, screenName state.DisplayScreenName, cfg ...func(sess *state.Session)) (*state.SessionInstance, error) {
+	// func(sess *state.Session)
+	_va := make([]interface{}, len(cfg))
+	for _i := range cfg {
+		_va[_i] = cfg[_i]
+	}
+	var _ca []interface{}
+	_ca = append(_ca, ctx, chatCookie, screenName)
+	_ca = append(_ca, _va...)
+	ret := _mock.Called(_ca...)
 
 	if len(ret) == 0 {
 		panic("no return value specified for AddSession")
@@ -48,18 +56,18 @@ func (_mock *mockChatSessionRegistry) AddSession(ctx context.Context, chatCookie
 
 	var r0 *state.SessionInstance
 	var r1 error
-	if returnFunc, ok := ret.Get(0).(func(context.Context, string, state.DisplayScreenName) (*state.SessionInstance, error)); ok {
-		return returnFunc(ctx, chatCookie, screenName)
+	if returnFunc, ok := ret.Get(0).(func(context.Context, string, state.DisplayScreenName, ...func(sess *state.Session)) (*state.SessionInstance, error)); ok {
+		return returnFunc(ctx, chatCookie, screenName, cfg...)
 	}
-	if returnFunc, ok := ret.Get(0).(func(context.Context, string, state.DisplayScreenName) *state.SessionInstance); ok {
-		r0 = returnFunc(ctx, chatCookie, screenName)
+	if returnFunc, ok := ret.Get(0).(func(context.Context, string, state.DisplayScreenName, ...func(sess *state.Session)) *state.SessionInstance); ok {
+		r0 = returnFunc(ctx, chatCookie, screenName, cfg...)
 	} else {
 		if ret.Get(0) != nil {
 			r0 = ret.Get(0).(*state.SessionInstance)
 		}
 	}
-	if returnFunc, ok := ret.Get(1).(func(context.Context, string, state.DisplayScreenName) error); ok {
-		r1 = returnFunc(ctx, chatCookie, screenName)
+	if returnFunc, ok := ret.Get(1).(func(context.Context, string, state.DisplayScreenName, ...func(sess *state.Session)) error); ok {
+		r1 = returnFunc(ctx, chatCookie, screenName, cfg...)
 	} else {
 		r1 = ret.Error(1)
 	}
@@ -75,11 +83,13 @@ type mockChatSessionRegistry_AddSession_Call struct {
 //   - ctx context.Context
 //   - chatCookie string
 //   - screenName state.DisplayScreenName
-func (_e *mockChatSessionRegistry_Expecter) AddSession(ctx interface{}, chatCookie interface{}, screenName interface{}) *mockChatSessionRegistry_AddSession_Call {
-	return &mockChatSessionRegistry_AddSession_Call{Call: _e.mock.On("AddSession", ctx, chatCookie, screenName)}
+//   - cfg ...func(sess *state.Session)
+func (_e *mockChatSessionRegistry_Expecter) AddSession(ctx interface{}, chatCookie interface{}, screenName interface{}, cfg ...interface{}) *mockChatSessionRegistry_AddSession_Call {
+	return &mockChatSessionRegistry_AddSession_Call{Call: _e.mock.On("AddSession",
+		append([]interface{}{ctx, chatCookie, screenName}, cfg...)...)}
 }
 
-func (_c *mockChatSessionRegistry_AddSession_Call) Run(run func(ctx context.Context, chatCookie string, screenName state.DisplayScreenName)) *mockChatSessionRegistry_AddSession_Call {
+func (_c *mockChatSessionRegistry_AddSession_Call) Run(run func(ctx context.Context, chatCookie string, screenName state.DisplayScreenName, cfg ...func(sess *state.Session))) *mockChatSessionRegistry_AddSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
 		var arg0 context.Context
 		if args[0] != nil {
@@ -93,10 +103,19 @@ func (_c *mockChatSessionRegistry_AddSession_Call) Run(run func(ctx context.Cont
 		if args[2] != nil {
 			arg2 = args[2].(state.DisplayScreenName)
 		}
+		var arg3 []func(sess *state.Session)
+		variadicArgs := make([]func(sess *state.Session), len(args)-3)
+		for i, a := range args[3:] {
+			if a != nil {
+				variadicArgs[i] = a.(func(sess *state.Session))
+			}
+		}
+		arg3 = variadicArgs
 		run(
 			arg0,
 			arg1,
 			arg2,
+			arg3...,
 		)
 	})
 	return _c
@@ -107,14 +126,14 @@ func (_c *mockChatSessionRegistry_AddSession_Call) Return(sessionInstance *state
 	return _c
 }
 
-func (_c *mockChatSessionRegistry_AddSession_Call) RunAndReturn(run func(ctx context.Context, chatCookie string, screenName state.DisplayScreenName) (*state.SessionInstance, error)) *mockChatSessionRegistry_AddSession_Call {
+func (_c *mockChatSessionRegistry_AddSession_Call) RunAndReturn(run func(ctx context.Context, chatCookie string, screenName state.DisplayScreenName, cfg ...func(sess *state.Session)) (*state.SessionInstance, error)) *mockChatSessionRegistry_AddSession_Call {
 	_c.Call.Return(run)
 	return _c
 }
 
 // RemoveSession provides a mock function for the type mockChatSessionRegistry
-func (_mock *mockChatSessionRegistry) RemoveSession(instance *state.SessionInstance) {
-	_mock.Called(instance)
+func (_mock *mockChatSessionRegistry) RemoveSession(sess *state.Session) {
+	_mock.Called(sess)
 	return
 }
 
@@ -124,16 +143,16 @@ type mockChatSessionRegistry_RemoveSession_Call struct {
 }
 
 // RemoveSession is a helper method to define mock.On call
-//   - instance *state.SessionInstance
-func (_e *mockChatSessionRegistry_Expecter) RemoveSession(instance interface{}) *mockChatSessionRegistry_RemoveSession_Call {
-	return &mockChatSessionRegistry_RemoveSession_Call{Call: _e.mock.On("RemoveSession", instance)}
+//   - sess *state.Session
+func (_e *mockChatSessionRegistry_Expecter) RemoveSession(sess interface{}) *mockChatSessionRegistry_RemoveSession_Call {
+	return &mockChatSessionRegistry_RemoveSession_Call{Call: _e.mock.On("RemoveSession", sess)}
 }
 
-func (_c *mockChatSessionRegistry_RemoveSession_Call) Run(run func(instance *state.SessionInstance)) *mockChatSessionRegistry_RemoveSession_Call {
+func (_c *mockChatSessionRegistry_RemoveSession_Call) Run(run func(sess *state.Session)) *mockChatSessionRegistry_RemoveSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
-		var arg0 *state.SessionInstance
+		var arg0 *state.Session
 		if args[0] != nil {
-			arg0 = args[0].(*state.SessionInstance)
+			arg0 = args[0].(*state.Session)
 		}
 		run(
 			arg0,
@@ -147,7 +166,7 @@ func (_c *mockChatSessionRegistry_RemoveSession_Call) Return() *mockChatSessionR
 	return _c
 }
 
-func (_c *mockChatSessionRegistry_RemoveSession_Call) RunAndReturn(run func(instance *state.SessionInstance)) *mockChatSessionRegistry_RemoveSession_Call {
+func (_c *mockChatSessionRegistry_RemoveSession_Call) RunAndReturn(run func(sess *state.Session)) *mockChatSessionRegistry_RemoveSession_Call {
 	_c.Run(run)
 	return _c
 }

+ 2 - 2
foodgroup/types.go

@@ -170,10 +170,10 @@ type ChatSessionRegistry interface {
 	// param identifies the chat room to which screenName is added. It returns
 	// the newly created session instance registered in the chat session
 	// manager.
-	AddSession(ctx context.Context, chatCookie string, screenName state.DisplayScreenName) (*state.SessionInstance, error)
+	AddSession(ctx context.Context, chatCookie string, screenName state.DisplayScreenName, cfg ...func(sess *state.Session)) (*state.SessionInstance, error)
 
 	// RemoveSession removes a session from the chat session manager.
-	RemoveSession(instance *state.SessionInstance)
+	RemoveSession(sess *state.Session)
 }
 
 // ClientSideBuddyListManager defines operations for managing a user's buddy list,

+ 24 - 18
server/oscar/mock_auth_test.go

@@ -463,8 +463,8 @@ func (_c *mockAuthService_RegisterBOSSession_Call) RunAndReturn(run func(ctx con
 }
 
 // RegisterChatSession provides a mock function for the type mockAuthService
-func (_mock *mockAuthService) RegisterChatSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error) {
-	ret := _mock.Called(ctx, authCookie)
+func (_mock *mockAuthService) RegisterChatSession(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error) {
+	ret := _mock.Called(ctx, authCookie, cfg)
 
 	if len(ret) == 0 {
 		panic("no return value specified for RegisterChatSession")
@@ -472,18 +472,18 @@ func (_mock *mockAuthService) RegisterChatSession(ctx context.Context, authCooki
 
 	var r0 *state.SessionInstance
 	var r1 error
-	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie) (*state.SessionInstance, error)); ok {
-		return returnFunc(ctx, authCookie)
+	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie, func(sess *state.Session)) (*state.SessionInstance, error)); ok {
+		return returnFunc(ctx, authCookie, cfg)
 	}
-	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie) *state.SessionInstance); ok {
-		r0 = returnFunc(ctx, authCookie)
+	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie, func(sess *state.Session)) *state.SessionInstance); ok {
+		r0 = returnFunc(ctx, authCookie, cfg)
 	} else {
 		if ret.Get(0) != nil {
 			r0 = ret.Get(0).(*state.SessionInstance)
 		}
 	}
-	if returnFunc, ok := ret.Get(1).(func(context.Context, state.ServerCookie) error); ok {
-		r1 = returnFunc(ctx, authCookie)
+	if returnFunc, ok := ret.Get(1).(func(context.Context, state.ServerCookie, func(sess *state.Session)) error); ok {
+		r1 = returnFunc(ctx, authCookie, cfg)
 	} else {
 		r1 = ret.Error(1)
 	}
@@ -498,11 +498,12 @@ type mockAuthService_RegisterChatSession_Call struct {
 // RegisterChatSession is a helper method to define mock.On call
 //   - ctx context.Context
 //   - authCookie state.ServerCookie
-func (_e *mockAuthService_Expecter) RegisterChatSession(ctx interface{}, authCookie interface{}) *mockAuthService_RegisterChatSession_Call {
-	return &mockAuthService_RegisterChatSession_Call{Call: _e.mock.On("RegisterChatSession", ctx, authCookie)}
+//   - cfg func(sess *state.Session)
+func (_e *mockAuthService_Expecter) RegisterChatSession(ctx interface{}, authCookie interface{}, cfg interface{}) *mockAuthService_RegisterChatSession_Call {
+	return &mockAuthService_RegisterChatSession_Call{Call: _e.mock.On("RegisterChatSession", ctx, authCookie, cfg)}
 }
 
-func (_c *mockAuthService_RegisterChatSession_Call) Run(run func(ctx context.Context, authCookie state.ServerCookie)) *mockAuthService_RegisterChatSession_Call {
+func (_c *mockAuthService_RegisterChatSession_Call) Run(run func(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session))) *mockAuthService_RegisterChatSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
 		var arg0 context.Context
 		if args[0] != nil {
@@ -512,9 +513,14 @@ func (_c *mockAuthService_RegisterChatSession_Call) Run(run func(ctx context.Con
 		if args[1] != nil {
 			arg1 = args[1].(state.ServerCookie)
 		}
+		var arg2 func(sess *state.Session)
+		if args[2] != nil {
+			arg2 = args[2].(func(sess *state.Session))
+		}
 		run(
 			arg0,
 			arg1,
+			arg2,
 		)
 	})
 	return _c
@@ -525,7 +531,7 @@ func (_c *mockAuthService_RegisterChatSession_Call) Return(sessionInstance *stat
 	return _c
 }
 
-func (_c *mockAuthService_RegisterChatSession_Call) RunAndReturn(run func(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)) *mockAuthService_RegisterChatSession_Call {
+func (_c *mockAuthService_RegisterChatSession_Call) RunAndReturn(run func(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error)) *mockAuthService_RegisterChatSession_Call {
 	_c.Call.Return(run)
 	return _c
 }
@@ -645,7 +651,7 @@ func (_c *mockAuthService_Signout_Call) RunAndReturn(run func(ctx context.Contex
 }
 
 // SignoutChat provides a mock function for the type mockAuthService
-func (_mock *mockAuthService) SignoutChat(ctx context.Context, instance *state.SessionInstance) {
+func (_mock *mockAuthService) SignoutChat(ctx context.Context, instance *state.Session) {
 	_mock.Called(ctx, instance)
 	return
 }
@@ -657,20 +663,20 @@ type mockAuthService_SignoutChat_Call struct {
 
 // SignoutChat is a helper method to define mock.On call
 //   - ctx context.Context
-//   - instance *state.SessionInstance
+//   - instance *state.Session
 func (_e *mockAuthService_Expecter) SignoutChat(ctx interface{}, instance interface{}) *mockAuthService_SignoutChat_Call {
 	return &mockAuthService_SignoutChat_Call{Call: _e.mock.On("SignoutChat", ctx, instance)}
 }
 
-func (_c *mockAuthService_SignoutChat_Call) Run(run func(ctx context.Context, instance *state.SessionInstance)) *mockAuthService_SignoutChat_Call {
+func (_c *mockAuthService_SignoutChat_Call) Run(run func(ctx context.Context, instance *state.Session)) *mockAuthService_SignoutChat_Call {
 	_c.Call.Run(func(args mock.Arguments) {
 		var arg0 context.Context
 		if args[0] != nil {
 			arg0 = args[0].(context.Context)
 		}
-		var arg1 *state.SessionInstance
+		var arg1 *state.Session
 		if args[1] != nil {
-			arg1 = args[1].(*state.SessionInstance)
+			arg1 = args[1].(*state.Session)
 		}
 		run(
 			arg0,
@@ -685,7 +691,7 @@ func (_c *mockAuthService_SignoutChat_Call) Return() *mockAuthService_SignoutCha
 	return _c
 }
 
-func (_c *mockAuthService_SignoutChat_Call) RunAndReturn(run func(ctx context.Context, instance *state.SessionInstance)) *mockAuthService_SignoutChat_Call {
+func (_c *mockAuthService_SignoutChat_Call) RunAndReturn(run func(ctx context.Context, instance *state.Session)) *mockAuthService_SignoutChat_Call {
 	_c.Run(run)
 	return _c
 }

+ 8 - 8
server/oscar/server.go

@@ -252,7 +252,6 @@ func (s oscarServer) connectToOSCARService(
 
 		fnCfg := func(sess *state.Session) {
 			sess.OnSessionClose(func() {
-				fmt.Println("CLOSING SESSION")
 				if !shuttingDown(ctx) {
 					instances := sess.Instances()
 					if len(instances) > 0 {
@@ -342,7 +341,14 @@ func (s oscarServer) connectToOSCARService(
 
 		go s.receiveSessMessages(ctx, instance, flapc)
 	case wire.Chat:
-		instance, err = s.AuthService.RegisterChatSession(ctx, cookie)
+		fnCfg := func(sess *state.Session) {
+			sess.OnSessionClose(func() {
+				ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
+				defer cancel()
+				s.SignoutChat(ctx, sess)
+			})
+		}
+		instance, err = s.AuthService.RegisterChatSession(ctx, cookie, fnCfg)
 		if err != nil {
 			return err
 		}
@@ -353,12 +359,6 @@ func (s oscarServer) connectToOSCARService(
 			instance.CloseInstance()
 		}()
 
-		instance.Session().OnSessionClose(func() {
-			ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
-			defer cancel()
-			s.SignoutChat(ctx, instance)
-		})
-
 		go s.receiveSessMessages(ctx, instance, flapc)
 	default:
 		instance, err = s.AuthService.RetrieveBOSSession(ctx, cookie)

+ 16 - 6
server/oscar/server_test.go

@@ -672,12 +672,17 @@ func TestOscarServer_RouteConnection_Chat(t *testing.T) {
 
 	authService := newMockAuthService(t)
 	authService.EXPECT().
-		RegisterChatSession(mock.Anything, state.ServerCookie{Service: wire.Chat}).
+		RegisterChatSession(mock.Anything, state.ServerCookie{Service: wire.Chat}, mock.Anything).
+		Run(func(_ context.Context, _ state.ServerCookie, cfg func(*state.Session)) {
+			if cfg != nil {
+				cfg(instance.Session())
+			}
+		}).
 		Return(instance, nil)
 	wg.Add(1)
 	authService.EXPECT().
-		SignoutChat(mock.Anything, instance).
-		Run(func(ctx context.Context, s *state.SessionInstance) {
+		SignoutChat(mock.Anything, instance.Session()).
+		Run(func(ctx context.Context, s *state.Session) {
 			defer wg.Done()
 		})
 
@@ -1069,14 +1074,19 @@ func Test_oscarServer_receiveSessMessages_Chat_integration(t *testing.T) {
 		CrackCookie(mock.Anything).
 		Return(state.ServerCookie{Service: wire.Chat}, nil)
 	authService.EXPECT().
-		RegisterChatSession(mock.Anything, state.ServerCookie{Service: wire.Chat}).
+		RegisterChatSession(mock.Anything, state.ServerCookie{Service: wire.Chat}, mock.Anything).
+		Run(func(_ context.Context, _ state.ServerCookie, cfg func(*state.Session)) {
+			if cfg != nil {
+				cfg(instance.Session())
+			}
+		}).
 		Return(instance, nil)
 
 	var signoutWG sync.WaitGroup
 	signoutWG.Add(1)
 	authService.EXPECT().
-		SignoutChat(mock.Anything, instance).
-		Run(func(ctx context.Context, s *state.SessionInstance) { signoutWG.Done() })
+		SignoutChat(mock.Anything, instance.Session()).
+		Run(func(ctx context.Context, s *state.Session) { signoutWG.Done() })
 
 	onlineNotifier := newMockOnlineNotifier(t)
 	onlineNotifier.EXPECT().

+ 2 - 2
server/oscar/types.go

@@ -52,10 +52,10 @@ type AuthService interface {
 	FLAPLogin(ctx context.Context, inFrame wire.FLAPSignonFrame, advertisedHost string) (wire.TLVRestBlock, error)
 	KerberosLogin(ctx context.Context, inBody wire.SNAC_0x050C_0x0002_KerberosLoginRequest, advertisedHost string) (wire.SNACMessage, error)
 	RegisterBOSSession(ctx context.Context, authCookie state.ServerCookie, conf func(sess *state.Session)) (*state.SessionInstance, error)
-	RegisterChatSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)
+	RegisterChatSession(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error)
 	RetrieveBOSSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)
 	Signout(ctx context.Context, session *state.Session)
-	SignoutChat(ctx context.Context, instance *state.SessionInstance)
+	SignoutChat(ctx context.Context, instance *state.Session)
 }
 
 type AdminService interface {

+ 16 - 15
server/toc/cmd_client.go

@@ -527,17 +527,18 @@ func (s OSCARProxy) ChatAccept(
 		return 0, s.runtimeErr(ctx, fmt.Errorf("AuthService.CrackCookie: %w", err))
 	}
 
-	chatSess, err := s.AuthService.RegisterChatSession(ctx, serverCookie)
+	fnCfg := func(sess *state.Session) {
+		sess.OnSessionClose(func() {
+			ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
+			defer cancel()
+			s.AuthService.SignoutChat(ctx, sess)
+		})
+	}
+	chatSess, err := s.AuthService.RegisterChatSession(ctx, serverCookie, fnCfg)
 	if err != nil {
 		return 0, s.runtimeErr(ctx, fmt.Errorf("AuthService.RegisterChatSession: %w", err))
 	}
 
-	chatSess.Session().OnSessionClose(func() {
-		ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
-		defer cancel()
-		s.AuthService.SignoutChat(ctx, chatSess)
-	})
-
 	if msg, isLimited := s.checkRateLimit(ctx, me, wire.OService, wire.OServiceClientOnline); isLimited {
 		return 0, msg
 	}
@@ -722,17 +723,18 @@ func (s OSCARProxy) ChatJoin(
 		return 0, s.runtimeErr(ctx, fmt.Errorf("AuthService.CrackCookie: %w", err))
 	}
 
-	chatSess, err := s.AuthService.RegisterChatSession(ctx, serverCookie)
+	fnCfg := func(sess *state.Session) {
+		sess.OnSessionClose(func() {
+			ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
+			defer cancel()
+			s.AuthService.SignoutChat(ctx, sess)
+		})
+	}
+	chatSess, err := s.AuthService.RegisterChatSession(ctx, serverCookie, fnCfg)
 	if err != nil {
 		return 0, s.runtimeErr(ctx, fmt.Errorf("AuthService.RegisterChatSession: %w", err))
 	}
 
-	chatSess.Session().OnSessionClose(func() {
-		ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
-		defer cancel()
-		s.AuthService.SignoutChat(ctx, chatSess)
-	})
-
 	if msg, isLimited := s.checkRateLimit(ctx, me, wire.OService, wire.OServiceClientOnline); isLimited {
 		return 0, msg
 	}
@@ -2351,7 +2353,6 @@ func (s OSCARProxy) Signon(ctx context.Context, args []byte, recalcWarning func(
 
 	fnCfg := func(sess *state.Session) {
 		sess.OnSessionClose(func() {
-			fmt.Println("closing session!")
 			if !shuttingDown(ctx) {
 				instances := sess.Instances()
 				if len(instances) > 0 {

+ 2 - 2
server/toc/cmd_client_test.go

@@ -950,7 +950,7 @@ func TestOSCARProxy_RecvClientCmd_ChatAccept(t *testing.T) {
 			authSvc := newMockAuthService(t)
 			for _, params := range tc.mockParams.authParams.registerChatSessionParams {
 				authSvc.EXPECT().
-					RegisterChatSession(ctx, params.authCookie).
+					RegisterChatSession(ctx, params.authCookie, mock.Anything).
 					Return(params.instance, params.err)
 			}
 			for _, params := range tc.mockParams.authParams.crackCookieParams {
@@ -1482,7 +1482,7 @@ func TestOSCARProxy_RecvClientCmd_ChatJoin(t *testing.T) {
 			authSvc := newMockAuthService(t)
 			for _, params := range tc.mockParams.authParams.registerChatSessionParams {
 				authSvc.EXPECT().
-					RegisterChatSession(ctx, params.authCookie).
+					RegisterChatSession(ctx, params.authCookie, mock.Anything).
 					Return(params.instance, params.err)
 			}
 			for _, params := range tc.mockParams.authParams.crackCookieParams {

+ 27 - 21
server/toc/mock_auth_service_test.go

@@ -391,8 +391,8 @@ func (_c *mockAuthService_RegisterBOSSession_Call) RunAndReturn(run func(ctx con
 }
 
 // RegisterChatSession provides a mock function for the type mockAuthService
-func (_mock *mockAuthService) RegisterChatSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error) {
-	ret := _mock.Called(ctx, authCookie)
+func (_mock *mockAuthService) RegisterChatSession(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error) {
+	ret := _mock.Called(ctx, authCookie, cfg)
 
 	if len(ret) == 0 {
 		panic("no return value specified for RegisterChatSession")
@@ -400,18 +400,18 @@ func (_mock *mockAuthService) RegisterChatSession(ctx context.Context, authCooki
 
 	var r0 *state.SessionInstance
 	var r1 error
-	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie) (*state.SessionInstance, error)); ok {
-		return returnFunc(ctx, authCookie)
+	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie, func(sess *state.Session)) (*state.SessionInstance, error)); ok {
+		return returnFunc(ctx, authCookie, cfg)
 	}
-	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie) *state.SessionInstance); ok {
-		r0 = returnFunc(ctx, authCookie)
+	if returnFunc, ok := ret.Get(0).(func(context.Context, state.ServerCookie, func(sess *state.Session)) *state.SessionInstance); ok {
+		r0 = returnFunc(ctx, authCookie, cfg)
 	} else {
 		if ret.Get(0) != nil {
 			r0 = ret.Get(0).(*state.SessionInstance)
 		}
 	}
-	if returnFunc, ok := ret.Get(1).(func(context.Context, state.ServerCookie) error); ok {
-		r1 = returnFunc(ctx, authCookie)
+	if returnFunc, ok := ret.Get(1).(func(context.Context, state.ServerCookie, func(sess *state.Session)) error); ok {
+		r1 = returnFunc(ctx, authCookie, cfg)
 	} else {
 		r1 = ret.Error(1)
 	}
@@ -426,11 +426,12 @@ type mockAuthService_RegisterChatSession_Call struct {
 // RegisterChatSession is a helper method to define mock.On call
 //   - ctx context.Context
 //   - authCookie state.ServerCookie
-func (_e *mockAuthService_Expecter) RegisterChatSession(ctx interface{}, authCookie interface{}) *mockAuthService_RegisterChatSession_Call {
-	return &mockAuthService_RegisterChatSession_Call{Call: _e.mock.On("RegisterChatSession", ctx, authCookie)}
+//   - cfg func(sess *state.Session)
+func (_e *mockAuthService_Expecter) RegisterChatSession(ctx interface{}, authCookie interface{}, cfg interface{}) *mockAuthService_RegisterChatSession_Call {
+	return &mockAuthService_RegisterChatSession_Call{Call: _e.mock.On("RegisterChatSession", ctx, authCookie, cfg)}
 }
 
-func (_c *mockAuthService_RegisterChatSession_Call) Run(run func(ctx context.Context, authCookie state.ServerCookie)) *mockAuthService_RegisterChatSession_Call {
+func (_c *mockAuthService_RegisterChatSession_Call) Run(run func(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session))) *mockAuthService_RegisterChatSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
 		var arg0 context.Context
 		if args[0] != nil {
@@ -440,9 +441,14 @@ func (_c *mockAuthService_RegisterChatSession_Call) Run(run func(ctx context.Con
 		if args[1] != nil {
 			arg1 = args[1].(state.ServerCookie)
 		}
+		var arg2 func(sess *state.Session)
+		if args[2] != nil {
+			arg2 = args[2].(func(sess *state.Session))
+		}
 		run(
 			arg0,
 			arg1,
+			arg2,
 		)
 	})
 	return _c
@@ -453,7 +459,7 @@ func (_c *mockAuthService_RegisterChatSession_Call) Return(sessionInstance *stat
 	return _c
 }
 
-func (_c *mockAuthService_RegisterChatSession_Call) RunAndReturn(run func(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)) *mockAuthService_RegisterChatSession_Call {
+func (_c *mockAuthService_RegisterChatSession_Call) RunAndReturn(run func(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error)) *mockAuthService_RegisterChatSession_Call {
 	_c.Call.Return(run)
 	return _c
 }
@@ -573,8 +579,8 @@ func (_c *mockAuthService_Signout_Call) RunAndReturn(run func(ctx context.Contex
 }
 
 // SignoutChat provides a mock function for the type mockAuthService
-func (_mock *mockAuthService) SignoutChat(ctx context.Context, instance *state.SessionInstance) {
-	_mock.Called(ctx, instance)
+func (_mock *mockAuthService) SignoutChat(ctx context.Context, sess *state.Session) {
+	_mock.Called(ctx, sess)
 	return
 }
 
@@ -585,20 +591,20 @@ type mockAuthService_SignoutChat_Call struct {
 
 // SignoutChat is a helper method to define mock.On call
 //   - ctx context.Context
-//   - instance *state.SessionInstance
-func (_e *mockAuthService_Expecter) SignoutChat(ctx interface{}, instance interface{}) *mockAuthService_SignoutChat_Call {
-	return &mockAuthService_SignoutChat_Call{Call: _e.mock.On("SignoutChat", ctx, instance)}
+//   - sess *state.Session
+func (_e *mockAuthService_Expecter) SignoutChat(ctx interface{}, sess interface{}) *mockAuthService_SignoutChat_Call {
+	return &mockAuthService_SignoutChat_Call{Call: _e.mock.On("SignoutChat", ctx, sess)}
 }
 
-func (_c *mockAuthService_SignoutChat_Call) Run(run func(ctx context.Context, instance *state.SessionInstance)) *mockAuthService_SignoutChat_Call {
+func (_c *mockAuthService_SignoutChat_Call) Run(run func(ctx context.Context, sess *state.Session)) *mockAuthService_SignoutChat_Call {
 	_c.Call.Run(func(args mock.Arguments) {
 		var arg0 context.Context
 		if args[0] != nil {
 			arg0 = args[0].(context.Context)
 		}
-		var arg1 *state.SessionInstance
+		var arg1 *state.Session
 		if args[1] != nil {
-			arg1 = args[1].(*state.SessionInstance)
+			arg1 = args[1].(*state.Session)
 		}
 		run(
 			arg0,
@@ -613,7 +619,7 @@ func (_c *mockAuthService_SignoutChat_Call) Return() *mockAuthService_SignoutCha
 	return _c
 }
 
-func (_c *mockAuthService_SignoutChat_Call) RunAndReturn(run func(ctx context.Context, instance *state.SessionInstance)) *mockAuthService_SignoutChat_Call {
+func (_c *mockAuthService_SignoutChat_Call) RunAndReturn(run func(ctx context.Context, sess *state.Session)) *mockAuthService_SignoutChat_Call {
 	_c.Run(run)
 	return _c
 }

+ 2 - 2
server/toc/types.go

@@ -50,10 +50,10 @@ type AuthService interface {
 	CrackCookie(authCookie []byte) (state.ServerCookie, error)
 	FLAPLogin(ctx context.Context, inFrame wire.FLAPSignonFrame, advertisedHost string) (wire.TLVRestBlock, error)
 	RegisterBOSSession(ctx context.Context, authCookie state.ServerCookie, cfg func(*state.Session)) (*state.SessionInstance, error)
-	RegisterChatSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)
+	RegisterChatSession(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error)
 	RetrieveBOSSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)
 	Signout(ctx context.Context, session *state.Session)
-	SignoutChat(ctx context.Context, instance *state.SessionInstance)
+	SignoutChat(ctx context.Context, sess *state.Session)
 }
 
 type LocateService interface {

+ 2 - 2
server/webapi/types.go

@@ -49,10 +49,10 @@ type AuthService interface {
 	CrackCookie(authCookie []byte) (state.ServerCookie, error)
 	FLAPLogin(ctx context.Context, inFrame wire.FLAPSignonFrame, advertisedHost string) (wire.TLVRestBlock, error)
 	RegisterBOSSession(ctx context.Context, authCookie state.ServerCookie, conf func(sess *state.Session)) (*state.SessionInstance, error)
-	RegisterChatSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)
+	RegisterChatSession(ctx context.Context, authCookie state.ServerCookie, cfg func(sess *state.Session)) (*state.SessionInstance, error)
 	RetrieveBOSSession(ctx context.Context, authCookie state.ServerCookie) (*state.SessionInstance, error)
 	Signout(ctx context.Context, session *state.Session)
-	SignoutChat(ctx context.Context, instance *state.SessionInstance)
+	SignoutChat(ctx context.Context, sess *state.Session)
 }
 
 type LocateService interface {

+ 6 - 6
state/session_manager.go

@@ -326,7 +326,7 @@ type InMemoryChatSessionManager struct {
 
 // AddSession adds a user to a chat room. If screenName already exists, the old
 // session is replaced by a new one.
-func (s *InMemoryChatSessionManager) AddSession(ctx context.Context, chatCookie string, screenName DisplayScreenName) (*SessionInstance, error) {
+func (s *InMemoryChatSessionManager) AddSession(ctx context.Context, chatCookie string, screenName DisplayScreenName, cfg ...func(sess *Session)) (*SessionInstance, error) {
 	s.mapMutex.Lock()
 	if _, ok := s.store[chatCookie]; !ok {
 		s.store[chatCookie] = NewInMemorySessionManager(s.logger)
@@ -337,7 +337,7 @@ func (s *InMemoryChatSessionManager) AddSession(ctx context.Context, chatCookie
 	ctx, cancel := context.WithTimeout(ctx, time.Second*5)
 	defer cancel()
 
-	sess, err := sessionManager.AddSession(ctx, screenName, false)
+	sess, err := sessionManager.AddSession(ctx, screenName, false, cfg...)
 	if err != nil {
 		return nil, fmt.Errorf("AddSession: %w", err)
 	}
@@ -366,18 +366,18 @@ func (s *InMemoryChatSessionManager) AddSession(ctx context.Context, chatCookie
 
 // RemoveSession removes a user session from a chat room. It panics if you
 // attempt to remove the session twice.
-func (s *InMemoryChatSessionManager) RemoveSession(instance *SessionInstance) {
+func (s *InMemoryChatSessionManager) RemoveSession(sess *Session) {
 	s.mapMutex.Lock()
 	defer s.mapMutex.Unlock()
 
-	sessionManager, ok := s.store[instance.ChatRoomCookie()]
+	sessionManager, ok := s.store[sess.ChatRoomCookie()]
 	if !ok {
 		panic("attempting to remove a session after its room has been deleted")
 	}
-	sessionManager.RemoveSession(instance.Session())
+	sessionManager.RemoveSession(sess)
 
 	if sessionManager.Empty() {
-		delete(s.store, instance.ChatRoomCookie())
+		delete(s.store, sess.ChatRoomCookie())
 	}
 }
 

+ 3 - 3
state/session_manager_test.go

@@ -635,8 +635,8 @@ func TestInMemoryChatSessionManager_RemoveSession(t *testing.T) {
 
 	assert.Len(t, sm.AllSessions("chat-room-1"), 2)
 
-	sm.RemoveSession(user1)
-	sm.RemoveSession(user2)
+	sm.RemoveSession(user1.Session())
+	sm.RemoveSession(user2.Session())
 
 	assert.Empty(t, sm.AllSessions("chat-room-1"))
 }
@@ -669,7 +669,7 @@ func TestInMemoryChatSessionManager_RemoveSession_DoubleLogin(t *testing.T) {
 			synctest.Wait()
 
 			// AddSession() is blocked waiting for the lock, now unblock it
-			sm.RemoveSession(chatSess1)
+			sm.RemoveSession(chatSess1.Session())
 
 			wg.Wait()
 		})