Просмотр исходного кода

skip incomplete signon sessions in manager operations

This fixes a bug where pollers to the GET /session endpoint can send
messages to users who have not yet signed on, causing them to crash.
Mike 9 месяцев назад
Родитель
Сommit
5ae9c42269
2 измененных файлов с 191 добавлено и 2 удалено
  1. 12 0
      state/session_manager.go
  2. 179 2
      state/session_manager_test.go

+ 12 - 0
state/session_manager.go

@@ -40,6 +40,9 @@ func (s *InMemorySessionManager) RelayToAll(ctx context.Context, msg wire.SNACMe
 	s.mapMutex.RLock()
 	defer s.mapMutex.RUnlock()
 	for _, rec := range s.store {
+		if !rec.sess.SignonComplete() {
+			continue
+		}
 		s.maybeRelayMessage(ctx, msg, rec.sess)
 	}
 }
@@ -137,6 +140,9 @@ func (s *InMemorySessionManager) RetrieveSession(screenName IdentScreenName) *Se
 	s.mapMutex.RLock()
 	defer s.mapMutex.RUnlock()
 	if rec, ok := s.store[screenName]; ok {
+		if !rec.sess.SignonComplete() {
+			return nil
+		}
 		return rec.sess
 	}
 	return nil
@@ -148,6 +154,9 @@ func (s *InMemorySessionManager) retrieveByScreenNames(screenNames []IdentScreen
 	var ret []*Session
 	for _, sn := range screenNames {
 		for _, rec := range s.store {
+			if !rec.sess.SignonComplete() {
+				continue
+			}
 			if sn == rec.sess.IdentScreenName() {
 				ret = append(ret, rec.sess)
 			}
@@ -169,6 +178,9 @@ func (s *InMemorySessionManager) AllSessions() []*Session {
 	defer s.mapMutex.RUnlock()
 	var sessions []*Session
 	for _, rec := range s.store {
+		if !rec.sess.SignonComplete() {
+			continue
+		}
 		sessions = append(sessions, rec.sess)
 	}
 	return sessions

+ 179 - 2
state/session_manager_test.go

@@ -17,6 +17,7 @@ func TestInMemorySessionManager_AddSession(t *testing.T) {
 	ctx := context.Background()
 	sess1, err := sm.AddSession(ctx, "user-screen-name")
 	assert.NoError(t, err)
+	sess1.SetSignonComplete()
 
 	go func() {
 		<-sess1.Closed()
@@ -25,6 +26,7 @@ func TestInMemorySessionManager_AddSession(t *testing.T) {
 
 	sess2, err := sm.AddSession(ctx, "user-screen-name")
 	assert.NoError(t, err)
+	sess2.SetSignonComplete()
 
 	assert.NotSame(t, sess1, sess2)
 	assert.Contains(t, sm.AllSessions(), sess2)
@@ -36,6 +38,7 @@ func TestInMemorySessionManager_AddSession_Timeout(t *testing.T) {
 	ctx, cancel := context.WithCancel(context.Background())
 	sess1, err := sm.AddSession(ctx, "user-screen-name")
 	assert.NoError(t, err)
+	sess1.SetSignonComplete()
 
 	go func() {
 		<-sess1.Closed()
@@ -53,6 +56,7 @@ func TestInMemorySessionManager_AddSession_SessionConflict(t *testing.T) {
 	ctx := context.Background()
 	sess1, err := sm.AddSession(ctx, "user-screen-name")
 	assert.NoError(t, err)
+	sess1.SetSignonComplete()
 
 	go func() {
 		<-sess1.Closed()
@@ -76,9 +80,11 @@ func TestInMemorySessionManager_Remove_Existing(t *testing.T) {
 
 	user1New, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1New.SetSignonComplete()
 
 	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	sm.RemoveSession(user1New)
 
@@ -98,9 +104,11 @@ func TestInMemorySessionManager_Remove_MissingSameScreenName(t *testing.T) {
 
 	user1New, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1New.SetSignonComplete()
 
 	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	sm.RemoveSession(user1Old)
 
@@ -135,8 +143,9 @@ func TestInMemorySessionManager_Empty(t *testing.T) {
 			sm := NewInMemorySessionManager(slog.Default())
 
 			for _, screenName := range tt.given {
-				_, err := sm.AddSession(context.Background(), screenName)
+				sess, err := sm.AddSession(context.Background(), screenName)
 				assert.NoError(t, err)
+				sess.SetSignonComplete()
 			}
 
 			have := sm.Empty()
@@ -173,8 +182,9 @@ func TestInMemorySessionManager_Retrieve(t *testing.T) {
 			sm := NewInMemorySessionManager(slog.Default())
 
 			for _, screenName := range tt.given {
-				_, err := sm.AddSession(context.Background(), screenName)
+				sess, err := sm.AddSession(context.Background(), screenName)
 				assert.NoError(t, err)
+				sess.SetSignonComplete()
 			}
 
 			have := sm.RetrieveSession(tt.lookupScreenName)
@@ -192,10 +202,13 @@ func TestInMemorySessionManager_RelayToScreenNames(t *testing.T) {
 
 	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 	user3, err := sm.AddSession(context.Background(), "user-screen-name-3")
 	assert.NoError(t, err)
+	user3.SetSignonComplete()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -227,8 +240,10 @@ func TestInMemorySessionManager_Broadcast(t *testing.T) {
 
 	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -250,8 +265,10 @@ func TestInMemorySessionManager_Broadcast_SkipClosedSession(t *testing.T) {
 
 	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 	user2.Close()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
@@ -275,8 +292,10 @@ func TestInMemorySessionManager_RelayToScreenName_SessionExists(t *testing.T) {
 
 	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -300,6 +319,7 @@ func TestInMemorySessionManager_RelayToScreenName_SessionNotExist(t *testing.T)
 
 	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -318,6 +338,7 @@ func TestInMemorySessionManager_RelayToScreenName_SkipFullSession(t *testing.T)
 
 	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	msg := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
 	wantCount := 0
@@ -351,10 +372,13 @@ func TestInMemoryChatSessionManager_RelayToAllExcept_HappyPath(t *testing.T) {
 	cookie := "the-cookie"
 	user1, err := sm.AddSession(context.Background(), cookie, "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), cookie, "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 	user3, err := sm.AddSession(context.Background(), cookie, "user-screen-name-3")
 	assert.NoError(t, err)
+	user3.SetSignonComplete()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -382,8 +406,10 @@ func TestInMemoryChatSessionManager_AllSessions_RoomExists(t *testing.T) {
 
 	user1, err := sm.AddSession(context.Background(), "the-cookie", "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "the-cookie", "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	sessions := sm.AllSessions("the-cookie")
 	assert.Len(t, sessions, 2)
@@ -402,8 +428,10 @@ func TestInMemoryChatSessionManager_RelayToScreenName_SessionAndChatRoomExist(t
 
 	user1, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
 
@@ -427,8 +455,10 @@ func TestInMemoryChatSessionManager_RemoveSession(t *testing.T) {
 
 	user1, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
 	assert.NoError(t, err)
+	user1.SetSignonComplete()
 	user2, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-2")
 	assert.NoError(t, err)
+	user2.SetSignonComplete()
 
 	assert.Len(t, sm.AllSessions("chat-room-1"), 2)
 
@@ -443,6 +473,7 @@ func TestInMemoryChatSessionManager_RemoveSession_DoubleLogin(t *testing.T) {
 
 	chatSess1, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
 	assert.NoError(t, err)
+	chatSess1.SetSignonComplete()
 
 	wg := &sync.WaitGroup{}
 	wg.Add(1)
@@ -452,6 +483,7 @@ func TestInMemoryChatSessionManager_RemoveSession_DoubleLogin(t *testing.T) {
 		// room for the new session
 		chatSess2, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
 		assert.NoError(t, err)
+		chatSess2.SetSignonComplete()
 		assert.Equal(t, chatSess1.DisplayScreenName(), chatSess2.DisplayScreenName())
 		wg.Done()
 	}()
@@ -474,8 +506,10 @@ func TestInMemoryChatSessionManager_RemoveUserFromAllChats(t *testing.T) {
 	user1 := NewIdentScreenName("user-screen-name-1")
 	user1sess, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-1")
 	assert.NoError(t, err)
+	user1sess.SetSignonComplete()
 	user2sess, err := sm.AddSession(context.Background(), "chat-room-1", "user-screen-name-2")
 	assert.NoError(t, err)
+	user2sess.SetSignonComplete()
 
 	assert.Len(t, sm.AllSessions("chat-room-1"), 2)
 
@@ -490,3 +524,146 @@ func TestInMemoryChatSessionManager_RemoveUserFromAllChats(t *testing.T) {
 	assert.True(t, lookup[user2sess])
 
 }
+
+func TestInMemorySessionManager_RelayToAll_SkipIncompleteSignon(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user1.SetSignonComplete()
+
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
+	// user2 has not completed signon
+
+	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
+
+	sm.RelayToAll(context.Background(), want)
+
+	select {
+	case have := <-user1.ReceiveMessage():
+		assert.Equal(t, want, have)
+	}
+
+	select {
+	case <-user2.ReceiveMessage():
+		assert.Fail(t, "user 2 should not receive a message because signon is incomplete")
+	default:
+	}
+}
+
+func TestInMemorySessionManager_RetrieveSession_IncompleteSignon(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	// user1 has not completed signon
+
+	sess := sm.RetrieveSession(NewIdentScreenName("user-screen-name-1"))
+	assert.Nil(t, sess, "should return nil for session with incomplete signon")
+
+	user1.SetSignonComplete()
+	sess = sm.RetrieveSession(NewIdentScreenName("user-screen-name-1"))
+	assert.NotNil(t, sess, "should return session after signon is complete")
+	assert.Equal(t, user1, sess)
+}
+
+func TestInMemorySessionManager_RetrieveSession_CompleteSignon(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user1.SetSignonComplete()
+
+	sess := sm.RetrieveSession(NewIdentScreenName("user-screen-name-1"))
+	assert.NotNil(t, sess)
+	assert.Equal(t, user1, sess)
+}
+
+func TestInMemorySessionManager_RelayToScreenNames_SkipIncompleteSignon(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user1.SetSignonComplete()
+
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
+	// user2 has not completed signon
+
+	user3, err := sm.AddSession(context.Background(), "user-screen-name-3")
+	assert.NoError(t, err)
+	user3.SetSignonComplete()
+
+	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
+
+	recips := []IdentScreenName{
+		NewIdentScreenName("user-screen-name-1"),
+		NewIdentScreenName("user-screen-name-2"), // incomplete signon
+		NewIdentScreenName("user-screen-name-3"),
+	}
+	sm.RelayToScreenNames(context.Background(), recips, want)
+
+	select {
+	case have := <-user1.ReceiveMessage():
+		assert.Equal(t, want, have)
+	}
+
+	select {
+	case <-user2.ReceiveMessage():
+		assert.Fail(t, "user 2 should not receive a message because signon is incomplete")
+	default:
+	}
+
+	select {
+	case have := <-user3.ReceiveMessage():
+		assert.Equal(t, want, have)
+	}
+}
+
+func TestInMemorySessionManager_AllSessions_SkipIncompleteSignon(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	user1.SetSignonComplete()
+
+	user2, err := sm.AddSession(context.Background(), "user-screen-name-2")
+	assert.NoError(t, err)
+	// user2 has not completed signon
+
+	user3, err := sm.AddSession(context.Background(), "user-screen-name-3")
+	assert.NoError(t, err)
+	user3.SetSignonComplete()
+
+	sessions := sm.AllSessions()
+	assert.Len(t, sessions, 2, "should only return sessions with complete signon")
+
+	lookup := make(map[*Session]bool)
+	for _, session := range sessions {
+		lookup[session] = true
+	}
+
+	assert.True(t, lookup[user1], "user1 should be included (complete signon)")
+	assert.False(t, lookup[user2], "user2 should not be included (incomplete signon)")
+	assert.True(t, lookup[user3], "user3 should be included (complete signon)")
+}
+
+func TestInMemorySessionManager_RelayToScreenName_IncompleteSignon(t *testing.T) {
+	sm := NewInMemorySessionManager(slog.Default())
+
+	user1, err := sm.AddSession(context.Background(), "user-screen-name-1")
+	assert.NoError(t, err)
+	// user1 has not completed signon
+
+	want := wire.SNACMessage{Frame: wire.SNACFrame{FoodGroup: wire.ICBM}}
+
+	recip := NewIdentScreenName("user-screen-name-1")
+	sm.RelayToScreenName(context.Background(), recip, want)
+
+	select {
+	case <-user1.ReceiveMessage():
+		assert.Fail(t, "user 1 should not receive a message because signon is incomplete")
+	default:
+	}
+}