فهرست منبع

fix: empty buddy list when AIM<=3.5 boots >3.5 client

This commit makes sure that signon/signoff does not overlap for the same
user. For example, when a user signs on to client A, then attempts to
sign on to client B with the same screen name, client B's session can
not start until client A's session has completely shut down.

This fixes a race condition that occurred when a client signed into
Windows AIM > 3.5, then signed into Windows AIM <= 3.5. During login,
the buddy list status would get erased due to a race condition caused
by the first session terminating after the second session started.
Mike 1 سال پیش
والد
کامیت
e984f5d063

+ 133 - 16
cmd/server/factory.go

@@ -56,8 +56,22 @@ func MakeCommonDeps() (Container, error) {
 func Admin(deps Container) oscar.AdminServer {
 	logger := deps.logger.With("svc", "ADMIN")
 
-	adminService := foodgroup.NewAdminService(deps.sqLiteUserStore, deps.sqLiteUserStore, deps.inMemorySessionManager, deps.inMemorySessionManager)
-	authService := foodgroup.NewAuthService(deps.cfg, deps.inMemorySessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, deps.inMemorySessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, deps.inMemorySessionManager)
+	adminService := foodgroup.NewAdminService(
+		deps.sqLiteUserStore,
+		deps.sqLiteUserStore,
+		deps.inMemorySessionManager,
+		deps.inMemorySessionManager,
+	)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		deps.inMemorySessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.inMemorySessionManager,
+	)
 	oServiceService := foodgroup.NewOServiceServiceForAdmin(
 		deps.cfg,
 		logger,
@@ -84,8 +98,23 @@ func Alert(deps Container) oscar.BOSServer {
 	logger := deps.logger.With("svc", "ALERT")
 
 	sessionManager := state.NewInMemorySessionManager(logger)
-	authService := foodgroup.NewAuthService(deps.cfg, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
-	oServiceService := foodgroup.NewOServiceServiceForAlert(deps.cfg, logger, sessionManager, deps.sqLiteUserStore, sessionManager)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		sessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
+	oServiceService := foodgroup.NewOServiceServiceForAlert(
+		deps.cfg,
+		logger,
+		sessionManager,
+		deps.sqLiteUserStore,
+		sessionManager,
+	)
 
 	return oscar.BOSServer{
 		AuthService: authService,
@@ -104,7 +133,16 @@ func Alert(deps Container) oscar.BOSServer {
 func Auth(deps Container) oscar.AuthServer {
 	logger := deps.logger.With("svc", "AUTH")
 
-	authHandler := foodgroup.NewAuthService(deps.cfg, deps.inMemorySessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, nil, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
+	authHandler := foodgroup.NewAuthService(
+		deps.cfg,
+		deps.inMemorySessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
 
 	return oscar.AuthServer{
 		AuthService: authHandler,
@@ -118,9 +156,30 @@ func BART(deps Container) oscar.BOSServer {
 	logger := deps.logger.With("svc", "BART")
 
 	sessionManager := state.NewInMemorySessionManager(logger)
-	bartService := foodgroup.NewBARTService(logger, deps.sqLiteUserStore, sessionManager, deps.sqLiteUserStore, sessionManager)
-	authService := foodgroup.NewAuthService(deps.cfg, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
-	oServiceService := foodgroup.NewOServiceServiceForBART(deps.cfg, logger, sessionManager, deps.sqLiteUserStore, sessionManager)
+	bartService := foodgroup.NewBARTService(
+		logger,
+		deps.sqLiteUserStore,
+		sessionManager,
+		deps.sqLiteUserStore,
+		sessionManager,
+	)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		sessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
+	oServiceService := foodgroup.NewOServiceServiceForBART(
+		deps.cfg,
+		logger,
+		sessionManager,
+		deps.sqLiteUserStore,
+		sessionManager,
+	)
 
 	return oscar.BOSServer{
 		AuthService: authService,
@@ -139,7 +198,16 @@ func BART(deps Container) oscar.BOSServer {
 func BOS(deps Container) oscar.BOSServer {
 	logger := deps.logger.With("svc", "BOS")
 
-	authService := foodgroup.NewAuthService(deps.cfg, deps.inMemorySessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, deps.inMemorySessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		deps.inMemorySessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
 	bartService := foodgroup.NewBARTService(
 		logger,
 		deps.sqLiteUserStore,
@@ -147,7 +215,12 @@ func BOS(deps Container) oscar.BOSServer {
 		deps.sqLiteUserStore,
 		deps.inMemorySessionManager,
 	)
-	buddyService := foodgroup.NewBuddyService(deps.inMemorySessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, deps.inMemorySessionManager)
+	buddyService := foodgroup.NewBuddyService(
+		deps.inMemorySessionManager,
+		deps.sqLiteUserStore,
+		deps.sqLiteUserStore,
+		deps.inMemorySessionManager,
+	)
 	chatNavService := foodgroup.NewChatNavService(logger, deps.sqLiteUserStore)
 	feedbagService := foodgroup.NewFeedbagService(
 		logger,
@@ -163,10 +236,20 @@ func BOS(deps Container) oscar.BOSServer {
 		deps.inMemorySessionManager,
 		deps.inMemorySessionManager,
 	)
-	icbmService := foodgroup.NewICBMService(deps.inMemorySessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, deps.inMemorySessionManager)
+	icbmService := foodgroup.NewICBMService(
+		deps.inMemorySessionManager,
+		deps.sqLiteUserStore,
+		deps.sqLiteUserStore,
+		deps.inMemorySessionManager,
+	)
 	icqService := foodgroup.NewICQService(deps.inMemorySessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore,
 		logger, deps.inMemorySessionManager, deps.sqLiteUserStore)
-	locateService := foodgroup.NewLocateService(deps.inMemorySessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, deps.inMemorySessionManager)
+	locateService := foodgroup.NewLocateService(
+		deps.inMemorySessionManager,
+		deps.sqLiteUserStore,
+		deps.sqLiteUserStore,
+		deps.inMemorySessionManager,
+	)
 	oServiceService := foodgroup.NewOServiceServiceForBOS(
 		deps.cfg,
 		deps.inMemorySessionManager,
@@ -182,6 +265,7 @@ func BOS(deps Container) oscar.BOSServer {
 		AuthService:       authService,
 		BuddyListRegistry: deps.sqLiteUserStore,
 		Config:            deps.cfg,
+		DepartureNotifier: buddyService,
 		Handler: handler.NewBOSRouter(handler.Handlers{
 			AlertHandler:      handler.NewAlertHandler(logger),
 			BARTHandler:       handler.NewBARTHandler(logger, bartService),
@@ -206,7 +290,16 @@ func Chat(deps Container) oscar.ChatServer {
 	logger := deps.logger.With("svc", "CHAT")
 
 	sessionManager := state.NewInMemorySessionManager(logger)
-	authService := foodgroup.NewAuthService(deps.cfg, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		sessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
 	chatService := foodgroup.NewChatService(deps.chatSessionManager)
 	oServiceService := foodgroup.NewOServiceServiceForChat(
 		deps.cfg,
@@ -235,9 +328,24 @@ func ChatNav(deps Container) oscar.BOSServer {
 	logger := deps.logger.With("svc", "CHAT_NAV")
 
 	sessionManager := state.NewInMemorySessionManager(logger)
-	authService := foodgroup.NewAuthService(deps.cfg, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		sessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
 	chatNavService := foodgroup.NewChatNavService(logger, deps.sqLiteUserStore)
-	oServiceService := foodgroup.NewOServiceServiceForChatNav(deps.cfg, logger, sessionManager, deps.sqLiteUserStore, sessionManager)
+	oServiceService := foodgroup.NewOServiceServiceForChatNav(
+		deps.cfg,
+		logger,
+		sessionManager,
+		deps.sqLiteUserStore,
+		sessionManager,
+	)
 
 	return oscar.BOSServer{
 		AuthService: authService,
@@ -269,7 +377,16 @@ func ODir(deps Container) oscar.BOSServer {
 	logger := deps.logger.With("svc", "ODIR")
 
 	sessionManager := state.NewInMemorySessionManager(logger)
-	authService := foodgroup.NewAuthService(deps.cfg, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, sessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.sqLiteUserStore, nil)
+	authService := foodgroup.NewAuthService(
+		deps.cfg,
+		sessionManager,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		deps.hmacCookieBaker,
+		deps.chatSessionManager,
+		deps.sqLiteUserStore,
+		nil,
+	)
 	oServiceService := foodgroup.NewOServiceServiceForODir(deps.cfg, logger)
 	oDirService := foodgroup.NewODirService(logger, deps.sqLiteUserStore)
 

+ 18 - 13
foodgroup/auth.go

@@ -7,6 +7,7 @@ import (
 	"fmt"
 	"net"
 	"strconv"
+	"time"
 
 	"github.com/mk6i/retro-aim-server/config"
 	"github.com/mk6i/retro-aim-server/state"
@@ -22,14 +23,11 @@ func NewAuthService(
 	chatSessionRegistry ChatSessionRegistry,
 	userManager UserManager,
 	cookieBaker CookieBaker,
-	messageRelayer MessageRelayer,
 	chatMessageRelayer ChatMessageRelayer,
 	accountManager AccountManager,
-	buddyListRetriever BuddyListRetriever,
 	adminServerSessionRetriever SessionRetriever,
 ) *AuthService {
 	return &AuthService{
-		buddyBroadcaster:    newBuddyNotifier(buddyListRetriever, messageRelayer, adminServerSessionRetriever),
 		chatSessionRegistry: chatSessionRegistry,
 		config:              cfg,
 		cookieBaker:         cookieBaker,
@@ -46,7 +44,6 @@ func NewAuthService(
 // supports both FLAP (AIM v1.0-v3.0) and BUCP (AIM v3.5-v5.9) authentication
 // modes.
 type AuthService struct {
-	buddyBroadcaster            buddyBroadcaster
 	chatMessageRelayer          ChatMessageRelayer
 	chatSessionRegistry         ChatSessionRegistry
 	config                      config.Config
@@ -73,7 +70,11 @@ func (s AuthService) RegisterChatSession(authCookie []byte) (*state.Session, err
 	if err := wire.UnmarshalBE(&c, bytes.NewBuffer(token)); err != nil {
 		return nil, err
 	}
-	return s.chatSessionRegistry.AddSession(c.ChatCookie, c.ScreenName), nil
+	sess, err := s.chatSessionRegistry.AddSession(nil, c.ChatCookie, c.ScreenName)
+	if err != nil {
+		return nil, fmt.Errorf("AddSession: %w", err)
+	}
+	return sess, err
 }
 
 // bosCookie represents a token containing client metadata passed to the BOS
@@ -84,7 +85,7 @@ type bosCookie struct {
 }
 
 // RegisterBOSSession adds a new session to the session registry.
-func (s AuthService) RegisterBOSSession(authCookie []byte) (*state.Session, error) {
+func (s AuthService) RegisterBOSSession(ctx context.Context, authCookie []byte) (*state.Session, error) {
 	buf, err := s.cookieBaker.Crack(authCookie)
 	if err != nil {
 		return nil, err
@@ -103,7 +104,14 @@ func (s AuthService) RegisterBOSSession(authCookie []byte) (*state.Session, erro
 		return nil, fmt.Errorf("user not found")
 	}
 
-	sess := s.sessionManager.AddSession(u.DisplayScreenName)
+	ctx, cancel := context.WithTimeout(ctx, time.Second*5)
+	defer cancel()
+
+	sess, err := s.sessionManager.AddSession(ctx, u.DisplayScreenName)
+	if err != nil {
+		return nil, fmt.Errorf("AddSession: %w", err)
+	}
+
 	// Set the unconfirmed user info flag if this account is unconfirmed
 	if confirmed, err := s.accountManager.ConfirmStatusByName(sess.IdentScreenName()); err != nil {
 		return nil, fmt.Errorf("error setting unconfirmed user flag: %w", err)
@@ -151,13 +159,10 @@ func (s AuthService) RetrieveBOSSession(authCookie []byte) (*state.Session, erro
 }
 
 // Signout removes this user's session and notifies users who have this user on
-// their buddy list about this user's departure.
-func (s AuthService) Signout(ctx context.Context, sess *state.Session) error {
-	if err := s.buddyBroadcaster.BroadcastBuddyDeparted(ctx, sess); err != nil {
-		return err
-	}
+// their buddy list about this user's departure. It's guaranteed that the
+// session is removed from the session pool.
+func (s AuthService) Signout(_ context.Context, sess *state.Session) {
 	s.sessionManager.RemoveSession(sess)
-	return nil
 }
 
 // SignoutChat removes user from chat room and notifies remaining participants

+ 13 - 20
foodgroup/auth_test.go

@@ -2,6 +2,7 @@ package foodgroup
 
 import (
 	"bytes"
+	"context"
 	"fmt"
 	"io"
 	"testing"
@@ -1003,8 +1004,8 @@ func TestAuthService_RegisterChatSession_HappyPath(t *testing.T) {
 	chatCookie := "the-chat-cookie"
 	chatSessionRegistry := newMockChatSessionRegistry(t)
 	chatSessionRegistry.EXPECT().
-		AddSession(chatCookie, sess.DisplayScreenName()).
-		Return(sess)
+		AddSession(mock.Anything, chatCookie, sess.DisplayScreenName()).
+		Return(sess, nil)
 
 	c := chatLoginCookie{
 		ChatCookie: chatCookie,
@@ -1019,7 +1020,7 @@ func TestAuthService_RegisterChatSession_HappyPath(t *testing.T) {
 		Crack(authCookie).
 		Return(chatCookieBuf.Bytes(), nil)
 
-	svc := NewAuthService(config.Config{}, nil, chatSessionRegistry, nil, cookieBaker, nil, nil, nil, nil, nil)
+	svc := NewAuthService(config.Config{}, nil, chatSessionRegistry, nil, cookieBaker, nil, nil, nil)
 
 	have, err := svc.RegisterChatSession(authCookie)
 	assert.NoError(t, err)
@@ -1153,8 +1154,8 @@ func TestAuthService_RegisterBOSSession(t *testing.T) {
 			sessionRegistry := newMockSessionRegistry(t)
 			for _, params := range tc.mockParams.addSessionParams {
 				sessionRegistry.EXPECT().
-					AddSession(params.screenName).
-					Return(params.result)
+					AddSession(mock.Anything, params.screenName).
+					Return(params.result, params.err)
 			}
 			cookieBaker := newMockCookieBaker(t)
 			for _, params := range tc.mockParams.cookieCrackParams {
@@ -1175,9 +1176,9 @@ func TestAuthService_RegisterBOSSession(t *testing.T) {
 					Return(params.confirmStatus, nil)
 			}
 
-			svc := NewAuthService(config.Config{}, sessionRegistry, nil, userManager, cookieBaker, nil, nil, accountManager, nil, nil)
+			svc := NewAuthService(config.Config{}, sessionRegistry, nil, userManager, cookieBaker, nil, accountManager, nil)
 
-			have, err := svc.RegisterBOSSession(tc.cookie)
+			have, err := svc.RegisterBOSSession(context.Background(), tc.cookie)
 			assert.NoError(t, err)
 
 			if tc.wantSess != nil {
@@ -1213,7 +1214,7 @@ func TestAuthService_RetrieveBOSSession_HappyPath(t *testing.T) {
 		User(sess.IdentScreenName()).
 		Return(&state.User{IdentScreenName: sess.IdentScreenName()}, nil)
 
-	svc := NewAuthService(config.Config{}, nil, nil, userManager, cookieBaker, nil, nil, nil, nil, sessionRetriever)
+	svc := NewAuthService(config.Config{}, nil, nil, userManager, cookieBaker, nil, nil, sessionRetriever)
 
 	have, err := svc.RetrieveBOSSession(authCookie)
 	assert.NoError(t, err)
@@ -1246,7 +1247,7 @@ func TestAuthService_RetrieveBOSSession_SessionNotFound(t *testing.T) {
 		User(sess.IdentScreenName()).
 		Return(&state.User{IdentScreenName: sess.IdentScreenName()}, nil)
 
-	svc := NewAuthService(config.Config{}, nil, nil, userManager, cookieBaker, nil, nil, nil, nil, sessionRetriever)
+	svc := NewAuthService(config.Config{}, nil, nil, userManager, cookieBaker, nil, nil, sessionRetriever)
 
 	have, err := svc.RetrieveBOSSession(authCookie)
 	assert.NoError(t, err)
@@ -1339,7 +1340,7 @@ func TestAuthService_SignoutChat(t *testing.T) {
 					RemoveSession(matchSession(params.screenName))
 			}
 
-			svc := NewAuthService(config.Config{}, nil, sessionManager, nil, nil, nil, chatMessageRelayer, nil, nil, nil)
+			svc := NewAuthService(config.Config{}, nil, sessionManager, nil, nil, chatMessageRelayer, nil, nil)
 			svc.SignoutChat(nil, tt.userSession)
 		})
 	}
@@ -1384,17 +1385,9 @@ func TestAuthService_Signout(t *testing.T) {
 			for _, params := range tt.mockParams.removeSessionParams {
 				sessionManager.EXPECT().RemoveSession(matchSession(params.screenName))
 			}
-			buddyUpdateBroadcaster := newMockbuddyBroadcaster(t)
-			for _, params := range tt.mockParams.broadcastBuddyDepartedParams {
-				buddyUpdateBroadcaster.EXPECT().
-					BroadcastBuddyDeparted(mock.Anything, matchSession(params.screenName)).
-					Return(params.err)
-			}
-			svc := NewAuthService(config.Config{}, sessionManager, nil, nil, nil, nil, nil, nil, nil, nil)
-			svc.buddyBroadcaster = buddyUpdateBroadcaster
+			svc := NewAuthService(config.Config{}, sessionManager, nil, nil, nil, nil, nil, nil)
 
-			err := svc.Signout(nil, tt.userSession)
-			assert.ErrorIs(t, err, tt.wantErr)
+			svc.Signout(nil, tt.userSession)
 		})
 	}
 }

+ 4 - 0
foodgroup/buddy.go

@@ -103,6 +103,10 @@ func (s BuddyService) DelBuddies(
 	return nil
 }
 
+func (s BuddyService) BroadcastBuddyDeparted(ctx context.Context, sess *state.Session) error {
+	return s.buddyBroadcaster.BroadcastBuddyDeparted(ctx, sess)
+}
+
 func newBuddyNotifier(
 	buddyListRetriever BuddyListRetriever,
 	messageRelayer MessageRelayer,

+ 26 - 13
foodgroup/mock_chat_session_registry_test.go

@@ -3,6 +3,8 @@
 package foodgroup
 
 import (
+	context "context"
+
 	state "github.com/mk6i/retro-aim-server/state"
 	mock "github.com/stretchr/testify/mock"
 )
@@ -20,24 +22,34 @@ func (_m *mockChatSessionRegistry) EXPECT() *mockChatSessionRegistry_Expecter {
 	return &mockChatSessionRegistry_Expecter{mock: &_m.Mock}
 }
 
-// AddSession provides a mock function with given fields: chatCookie, screenName
-func (_m *mockChatSessionRegistry) AddSession(chatCookie string, screenName state.DisplayScreenName) *state.Session {
-	ret := _m.Called(chatCookie, screenName)
+// AddSession provides a mock function with given fields: ctx, chatCookie, screenName
+func (_m *mockChatSessionRegistry) AddSession(ctx context.Context, chatCookie string, screenName state.DisplayScreenName) (*state.Session, error) {
+	ret := _m.Called(ctx, chatCookie, screenName)
 
 	if len(ret) == 0 {
 		panic("no return value specified for AddSession")
 	}
 
 	var r0 *state.Session
-	if rf, ok := ret.Get(0).(func(string, state.DisplayScreenName) *state.Session); ok {
-		r0 = rf(chatCookie, screenName)
+	var r1 error
+	if rf, ok := ret.Get(0).(func(context.Context, string, state.DisplayScreenName) (*state.Session, error)); ok {
+		return rf(ctx, chatCookie, screenName)
+	}
+	if rf, ok := ret.Get(0).(func(context.Context, string, state.DisplayScreenName) *state.Session); ok {
+		r0 = rf(ctx, chatCookie, screenName)
 	} else {
 		if ret.Get(0) != nil {
 			r0 = ret.Get(0).(*state.Session)
 		}
 	}
 
-	return r0
+	if rf, ok := ret.Get(1).(func(context.Context, string, state.DisplayScreenName) error); ok {
+		r1 = rf(ctx, chatCookie, screenName)
+	} else {
+		r1 = ret.Error(1)
+	}
+
+	return r0, r1
 }
 
 // mockChatSessionRegistry_AddSession_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddSession'
@@ -46,25 +58,26 @@ type mockChatSessionRegistry_AddSession_Call struct {
 }
 
 // AddSession is a helper method to define mock.On call
+//   - ctx context.Context
 //   - chatCookie string
 //   - screenName state.DisplayScreenName
-func (_e *mockChatSessionRegistry_Expecter) AddSession(chatCookie interface{}, screenName interface{}) *mockChatSessionRegistry_AddSession_Call {
-	return &mockChatSessionRegistry_AddSession_Call{Call: _e.mock.On("AddSession", chatCookie, screenName)}
+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)}
 }
 
-func (_c *mockChatSessionRegistry_AddSession_Call) Run(run func(chatCookie string, screenName state.DisplayScreenName)) *mockChatSessionRegistry_AddSession_Call {
+func (_c *mockChatSessionRegistry_AddSession_Call) Run(run func(ctx context.Context, chatCookie string, screenName state.DisplayScreenName)) *mockChatSessionRegistry_AddSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
-		run(args[0].(string), args[1].(state.DisplayScreenName))
+		run(args[0].(context.Context), args[1].(string), args[2].(state.DisplayScreenName))
 	})
 	return _c
 }
 
-func (_c *mockChatSessionRegistry_AddSession_Call) Return(_a0 *state.Session) *mockChatSessionRegistry_AddSession_Call {
-	_c.Call.Return(_a0)
+func (_c *mockChatSessionRegistry_AddSession_Call) Return(_a0 *state.Session, _a1 error) *mockChatSessionRegistry_AddSession_Call {
+	_c.Call.Return(_a0, _a1)
 	return _c
 }
 
-func (_c *mockChatSessionRegistry_AddSession_Call) RunAndReturn(run func(string, state.DisplayScreenName) *state.Session) *mockChatSessionRegistry_AddSession_Call {
+func (_c *mockChatSessionRegistry_AddSession_Call) RunAndReturn(run func(context.Context, string, state.DisplayScreenName) (*state.Session, error)) *mockChatSessionRegistry_AddSession_Call {
 	_c.Call.Return(run)
 	return _c
 }

+ 26 - 13
foodgroup/mock_session_registry_test.go

@@ -3,6 +3,8 @@
 package foodgroup
 
 import (
+	context "context"
+
 	state "github.com/mk6i/retro-aim-server/state"
 	mock "github.com/stretchr/testify/mock"
 )
@@ -20,24 +22,34 @@ func (_m *mockSessionRegistry) EXPECT() *mockSessionRegistry_Expecter {
 	return &mockSessionRegistry_Expecter{mock: &_m.Mock}
 }
 
-// AddSession provides a mock function with given fields: screenName
-func (_m *mockSessionRegistry) AddSession(screenName state.DisplayScreenName) *state.Session {
-	ret := _m.Called(screenName)
+// AddSession provides a mock function with given fields: ctx, screenName
+func (_m *mockSessionRegistry) AddSession(ctx context.Context, screenName state.DisplayScreenName) (*state.Session, error) {
+	ret := _m.Called(ctx, screenName)
 
 	if len(ret) == 0 {
 		panic("no return value specified for AddSession")
 	}
 
 	var r0 *state.Session
-	if rf, ok := ret.Get(0).(func(state.DisplayScreenName) *state.Session); ok {
-		r0 = rf(screenName)
+	var r1 error
+	if rf, ok := ret.Get(0).(func(context.Context, state.DisplayScreenName) (*state.Session, error)); ok {
+		return rf(ctx, screenName)
+	}
+	if rf, ok := ret.Get(0).(func(context.Context, state.DisplayScreenName) *state.Session); ok {
+		r0 = rf(ctx, screenName)
 	} else {
 		if ret.Get(0) != nil {
 			r0 = ret.Get(0).(*state.Session)
 		}
 	}
 
-	return r0
+	if rf, ok := ret.Get(1).(func(context.Context, state.DisplayScreenName) error); ok {
+		r1 = rf(ctx, screenName)
+	} else {
+		r1 = ret.Error(1)
+	}
+
+	return r0, r1
 }
 
 // mockSessionRegistry_AddSession_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddSession'
@@ -46,24 +58,25 @@ type mockSessionRegistry_AddSession_Call struct {
 }
 
 // AddSession is a helper method to define mock.On call
+//   - ctx context.Context
 //   - screenName state.DisplayScreenName
-func (_e *mockSessionRegistry_Expecter) AddSession(screenName interface{}) *mockSessionRegistry_AddSession_Call {
-	return &mockSessionRegistry_AddSession_Call{Call: _e.mock.On("AddSession", screenName)}
+func (_e *mockSessionRegistry_Expecter) AddSession(ctx interface{}, screenName interface{}) *mockSessionRegistry_AddSession_Call {
+	return &mockSessionRegistry_AddSession_Call{Call: _e.mock.On("AddSession", ctx, screenName)}
 }
 
-func (_c *mockSessionRegistry_AddSession_Call) Run(run func(screenName state.DisplayScreenName)) *mockSessionRegistry_AddSession_Call {
+func (_c *mockSessionRegistry_AddSession_Call) Run(run func(ctx context.Context, screenName state.DisplayScreenName)) *mockSessionRegistry_AddSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
-		run(args[0].(state.DisplayScreenName))
+		run(args[0].(context.Context), args[1].(state.DisplayScreenName))
 	})
 	return _c
 }
 
-func (_c *mockSessionRegistry_AddSession_Call) Return(_a0 *state.Session) *mockSessionRegistry_AddSession_Call {
-	_c.Call.Return(_a0)
+func (_c *mockSessionRegistry_AddSession_Call) Return(_a0 *state.Session, _a1 error) *mockSessionRegistry_AddSession_Call {
+	_c.Call.Return(_a0, _a1)
 	return _c
 }
 
-func (_c *mockSessionRegistry_AddSession_Call) RunAndReturn(run func(state.DisplayScreenName) *state.Session) *mockSessionRegistry_AddSession_Call {
+func (_c *mockSessionRegistry_AddSession_Call) RunAndReturn(run func(context.Context, state.DisplayScreenName) (*state.Session, error)) *mockSessionRegistry_AddSession_Call {
 	_c.Call.Return(run)
 	return _c
 }

+ 1 - 0
foodgroup/test_helpers.go

@@ -276,6 +276,7 @@ type sessionRegistryParams struct {
 type addSessionParams []struct {
 	screenName state.DisplayScreenName
 	result     *state.Session
+	err        error
 }
 
 // removeSessionParams is the list of parameters passed at the mock

+ 2 - 2
foodgroup/types.go

@@ -112,7 +112,7 @@ 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(chatCookie string, screenName state.DisplayScreenName) *state.Session
+	AddSession(ctx context.Context, chatCookie string, screenName state.DisplayScreenName) (*state.Session, error)
 
 	// RemoveSession removes a session from the chat session manager.
 	RemoveSession(sess *state.Session)
@@ -190,7 +190,7 @@ type ProfileManager interface {
 }
 
 type SessionRegistry interface {
-	AddSession(screenName state.DisplayScreenName) *state.Session
+	AddSession(ctx context.Context, screenName state.DisplayScreenName) (*state.Session, error)
 	RemoveSession(sess *state.Session)
 }
 

+ 2 - 2
server/oscar/auth.go

@@ -20,10 +20,10 @@ type AuthService interface {
 	BUCPChallenge(bodyIn wire.SNAC_0x17_0x06_BUCPChallengeRequest, newUUID func() uuid.UUID) (wire.SNACMessage, error)
 	BUCPLogin(bodyIn wire.SNAC_0x17_0x02_BUCPLoginRequest, newUserFn func(screenName state.DisplayScreenName) (state.User, error)) (wire.SNACMessage, error)
 	FLAPLogin(frame wire.FLAPSignonFrame, newUserFn func(screenName state.DisplayScreenName) (state.User, error)) (wire.TLVRestBlock, error)
-	RegisterBOSSession(authCookie []byte) (*state.Session, error)
+	RegisterBOSSession(ctx context.Context, authCookie []byte) (*state.Session, error)
 	RetrieveBOSSession(authCookie []byte) (*state.Session, error)
 	RegisterChatSession(authCookie []byte) (*state.Session, error)
-	Signout(ctx context.Context, sess *state.Session) error
+	Signout(ctx context.Context, sess *state.Session)
 	SignoutChat(ctx context.Context, sess *state.Session)
 }
 

+ 16 - 3
server/oscar/bos.go

@@ -31,11 +31,18 @@ type BuddyListRegistry interface {
 	UnregisterBuddyList(user state.IdentScreenName) error
 }
 
+// DepartureNotifier is the interface for sending buddy departure notifications
+// when a client disconnects.
+type DepartureNotifier interface {
+	BroadcastBuddyDeparted(ctx context.Context, sess *state.Session) error
+}
+
 // BOSServer provides client connection lifecycle management for the BOS
 // service.
 type BOSServer struct {
 	AuthService
 	BuddyListRegistry
+	DepartureNotifier
 	Handler
 	ListenAddr string
 	Logger     *slog.Logger
@@ -132,7 +139,7 @@ func (rt BOSServer) handleNewConnection(ctx context.Context, rwc io.ReadWriteClo
 		return errors.New("unable to get session id from payload")
 	}
 
-	sess, err := rt.RegisterBOSSession(authCookie)
+	sess, err := rt.RegisterBOSSession(ctx, authCookie)
 	if err != nil {
 		return err
 	}
@@ -149,14 +156,20 @@ func (rt BOSServer) handleNewConnection(ctx context.Context, rwc io.ReadWriteClo
 	defer func() {
 		sess.Close()
 		rwc.Close()
-		if err := rt.Signout(ctx, sess); err != nil {
-			rt.Logger.ErrorContext(ctx, "error notifying departure", "err", err.Error())
+		if rt.DepartureNotifier != nil {
+			if err := rt.DepartureNotifier.BroadcastBuddyDeparted(ctx, sess); err != nil {
+				rt.Logger.ErrorContext(ctx, "error sending buddy departure notifications", "err", err.Error())
+			}
 		}
 		if rt.BuddyListRegistry != nil { // nil check is a hack until server refactor
+			// buddy list must be cleared before session is closed, otherwise
+			// there will be a race condition that could cause the buddy list
+			// be prematurely deleted.
 			if err := rt.BuddyListRegistry.UnregisterBuddyList(sess.IdentScreenName()); err != nil {
 				rt.Logger.ErrorContext(ctx, "error removing buddy list entry", "err", err.Error())
 			}
 		}
+		rt.Signout(ctx, sess)
 	}()
 
 	ctx = context.WithValue(ctx, "screenName", sess.IdentScreenName())

+ 2 - 3
server/oscar/bos_test.go

@@ -73,11 +73,10 @@ func TestBOSService_handleNewConnection(t *testing.T) {
 
 	authService := newMockAuthService(t)
 	authService.EXPECT().
-		RegisterBOSSession([]byte("the-cookie")).
+		RegisterBOSSession(mock.Anything, []byte("the-cookie")).
 		Return(sess, nil)
 	authService.EXPECT().
-		Signout(mock.Anything, sess).
-		Return(nil)
+		Signout(mock.Anything, sess)
 
 	onlineNotifier := newMockOnlineNotifier(t)
 	onlineNotifier.EXPECT().

+ 2 - 2
server/oscar/connection_test.go

@@ -18,7 +18,7 @@ import (
 func TestHandleChatConnection_MessageRelay(t *testing.T) {
 	sessionManager := state.NewInMemorySessionManager(slog.Default())
 	// add a user to session that will receive relayed messages
-	sess := sessionManager.AddSession("bob")
+	sess, _ := sessionManager.AddSession(nil, "bob")
 
 	// start the server connection handler in the background
 	serverReader, _ := io.Pipe()
@@ -90,7 +90,7 @@ func TestHandleChatConnection_MessageRelay(t *testing.T) {
 func TestHandleChatConnection_ClientRequest(t *testing.T) {
 	sessionManager := state.NewInMemorySessionManager(slog.Default())
 	// add session so that the function can terminate upon closure
-	sess := sessionManager.AddSession("bob")
+	sess, _ := sessionManager.AddSession(nil, "bob")
 
 	inboundMsgs := []wire.SNACMessage{
 		{

+ 20 - 32
server/oscar/mock_auth_test.go

@@ -197,9 +197,9 @@ func (_c *mockAuthService_FLAPLogin_Call) RunAndReturn(run func(wire.FLAPSignonF
 	return _c
 }
 
-// RegisterBOSSession provides a mock function with given fields: authCookie
-func (_m *mockAuthService) RegisterBOSSession(authCookie []byte) (*state.Session, error) {
-	ret := _m.Called(authCookie)
+// RegisterBOSSession provides a mock function with given fields: ctx, authCookie
+func (_m *mockAuthService) RegisterBOSSession(ctx context.Context, authCookie []byte) (*state.Session, error) {
+	ret := _m.Called(ctx, authCookie)
 
 	if len(ret) == 0 {
 		panic("no return value specified for RegisterBOSSession")
@@ -207,19 +207,19 @@ func (_m *mockAuthService) RegisterBOSSession(authCookie []byte) (*state.Session
 
 	var r0 *state.Session
 	var r1 error
-	if rf, ok := ret.Get(0).(func([]byte) (*state.Session, error)); ok {
-		return rf(authCookie)
+	if rf, ok := ret.Get(0).(func(context.Context, []byte) (*state.Session, error)); ok {
+		return rf(ctx, authCookie)
 	}
-	if rf, ok := ret.Get(0).(func([]byte) *state.Session); ok {
-		r0 = rf(authCookie)
+	if rf, ok := ret.Get(0).(func(context.Context, []byte) *state.Session); ok {
+		r0 = rf(ctx, authCookie)
 	} else {
 		if ret.Get(0) != nil {
 			r0 = ret.Get(0).(*state.Session)
 		}
 	}
 
-	if rf, ok := ret.Get(1).(func([]byte) error); ok {
-		r1 = rf(authCookie)
+	if rf, ok := ret.Get(1).(func(context.Context, []byte) error); ok {
+		r1 = rf(ctx, authCookie)
 	} else {
 		r1 = ret.Error(1)
 	}
@@ -233,14 +233,15 @@ type mockAuthService_RegisterBOSSession_Call struct {
 }
 
 // RegisterBOSSession is a helper method to define mock.On call
+//   - ctx context.Context
 //   - authCookie []byte
-func (_e *mockAuthService_Expecter) RegisterBOSSession(authCookie interface{}) *mockAuthService_RegisterBOSSession_Call {
-	return &mockAuthService_RegisterBOSSession_Call{Call: _e.mock.On("RegisterBOSSession", authCookie)}
+func (_e *mockAuthService_Expecter) RegisterBOSSession(ctx interface{}, authCookie interface{}) *mockAuthService_RegisterBOSSession_Call {
+	return &mockAuthService_RegisterBOSSession_Call{Call: _e.mock.On("RegisterBOSSession", ctx, authCookie)}
 }
 
-func (_c *mockAuthService_RegisterBOSSession_Call) Run(run func(authCookie []byte)) *mockAuthService_RegisterBOSSession_Call {
+func (_c *mockAuthService_RegisterBOSSession_Call) Run(run func(ctx context.Context, authCookie []byte)) *mockAuthService_RegisterBOSSession_Call {
 	_c.Call.Run(func(args mock.Arguments) {
-		run(args[0].([]byte))
+		run(args[0].(context.Context), args[1].([]byte))
 	})
 	return _c
 }
@@ -250,7 +251,7 @@ func (_c *mockAuthService_RegisterBOSSession_Call) Return(_a0 *state.Session, _a
 	return _c
 }
 
-func (_c *mockAuthService_RegisterBOSSession_Call) RunAndReturn(run func([]byte) (*state.Session, error)) *mockAuthService_RegisterBOSSession_Call {
+func (_c *mockAuthService_RegisterBOSSession_Call) RunAndReturn(run func(context.Context, []byte) (*state.Session, error)) *mockAuthService_RegisterBOSSession_Call {
 	_c.Call.Return(run)
 	return _c
 }
@@ -372,21 +373,8 @@ func (_c *mockAuthService_RetrieveBOSSession_Call) RunAndReturn(run func([]byte)
 }
 
 // Signout provides a mock function with given fields: ctx, sess
-func (_m *mockAuthService) Signout(ctx context.Context, sess *state.Session) error {
-	ret := _m.Called(ctx, sess)
-
-	if len(ret) == 0 {
-		panic("no return value specified for Signout")
-	}
-
-	var r0 error
-	if rf, ok := ret.Get(0).(func(context.Context, *state.Session) error); ok {
-		r0 = rf(ctx, sess)
-	} else {
-		r0 = ret.Error(0)
-	}
-
-	return r0
+func (_m *mockAuthService) Signout(ctx context.Context, sess *state.Session) {
+	_m.Called(ctx, sess)
 }
 
 // mockAuthService_Signout_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Signout'
@@ -408,12 +396,12 @@ func (_c *mockAuthService_Signout_Call) Run(run func(ctx context.Context, sess *
 	return _c
 }
 
-func (_c *mockAuthService_Signout_Call) Return(_a0 error) *mockAuthService_Signout_Call {
-	_c.Call.Return(_a0)
+func (_c *mockAuthService_Signout_Call) Return() *mockAuthService_Signout_Call {
+	_c.Call.Return()
 	return _c
 }
 
-func (_c *mockAuthService_Signout_Call) RunAndReturn(run func(context.Context, *state.Session) error) *mockAuthService_Signout_Call {
+func (_c *mockAuthService_Signout_Call) RunAndReturn(run func(context.Context, *state.Session)) *mockAuthService_Signout_Call {
 	_c.Call.Return(run)
 	return _c
 }

+ 80 - 37
state/session_manager.go

@@ -2,17 +2,26 @@ package state
 
 import (
 	"context"
+	"errors"
+	"fmt"
 	"log/slog"
 	"sync"
 
 	"github.com/mk6i/retro-aim-server/wire"
 )
 
+type sessionSlot struct {
+	sess    *Session
+	removed chan bool
+}
+
+var errSessConflict = errors.New("session conflict: another session was created concurrently for this user")
+
 // InMemorySessionManager handles the lifecycle of a user session and provides
 // synchronized message relay between sessions in the session pool. An
 // InMemorySessionManager is safe for concurrent use by multiple goroutines.
 type InMemorySessionManager struct {
-	store    map[IdentScreenName]*Session
+	store    map[IdentScreenName]*sessionSlot
 	mapMutex sync.RWMutex
 	logger   *slog.Logger
 }
@@ -21,7 +30,7 @@ type InMemorySessionManager struct {
 func NewInMemorySessionManager(logger *slog.Logger) *InMemorySessionManager {
 	return &InMemorySessionManager{
 		logger: logger,
-		store:  make(map[IdentScreenName]*Session),
+		store:  make(map[IdentScreenName]*sessionSlot),
 	}
 }
 
@@ -29,8 +38,8 @@ func NewInMemorySessionManager(logger *slog.Logger) *InMemorySessionManager {
 func (s *InMemorySessionManager) RelayToAll(ctx context.Context, msg wire.SNACMessage) {
 	s.mapMutex.RLock()
 	defer s.mapMutex.RUnlock()
-	for _, sess := range s.store {
-		s.maybeRelayMessage(ctx, msg, sess)
+	for _, rec := range s.store {
+		s.maybeRelayMessage(ctx, msg, rec.sess)
 	}
 }
 
@@ -61,42 +70,69 @@ func (s *InMemorySessionManager) maybeRelayMessage(ctx context.Context, msg wire
 	}
 }
 
-// AddSession adds a new session to the pool. It replaces an existing session
-// with a matching screen name, ensuring that each screen name is unique in the
-// pool. This method does not return a nil session value. It's possible to add
-// a non-existent user to the session. Callers should ensure that the account
-// represented by displayScreenName is valid.
-func (s *InMemorySessionManager) AddSession(displayScreenName DisplayScreenName) *Session {
+// AddSession adds a new session to the pool, ensuring only one session exists
+// for a given screen name. If a session with the same screen name is already
+// active, the call blocks until the active session is terminated by
+// [InMemorySessionManager.RemoveSession] or the context is canceled. When
+// concurrent calls are made for the same screen name, only one call succeeds
+// and the others return an error.
+func (s *InMemorySessionManager) AddSession(ctx context.Context, screenName DisplayScreenName) (*Session, error) {
 	s.mapMutex.Lock()
-	defer s.mapMutex.Unlock()
 
-	identScreenName := NewIdentScreenName(string(displayScreenName))
-	// Only allow one session at a time per screen name. A session may already
-	// exist because:
-	// 1) the user is signing on using an already logged-on screen name.
-	// 2) the session might be orphaned due to an undetected client
-	// disconnection.
-	for _, sess := range s.store {
-		if identScreenName == sess.IdentScreenName() {
-			sess.Close()
-			delete(s.store, identScreenName)
-			break
+	active := s.findRec(screenName.IdentScreenName())
+	if active != nil {
+		// there's an active session that needs to be removed. don't hold the
+		// lock while we wait.
+		s.mapMutex.Unlock()
+
+		// signal to callers that this session has to go
+		active.sess.Close()
+
+		select {
+		case <-active.removed: // wait for RemoveSession to be called
+		case <-ctx.Done():
+			return nil, fmt.Errorf("waiting for previous session to terminate: %w", ctx.Err())
 		}
+
+		// the session has been removed, let's try to replace it
+		s.mapMutex.Lock()
+	}
+
+	defer s.mapMutex.Unlock()
+
+	// make sure a concurrent call didn't already add a session
+	if active != nil && s.findRec(screenName.IdentScreenName()) != nil {
+		return nil, errSessConflict
 	}
 
 	sess := NewSession()
-	sess.SetIdentScreenName(identScreenName)
-	sess.SetDisplayScreenName(displayScreenName)
-	s.store[identScreenName] = sess
-	return sess
+	sess.SetIdentScreenName(screenName.IdentScreenName())
+	sess.SetDisplayScreenName(screenName)
+
+	s.store[sess.IdentScreenName()] = &sessionSlot{
+		sess:    sess,
+		removed: make(chan bool),
+	}
+
+	return sess, nil
+}
+
+func (s *InMemorySessionManager) findRec(identScreenName IdentScreenName) *sessionSlot {
+	for _, rec := range s.store {
+		if identScreenName == rec.sess.IdentScreenName() {
+			return rec
+		}
+	}
+	return nil
 }
 
 // RemoveSession takes a session out of the session pool.
 func (s *InMemorySessionManager) RemoveSession(sess *Session) {
 	s.mapMutex.Lock()
 	defer s.mapMutex.Unlock()
-	if sess == s.store[sess.IdentScreenName()] {
+	if rec, ok := s.store[sess.IdentScreenName()]; ok && rec.sess == sess {
 		delete(s.store, sess.IdentScreenName())
+		close(rec.removed)
 	}
 }
 
@@ -105,7 +141,10 @@ func (s *InMemorySessionManager) RemoveSession(sess *Session) {
 func (s *InMemorySessionManager) RetrieveSession(screenName IdentScreenName) *Session {
 	s.mapMutex.RLock()
 	defer s.mapMutex.RUnlock()
-	return s.store[screenName]
+	if rec, ok := s.store[screenName]; ok {
+		return rec.sess
+	}
+	return nil
 }
 
 func (s *InMemorySessionManager) retrieveByScreenNames(screenNames []IdentScreenName) []*Session {
@@ -113,9 +152,9 @@ func (s *InMemorySessionManager) retrieveByScreenNames(screenNames []IdentScreen
 	defer s.mapMutex.RUnlock()
 	var ret []*Session
 	for _, sn := range screenNames {
-		for _, sess := range s.store {
-			if sn == sess.IdentScreenName() {
-				ret = append(ret, sess)
+		for _, rec := range s.store {
+			if sn == rec.sess.IdentScreenName() {
+				ret = append(ret, rec.sess)
 			}
 		}
 	}
@@ -134,8 +173,8 @@ func (s *InMemorySessionManager) AllSessions() []*Session {
 	s.mapMutex.RLock()
 	defer s.mapMutex.RUnlock()
 	var sessions []*Session
-	for _, sess := range s.store {
-		sessions = append(sessions, sess)
+	for _, rec := range s.store {
+		sessions = append(sessions, rec.sess)
 	}
 	return sessions
 }
@@ -159,8 +198,8 @@ type InMemoryChatSessionManager struct {
 }
 
 // AddSession adds a user to a chat room. If screenName already exists, the old
-// session is closed replaced by a new one.
-func (s *InMemoryChatSessionManager) AddSession(chatCookie string, screenName DisplayScreenName) *Session {
+// session is replaced by a new one.
+func (s *InMemoryChatSessionManager) AddSession(ctx context.Context, chatCookie string, screenName DisplayScreenName) (*Session, error) {
 	s.mapMutex.Lock()
 	defer s.mapMutex.Unlock()
 
@@ -170,10 +209,14 @@ func (s *InMemoryChatSessionManager) AddSession(chatCookie string, screenName Di
 
 	sessionManager := s.store[chatCookie]
 
-	sess := sessionManager.AddSession(screenName)
+	sess, err := sessionManager.AddSession(ctx, screenName)
+	if err != nil {
+		return nil, fmt.Errorf("AddSession: %w", err)
+	}
+
 	sess.SetChatRoomCookie(chatCookie)
 
-	return sess
+	return sess, nil
 }
 
 // RemoveSession removes a user session from a chat room. It panics if you

+ 124 - 43
state/session_manager_test.go

@@ -3,6 +3,7 @@ package state
 import (
 	"context"
 	"log/slog"
+	"sync"
 	"testing"
 
 	"github.com/mk6i/retro-aim-server/wire"
@@ -13,25 +14,71 @@ import (
 func TestInMemorySessionManager_AddSession(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	want1 := sm.AddSession("user-screen-name")
-	have1 := sm.RetrieveSession(NewIdentScreenName("user-screen-name"))
-	assert.Same(t, want1, have1)
+	ctx := context.Background()
+	sess1, err := sm.AddSession(ctx, "user-screen-name")
+	assert.NoError(t, err)
 
-	want2 := sm.AddSession("user-screen-name")
-	have2 := sm.RetrieveSession(NewIdentScreenName("user-screen-name"))
-	assert.Same(t, want2, have2)
+	go func() {
+		<-sess1.Closed()
+		sm.RemoveSession(sess1)
+	}()
 
-	// ensure that the second session created with the same screen name as the
-	// first session clobbers the previous session in the session manager store
-	assert.NotSame(t, have1, have2)
+	sess2, err := sm.AddSession(ctx, "user-screen-name")
+	assert.NoError(t, err)
+
+	assert.NotSame(t, sess1, sess2)
+	assert.Contains(t, sm.AllSessions(), sess2)
+}
+
+func TestInMemorySessionManager_AddSession_Timeout(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	ctx, cancel := context.WithCancel(context.Background())
+	sess1, err := sm.AddSession(ctx, "user-screen-name")
+	assert.NoError(t, err)
+
+	go func() {
+		<-sess1.Closed()
+		cancel()
+	}()
+
+	sess2, err := sm.AddSession(ctx, "user-screen-name")
+	assert.Nil(t, sess2)
+	assert.ErrorIs(t, err, context.Canceled)
+}
+
+func TestInMemorySessionManager_AddSession_SessionConflict(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	ctx := context.Background()
+	sess1, err := sm.AddSession(ctx, "user-screen-name")
+	assert.NoError(t, err)
+
+	go func() {
+		<-sess1.Closed()
+		rec, ok := sm.store[NewIdentScreenName("user-screen-name")]
+		if assert.True(t, ok) {
+			close(rec.removed)
+		}
+	}()
+
+	sess2, err := sm.AddSession(ctx, "user-screen-name")
+	assert.Nil(t, sess2)
+	assert.ErrorIs(t, err, errSessConflict)
 }
 
 func TestInMemorySessionManager_Remove_Existing(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1Old := sm.AddSession("user-screen-name-1")
-	user1New := sm.AddSession("user-screen-name-1")
-	user2 := sm.AddSession("user-screen-name-2")
+	user1Old, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	sm.RemoveSession(user1Old)
+
+	user1New, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
 
 	sm.RemoveSession(user1New)
 
@@ -45,9 +92,15 @@ func TestInMemorySessionManager_Remove_Existing(t *testing.T) {
 func TestInMemorySessionManager_Remove_MissingSameScreenName(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1Old := sm.AddSession("user-screen-name-1")
-	user1New := sm.AddSession("user-screen-name-1")
-	user2 := sm.AddSession("user-screen-name-2")
+	user1Old, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	sm.RemoveSession(user1Old)
+
+	user1New, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
 
 	sm.RemoveSession(user1Old)
 
@@ -82,7 +135,8 @@ func TestInMemorySessionManager_Empty(t *testing.T) {
 			sm := NewInMemorySessionManager(slog.Default())
 
 			for _, screenName := range tt.given {
-				sm.AddSession(screenName)
+				_, err := sm.AddSession(context.Background(), screenName)
+				assert.NoError(t, err)
 			}
 
 			have := sm.Empty()
@@ -119,7 +173,8 @@ func TestInMemorySessionManager_Retrieve(t *testing.T) {
 			sm := NewInMemorySessionManager(slog.Default())
 
 			for _, screenName := range tt.given {
-				sm.AddSession(screenName)
+				_, err := sm.AddSession(context.Background(), screenName)
+				assert.NoError(t, err)
 			}
 
 			have := sm.RetrieveSession(tt.lookupScreenName)
@@ -135,9 +190,12 @@ func TestInMemorySessionManager_Retrieve(t *testing.T) {
 func TestInMemorySessionManager_RelayToScreenNames(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1 := sm.AddSession("user-screen-name-1")
-	user2 := sm.AddSession("user-screen-name-2")
-	user3 := sm.AddSession("user-screen-name-3")
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
+	user3, err := sm.AddSession(context.Background(), "user-screen-name-3")
+	assert.NoError(t, err)
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -167,8 +225,10 @@ func TestInMemorySessionManager_RelayToScreenNames(t *testing.T) {
 func TestInMemorySessionManager_Broadcast(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1 := sm.AddSession("user-screen-name-1")
-	user2 := sm.AddSession("user-screen-name-2")
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -188,8 +248,10 @@ func TestInMemorySessionManager_Broadcast(t *testing.T) {
 func TestInMemorySessionManager_Broadcast_SkipClosedSession(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1 := sm.AddSession("user-screen-name-1")
-	user2 := sm.AddSession("user-screen-name-2")
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
 	user2.Close()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
@@ -211,8 +273,10 @@ func TestInMemorySessionManager_Broadcast_SkipClosedSession(t *testing.T) {
 func TestInMemorySessionManager_RelayToScreenName_SessionExists(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1 := sm.AddSession("user-screen-name-1")
-	user2 := sm.AddSession("user-screen-name-2")
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -234,7 +298,8 @@ func TestInMemorySessionManager_RelayToScreenName_SessionExists(t *testing.T) {
 func TestInMemorySessionManager_RelayToScreenName_SessionNotExist(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1 := sm.AddSession("user-screen-name-1")
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -251,7 +316,8 @@ func TestInMemorySessionManager_RelayToScreenName_SessionNotExist(t *testing.T)
 func TestInMemorySessionManager_RelayToScreenName_SkipFullSession(t *testing.T) {
 	sm := NewInMemorySessionManager(slog.Default())
 
-	user1 := sm.AddSession("user-screen-name-1")
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
 	msg := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
 	wantCount := 0
@@ -283,9 +349,12 @@ func TestInMemoryChatSessionManager_RelayToAllExcept_HappyPath(t *testing.T) {
 	sm := NewInMemoryChatSessionManager(slog.Default())
 
 	cookie := "the-cookie"
-	user1 := sm.AddSession(cookie, "user-screen-name-1")
-	user2 := sm.AddSession(cookie, "user-screen-name-2")
-	user3 := sm.AddSession(cookie, "user-screen-name-3")
+	user1, err := sm.AddSession(context.Background(), cookie, "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), cookie, "user-screen-name-2")
+	assert.NoError(t, err)
+	user3, err := sm.AddSession(context.Background(), cookie, "user-screen-name-3")
+	assert.NoError(t, err)
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -311,8 +380,10 @@ func TestInMemoryChatSessionManager_RelayToAllExcept_HappyPath(t *testing.T) {
 func TestInMemoryChatSessionManager_AllSessions_RoomExists(t *testing.T) {
 	sm := NewInMemoryChatSessionManager(slog.Default())
 
-	user1 := sm.AddSession("the-cookie", "user-screen-name-1")
-	user2 := sm.AddSession("the-cookie", "user-screen-name-2")
+	user1, err := sm.AddSession(context.Background(), "the-cookie", "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "the-cookie", "user-screen-name-2")
+	assert.NoError(t, err)
 
 	sessions := sm.AllSessions("the-cookie")
 	assert.Len(t, sessions, 2)
@@ -329,8 +400,10 @@ func TestInMemoryChatSessionManager_AllSessions_RoomExists(t *testing.T) {
 func TestInMemoryChatSessionManager_RelayToScreenName_SessionAndChatRoomExist(t *testing.T) {
 	sm := NewInMemoryChatSessionManager(slog.Default())
 
-	user1 := sm.AddSession("chat-room-1", "user-screen-name-1")
-	user2 := sm.AddSession("chat-room-1", "user-screen-name-2")
+	user1, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-2")
+	assert.NoError(t, err)
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -352,8 +425,10 @@ func TestInMemoryChatSessionManager_RelayToScreenName_SessionAndChatRoomExist(t
 func TestInMemoryChatSessionManager_RemoveSession(t *testing.T) {
 	sm := NewInMemoryChatSessionManager(slog.Default())
 
-	user1 := sm.AddSession("chat-room-1", "user-screen-name-1")
-	user2 := sm.AddSession("chat-room-1", "user-screen-name-2")
+	user1, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
+	assert.NoError(t, err)
+	user2, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-2")
+	assert.NoError(t, err)
 
 	assert.Len(t, sm.AllSessions("chat-room-1"), 2)
 
@@ -366,14 +441,20 @@ func TestInMemoryChatSessionManager_RemoveSession(t *testing.T) {
 func TestInMemoryChatSessionManager_RemoveSession_DoubleLogin(t *testing.T) {
 	sm := NewInMemoryChatSessionManager(slog.Default())
 
-	user1 := sm.AddSession("chat-room-1", "user-screen-name-1")
-	user2 := sm.AddSession("chat-room-1", "user-screen-name-1")
-	assert.NotSame(t, user1, user2)
+	user1, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
+	assert.NoError(t, err)
 
-	assert.Len(t, sm.AllSessions("chat-room-1"), 1)
+	var wg sync.WaitGroup
+	wg.Add(1)
+	go func() {
+		defer wg.Done()
+		user2, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
+		assert.NoError(t, err)
+		assert.NotSame(t, user1, user2)
+	}()
 
 	sm.RemoveSession(user1)
-	sm.RemoveSession(user2)
+	wg.Wait()
 
-	assert.Empty(t, sm.AllSessions("chat-room-1"))
+	assert.Len(t, sm.AllSessions("chat-room-1"), 1)
 }