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

preserve rate limits caused by warnings between sessions

Mike 9 месяцев назад
Родитель
Сommit
0b6306f270
8 измененных файлов с 317 добавлено и 245 удалено
  1. 0 3
      cmd/server/factory.go
  2. 9 0
      foodgroup/helpers_test.go
  3. 11 11
      foodgroup/icbm.go
  4. 15 3
      foodgroup/icbm_test.go
  5. 2 5
      foodgroup/oservice.go
  6. 55 54
      foodgroup/oservice_test.go
  7. 51 46
      state/session.go
  8. 174 123
      state/session_test.go

+ 0 - 3
cmd/server/factory.go

@@ -284,7 +284,6 @@ func OSCAR(deps Container) *oscar.Server {
 		deps.sqLiteUserStore,
 		deps.inMemorySessionManager,
 		deps.sqLiteUserStore,
-		deps.rateLimitClasses,
 		deps.snacRateLimits,
 		deps.chatSessionManager,
 	)
@@ -421,7 +420,6 @@ func TOC(deps Container) *toc.Server {
 				deps.sqLiteUserStore,
 				deps.inMemorySessionManager,
 				deps.sqLiteUserStore,
-				deps.rateLimitClasses,
 				deps.snacRateLimits,
 				deps.chatSessionManager,
 			),
@@ -517,7 +515,6 @@ func WebAPI(deps Container) *webapi.Server {
 			deps.sqLiteUserStore,
 			deps.inMemorySessionManager,
 			deps.sqLiteUserStore,
-			deps.rateLimitClasses,
 			deps.snacRateLimits,
 			deps.chatSessionManager,
 		),

+ 9 - 0
foodgroup/helpers_test.go

@@ -782,6 +782,7 @@ func sessOptWantTypingEvents(session *state.Session) {
 	session.SetTypingEventsEnabled(true)
 }
 
+// sessOptSetFoodGroupVersion sets food group versions
 func sessOptSetFoodGroupVersion(foodGroup uint16, version uint16) func(session *state.Session) {
 	return func(session *state.Session) {
 		var versions [wire.MDir + 1]uint16
@@ -790,6 +791,13 @@ func sessOptSetFoodGroupVersion(foodGroup uint16, version uint16) func(session *
 	}
 }
 
+// sessOptSetRateClasses sets rate limit classes
+func sessOptSetRateClasses(classes wire.RateLimitClasses) func(session *state.Session) {
+	return func(session *state.Session) {
+		session.SetRateClasses(time.Now(), classes)
+	}
+}
+
 // sessClientID sets the client ID
 func sessClientID(clientID string) func(session *state.Session) {
 	return func(session *state.Session) {
@@ -810,6 +818,7 @@ func newTestSession(screenName state.DisplayScreenName, options ...func(session
 	s := state.NewSession()
 	s.SetIdentScreenName(screenName.IdentScreenName())
 	s.SetDisplayScreenName(screenName)
+	s.SetRateClasses(time.Now(), wire.DefaultRateLimitClasses())
 	for _, op := range options {
 		op(s)
 	}

+ 11 - 11
foodgroup/icbm.go

@@ -15,9 +15,10 @@ import (
 )
 
 const (
-	evilDelta       = uint16(100)
-	evilDeltaAnon   = uint16(30)
-	warningDecayPct = -50
+	evilDelta         = uint16(100)
+	evilDeltaAnon     = uint16(30)
+	warningDecayPct   = -50
+	rateDecayInterval = 5 * time.Minute
 )
 
 // NewICBMService returns a new instance of ICBMService.
@@ -42,7 +43,7 @@ func NewICBMService(
 		snacRateLimits:      snacRateLimits,
 		convoTracker:        newConvoTracker(),
 		logger:              logger,
-		interval:            5 * time.Minute,
+		interval:            rateDecayInterval,
 	}
 }
 
@@ -327,7 +328,7 @@ func (s ICBMService) EvilRequest(ctx context.Context, sess *state.Session, inFra
 		panic("failed to retrieve rate class for ICBMChannelMsgToHost")
 	}
 
-	ok, newLevel := recipSess.IncrementWarning(int16(increase), classID)
+	ok, newLevel := recipSess.ScaleWarningAndRateLimit(int16(increase), classID)
 	if !ok {
 		// target's warning is at 100%
 		return *newICBMErr(inFrame.RequestID, wire.ErrorCodeRequestDenied), nil
@@ -386,10 +387,6 @@ func (s ICBMService) RestoreWarningLevel(ctx context.Context, sess *state.Sessio
 		return nil
 	}
 
-	sess.SetWarning(u.LastWarnLevel)
-
-	warnDelta := calcElapsedWarningLevel(u.LastWarnUpdate, s.timeNow(), s.interval)
-
 	// get the rate class for sending IMs, which gets limited when the user gets warned
 	classID, ok := s.snacRateLimits.RateClassLookup(wire.ICBM, wire.ICBMChannelMsgToHost)
 	if !ok {
@@ -398,7 +395,10 @@ func (s ICBMService) RestoreWarningLevel(ctx context.Context, sess *state.Sessio
 
 	// increment warning level by the amount of time that has passed since last
 	// login, proportionally increasing the warning level
-	sess.IncrementWarning(warnDelta, classID)
+	warnDelta := calcElapsedWarningLevel(u.LastWarnUpdate, s.timeNow(), s.interval)
+	newWarning := int16(u.LastWarnLevel) + warnDelta
+	sess.SetWarning(0)
+	sess.ScaleWarningAndRateLimit(newWarning, classID)
 
 	if sess.Warning() > 0 {
 		s.logger.DebugContext(ctx, "restored warning level with time decay applied since last login",
@@ -529,7 +529,7 @@ func (s ICBMService) UpdateWarnLevel(ctx context.Context, sess *state.Session) {
 				doReset = false
 			}
 
-			ok, warning := sess.IncrementWarning(warningDecayPct, classID)
+			ok, warning := sess.ScaleWarningAndRateLimit(warningDecayPct, classID)
 			if !ok {
 				s.logger.ErrorContext(ctx, "warning increment out of rage", "level", warning)
 				stopTicker()

+ 15 - 3
foodgroup/icbm_test.go

@@ -1763,19 +1763,19 @@ func TestICBMService_UpdateWarnLevel(t *testing.T) {
 			svc.UpdateWarnLevel(ctx, sess) // do a sync test here?
 		}()
 
-		ok, _ := sess.IncrementWarning(100, 3)
+		ok, _ := sess.ScaleWarningAndRateLimit(100, 3)
 		assert.True(t, ok)
 		assert.Equal(t, uint16(100), <-warnCh)
 		assert.Equal(t, uint16(50), <-warnCh)
 		assert.Equal(t, uint16(0), <-warnCh)
 
-		ok, _ = sess.IncrementWarning(100, 3)
+		ok, _ = sess.ScaleWarningAndRateLimit(100, 3)
 		assert.True(t, ok)
 		assert.Equal(t, uint16(100), <-warnCh)
 		assert.Equal(t, uint16(50), <-warnCh)
 		assert.Equal(t, uint16(0), <-warnCh)
 
-		sess.IncrementWarning(30, 3)
+		sess.ScaleWarningAndRateLimit(30, 3)
 		assert.Equal(t, uint16(30), <-warnCh)
 		assert.Equal(t, uint16(0), <-warnCh)
 
@@ -1849,10 +1849,22 @@ func TestICBMService_RestoreWarningLevel(t *testing.T) {
 			ctx, cancel := context.WithCancel(context.Background())
 			defer cancel()
 
+			statesBefore := sess.RateLimitStates()
+
 			err := svc.RestoreWarningLevel(ctx, sess)
 			assert.NoError(t, err)
 
 			assert.Equal(t, tt.expectedWarn, sess.Warning())
+
+			statesAfter := sess.RateLimitStates()
+
+			if tt.expectedWarn > 0 {
+				// make sure the rate limits changed
+				assert.NotEqual(t, statesBefore, statesAfter)
+			} else {
+				// make sure the rate limits have been restored
+				assert.Equal(t, statesBefore, statesAfter)
+			}
 		})
 	}
 }

+ 2 - 5
foodgroup/oservice.go

@@ -19,7 +19,6 @@ type OServiceService struct {
 	buddyBroadcaster buddyBroadcaster
 	cfg              config.Config // todo remove
 	logger           *slog.Logger
-	rateLimitClasses wire.RateLimitClasses
 	snacRateLimits   wire.SNACRateLimits
 	timeNow          func() time.Time
 
@@ -39,7 +38,6 @@ func NewOServiceService(
 	relationshipFetcher RelationshipFetcher,
 	sessionRetriever SessionRetriever,
 	bartItemManager BARTItemManager,
-	rateLimitClasses wire.RateLimitClasses,
 	snacRateLimits wire.SNACRateLimits,
 	chatMessageRelayer ChatMessageRelayer,
 ) *OServiceService {
@@ -49,7 +47,6 @@ func NewOServiceService(
 		buddyBroadcaster:   newBuddyNotifier(bartItemManager, relationshipFetcher, messageRelayer, sessionRetriever),
 		cfg:                cfg,
 		logger:             logger,
-		rateLimitClasses:   rateLimitClasses,
 		snacRateLimits:     snacRateLimits,
 		timeNow:            time.Now,
 		chatRoomManager:    chatRoomManager,
@@ -170,7 +167,7 @@ func (s OServiceService) RateParamsQuery(ctx context.Context, sess *state.Sessio
 		},
 	}
 
-	for _, class := range s.rateLimitClasses.All() {
+	for _, class := range sess.RateLimitStates() {
 		str := wire.RateParamsSNAC{
 			ID:              uint16(class.ID),
 			WindowSize:      uint32(class.WindowSize),
@@ -178,7 +175,7 @@ func (s OServiceService) RateParamsQuery(ctx context.Context, sess *state.Sessio
 			AlertLevel:      uint32(class.AlertLevel),
 			LimitLevel:      uint32(class.LimitLevel),
 			DisconnectLevel: uint32(class.DisconnectLevel),
-			CurrentLevel:    uint32(class.MaxLevel),
+			CurrentLevel:    uint32(class.CurrentLevel),
 			MaxLevel:        uint32(class.MaxLevel),
 		}
 		if sess.FoodGroupVersions()[wire.OService] > 1 {

+ 55 - 54
foodgroup/oservice_test.go

@@ -801,7 +801,7 @@ func TestOServiceService_ServiceRequest(t *testing.T) {
 			//
 			// send input SNAC
 			//
-			svc := NewOServiceService(config.Config{}, nil, slog.Default(), cookieIssuer, chatRoomManager, nil, nil, nil, wire.DefaultRateLimitClasses(), wire.DefaultSNACRateLimits(), chatMessageRelayer)
+			svc := NewOServiceService(config.Config{}, nil, slog.Default(), cookieIssuer, chatRoomManager, nil, nil, nil, wire.DefaultSNACRateLimits(), chatMessageRelayer)
 
 			outputSNAC, err := svc.ServiceRequest(context.Background(), tc.service, tc.userSession, tc.inputSNAC.Frame,
 				tc.inputSNAC.Body.(wire.SNAC_0x01_0x04_OServiceServiceRequest), tc.listener)
@@ -1016,6 +1016,54 @@ func TestOServiceService_SetUserInfoFields(t *testing.T) {
 }
 
 func TestOServiceService_RateParamsQuery(t *testing.T) {
+	rateClasses := wire.NewRateLimitClasses([5]wire.RateClass{
+		{
+			ID:              1,
+			WindowSize:      80,
+			ClearLevel:      2500,
+			AlertLevel:      2000,
+			LimitLevel:      1500,
+			DisconnectLevel: 800,
+			MaxLevel:        6000,
+		},
+		{
+			ID:              2,
+			WindowSize:      80,
+			ClearLevel:      3000,
+			AlertLevel:      2000,
+			LimitLevel:      1500,
+			DisconnectLevel: 1000,
+			MaxLevel:        6000,
+		},
+		{
+			ID:              3,
+			WindowSize:      20,
+			ClearLevel:      5100,
+			AlertLevel:      5000,
+			LimitLevel:      4000,
+			DisconnectLevel: 3000,
+			MaxLevel:        6000,
+		},
+		{
+			ID:              4,
+			WindowSize:      20,
+			ClearLevel:      5500,
+			AlertLevel:      5300,
+			LimitLevel:      4200,
+			DisconnectLevel: 3000,
+			MaxLevel:        8000,
+		},
+		{
+			ID:              5,
+			WindowSize:      10,
+			ClearLevel:      5500,
+			AlertLevel:      5300,
+			LimitLevel:      4200,
+			DisconnectLevel: 3000,
+			MaxLevel:        8000,
+		},
+	})
+
 	expectRateGroups := []struct {
 		ID    uint16
 		Pairs []struct {
@@ -1302,7 +1350,7 @@ func TestOServiceService_RateParamsQuery(t *testing.T) {
 	}{
 		{
 			name:        "get rate limits for AIM > 1.x clients",
-			userSession: newTestSession("me", sessOptSetFoodGroupVersion(wire.OService, 3)),
+			userSession: newTestSession("me", sessOptSetFoodGroupVersion(wire.OService, 3), sessOptSetRateClasses(rateClasses)),
 			inputSNAC: wire.SNACMessage{
 				Frame: wire.SNACFrame{RequestID: 1234},
 			},
@@ -1409,7 +1457,7 @@ func TestOServiceService_RateParamsQuery(t *testing.T) {
 		},
 		{
 			name:        "get rate limits for AIM 1.x client",
-			userSession: newTestSession("me", sessClientID("AOL Instant Messenger (TM), version 1.")),
+			userSession: newTestSession("me", sessClientID("AOL Instant Messenger (TM), version 1."), sessOptSetRateClasses(rateClasses)),
 			inputSNAC: wire.SNACMessage{
 				Frame: wire.SNACFrame{RequestID: 1234},
 			},
@@ -1476,55 +1524,8 @@ func TestOServiceService_RateParamsQuery(t *testing.T) {
 	for _, tc := range cases {
 		t.Run(tc.name, func(t *testing.T) {
 			svc := OServiceService{
-				cfg:    config.Config{},
-				logger: slog.Default(),
-				rateLimitClasses: wire.NewRateLimitClasses([5]wire.RateClass{
-					{
-						ID:              1,
-						WindowSize:      80,
-						ClearLevel:      2500,
-						AlertLevel:      2000,
-						LimitLevel:      1500,
-						DisconnectLevel: 800,
-						MaxLevel:        6000,
-					},
-					{
-						ID:              2,
-						WindowSize:      80,
-						ClearLevel:      3000,
-						AlertLevel:      2000,
-						LimitLevel:      1500,
-						DisconnectLevel: 1000,
-						MaxLevel:        6000,
-					},
-					{
-						ID:              3,
-						WindowSize:      20,
-						ClearLevel:      5100,
-						AlertLevel:      5000,
-						LimitLevel:      4000,
-						DisconnectLevel: 3000,
-						MaxLevel:        6000,
-					},
-					{
-						ID:              4,
-						WindowSize:      20,
-						ClearLevel:      5500,
-						AlertLevel:      5300,
-						LimitLevel:      4200,
-						DisconnectLevel: 3000,
-						MaxLevel:        8000,
-					},
-					{
-						ID:              5,
-						WindowSize:      10,
-						ClearLevel:      5500,
-						AlertLevel:      5300,
-						LimitLevel:      4200,
-						DisconnectLevel: 3000,
-						MaxLevel:        8000,
-					},
-				}),
+				cfg:            config.Config{},
+				logger:         slog.Default(),
 				snacRateLimits: wire.DefaultSNACRateLimits(),
 				timeNow:        tc.timeNow,
 			}
@@ -1697,7 +1698,7 @@ func TestOServiceService_HostOnline(t *testing.T) {
 
 	for _, tc := range cases {
 		t.Run(tc.name, func(t *testing.T) {
-			svc := NewOServiceService(config.Config{}, nil, slog.Default(), nil, nil, nil, nil, nil, wire.DefaultRateLimitClasses(), wire.DefaultSNACRateLimits(), nil)
+			svc := NewOServiceService(config.Config{}, nil, slog.Default(), nil, nil, nil, nil, nil, wire.DefaultSNACRateLimits(), nil)
 			have := svc.HostOnline(tc.service)
 			assert.Equal(t, tc.expectOutput, have)
 		})
@@ -2026,7 +2027,7 @@ func TestOServiceService_ClientOnline(t *testing.T) {
 					RelayToScreenName(mock.Anything, params.cookie, params.screenName, params.message)
 			}
 
-			svc := NewOServiceService(config.Config{}, messageRelayer, slog.Default(), nil, chatRoomManager, nil, nil, nil, wire.DefaultRateLimitClasses(), wire.DefaultSNACRateLimits(), chatMessageRelayer)
+			svc := NewOServiceService(config.Config{}, messageRelayer, slog.Default(), nil, chatRoomManager, nil, nil, nil, wire.DefaultSNACRateLimits(), chatMessageRelayer)
 			svc.buddyBroadcaster = buddyUpdateBroadcaster
 			haveErr := svc.ClientOnline(context.Background(), tt.service, tt.bodyIn, tt.sess)
 			assert.ErrorIs(t, tt.wantErr, haveErr)

+ 51 - 46
state/session.go

@@ -49,34 +49,34 @@ const (
 // Session represents a user's current session. Unless stated otherwise, all
 // methods may be safely accessed by multiple goroutines.
 type Session struct {
-	awayMessage           string
-	caps                  [][16]byte
-	chatRoomCookie        string
-	clientID              string
-	closed                bool
-	displayScreenName     DisplayScreenName
-	foodGroupVersions     [wire.MDir + 1]uint16
-	identScreenName       IdentScreenName
-	idle                  bool
-	idleTime              time.Time
-	lastObservedStates    [5]RateClassState
-	msgCh                 chan wire.SNACMessage
-	multiConnFlag         wire.MultiConnFlag
-	mutex                 sync.RWMutex
-	nowFn                 func() time.Time
-	rateByClassID         [5]RateClassState
-	rateByClassIDOriginal [5]RateClassState
-	remoteAddr            *netip.AddrPort
-	signonComplete        bool
-	signonTime            time.Time
-	stopCh                chan struct{}
-	typingEventsEnabled   bool
-	uin                   uint32
-	userInfoBitmask       uint16
-	userStatusBitmask     uint32
-	warning               uint16
-	warningCh             chan uint16
-	lastWarnUpdate        time.Time
+	awayMessage             string
+	caps                    [][16]byte
+	chatRoomCookie          string
+	clientID                string
+	closed                  bool
+	displayScreenName       DisplayScreenName
+	foodGroupVersions       [wire.MDir + 1]uint16
+	identScreenName         IdentScreenName
+	idle                    bool
+	idleTime                time.Time
+	lastObservedStates      [5]RateClassState
+	msgCh                   chan wire.SNACMessage
+	multiConnFlag           wire.MultiConnFlag
+	mutex                   sync.RWMutex
+	nowFn                   func() time.Time
+	rateLimitStates         [5]RateClassState
+	rateLimitStatesOriginal [5]RateClassState
+	remoteAddr              *netip.AddrPort
+	signonComplete          bool
+	signonTime              time.Time
+	stopCh                  chan struct{}
+	typingEventsEnabled     bool
+	uin                     uint32
+	userInfoBitmask         uint16
+	userStatusBitmask       uint32
+	warning                 uint16
+	warningCh               chan uint16
+	lastWarnUpdate          time.Time
 }
 
 // NewSession returns a new instance of Session. By default, the user may have
@@ -141,11 +141,11 @@ func (s *Session) SetRateClasses(now time.Time, classes wire.RateLimitClasses) {
 	if s.lastObservedStates[0].ID == 0 {
 		s.lastObservedStates = newStates
 	} else {
-		s.lastObservedStates = s.rateByClassID
+		s.lastObservedStates = s.rateLimitStates
 	}
 
-	s.rateByClassID = newStates
-	s.rateByClassIDOriginal = newStates
+	s.rateLimitStates = newStates
+	s.rateLimitStatesOriginal = newStates
 }
 
 // SetRemoteAddr sets the user's remote IP address
@@ -185,6 +185,11 @@ func (s *Session) UserInfoBitmask() (flags uint16) {
 	return s.userInfoBitmask
 }
 
+// RateLimitStates returns the current session rate limits
+func (s *Session) RateLimitStates() [5]RateClassState {
+	return s.rateLimitStates
+}
+
 // SetUserStatusBitmask sets the user status bitmask from the client.
 func (s *Session) SetUserStatusBitmask(bitmask uint32) {
 	s.mutex.Lock()
@@ -199,11 +204,11 @@ func (s *Session) UserStatusBitmask() uint32 {
 	return s.userStatusBitmask
 }
 
-// IncrementWarning increments the user's warning level and scales rate limits accordingly.
+// ScaleWarningAndRateLimit increments the user's warning level and scales a rate limit accordingly.
 // The incr parameter is the warning increment (negative to decrease), and classID specifies
 // which rate limit class to scale. The incr param is a percentage represented as an integer
 // where 30 = 3.0%, 100 = 10.0%, etc.
-func (s *Session) IncrementWarning(incr int16, classID wire.RateLimitClassID) (bool, uint16) {
+func (s *Session) ScaleWarningAndRateLimit(incr int16, classID wire.RateLimitClassID) (bool, uint16) {
 	s.mutex.Lock()
 	defer s.mutex.Unlock()
 
@@ -221,8 +226,8 @@ func (s *Session) IncrementWarning(incr int16, classID wire.RateLimitClassID) (b
 	pct := float32(incr) / 1000.0
 
 	// create reference variables for better readability
-	rateClass := &s.rateByClassID[classID-1]
-	originalRateClass := &s.rateByClassIDOriginal[classID-1]
+	rateClass := &s.rateLimitStates[classID-1]
+	originalRateClass := &s.rateLimitStatesOriginal[classID-1]
 
 	// clamp function to constrain values between min and max
 	clamp := func(value, min, max int32) int32 {
@@ -546,7 +551,7 @@ func (s *Session) SubscribeRateLimits(classes []wire.RateLimitClassID) {
 	defer s.mutex.Unlock()
 
 	for _, classID := range classes {
-		s.rateByClassID[classID-1].Subscribed = true
+		s.rateLimitStates[classID-1].Subscribed = true
 	}
 }
 
@@ -556,32 +561,32 @@ func (s *Session) ObserveRateChanges(now time.Time) (classDelta []RateClassState
 	s.mutex.Lock()
 	defer s.mutex.Unlock()
 
-	for i, params := range s.rateByClassID {
+	for i, params := range s.rateLimitStates {
 		if !params.Subscribed {
 			continue
 		}
 
 		state, level := wire.CheckRateLimit(params.LastTime, now, params.RateClass, params.CurrentLevel, params.LimitedNow)
-		s.rateByClassID[i].CurrentStatus = state
+		s.rateLimitStates[i].CurrentStatus = state
 
 		// clear limited now flag if passing from limited state to clear state
-		if s.rateByClassID[i].LimitedNow && state == wire.RateLimitStatusClear {
-			s.rateByClassID[i].LimitedNow = false
-			s.rateByClassID[i].CurrentLevel = level
+		if s.rateLimitStates[i].LimitedNow && state == wire.RateLimitStatusClear {
+			s.rateLimitStates[i].LimitedNow = false
+			s.rateLimitStates[i].CurrentLevel = level
 		}
 
 		// did rate class change?
 		if params.RateClass != s.lastObservedStates[i].RateClass {
-			classDelta = append(classDelta, s.rateByClassID[i])
+			classDelta = append(classDelta, s.rateLimitStates[i])
 		}
 
 		// did rate limit status change?
-		if s.lastObservedStates[i].CurrentStatus != s.rateByClassID[i].CurrentStatus {
-			stateDelta = append(stateDelta, s.rateByClassID[i])
+		if s.lastObservedStates[i].CurrentStatus != s.rateLimitStates[i].CurrentStatus {
+			stateDelta = append(stateDelta, s.rateLimitStates[i])
 		}
 
 		// save it for next time
-		s.lastObservedStates[i] = s.rateByClassID[i]
+		s.lastObservedStates[i] = s.rateLimitStates[i]
 	}
 
 	return classDelta, stateDelta
@@ -599,7 +604,7 @@ func (s *Session) EvaluateRateLimit(now time.Time, rateClassID wire.RateLimitCla
 		return wire.RateLimitStatusClear // don't rate limit bots
 	}
 
-	rateClass := &s.rateByClassID[rateClassID-1]
+	rateClass := &s.rateLimitStates[rateClassID-1]
 
 	status, newLevel := wire.CheckRateLimit(rateClass.LastTime, now, rateClass.RateClass, rateClass.CurrentLevel, rateClass.LimitedNow)
 	rateClass.CurrentLevel = newLevel

+ 174 - 123
state/session_test.go

@@ -29,9 +29,9 @@ func TestSession_IncrementAndGetWarning(t *testing.T) {
 	wg.Add(1)
 	go func() {
 		defer wg.Done()
-		s.IncrementWarning(1, 1)
-		s.IncrementWarning(2, 1)
-		s.IncrementWarning(3, 1)
+		s.ScaleWarningAndRateLimit(1, 1)
+		s.ScaleWarningAndRateLimit(2, 1)
+		s.ScaleWarningAndRateLimit(3, 1)
 	}()
 
 	assert.Equal(t, uint16(1), <-s.WarningCh())
@@ -101,7 +101,7 @@ func TestSession_TLVUserInfo(t *testing.T) {
 				s.SetSignonTime(time.Unix(1, 0))
 				s.SetIdentScreenName(NewIdentScreenName("xXAIMUSERXx"))
 				s.SetDisplayScreenName("xXAIMUSERXx")
-				s.IncrementWarning(10, 1)
+				s.ScaleWarningAndRateLimit(10, 1)
 				s.SetUserInfoFlag(wire.OServiceUserFlagOSCARFree)
 				return s
 			},
@@ -529,7 +529,7 @@ func TestSession_EvaluateRateLimit_ObserveRateChanges(t *testing.T) {
 
 		// this is a rearranged moving average formula that determines how many
 		// milliseconds it will take to reach the clear threshold
-		timeToRecover := int(math.Ceil((time.Duration(rateClass.ClearLevel*rateClass.WindowSize-sess.rateByClassID[rateClass.ID-1].CurrentLevel*(rateClass.WindowSize-1)) * time.Millisecond).Seconds()))
+		timeToRecover := int(math.Ceil((time.Duration(rateClass.ClearLevel*rateClass.WindowSize-sess.rateLimitStates[rateClass.ID-1].CurrentLevel*(rateClass.WindowSize-1)) * time.Millisecond).Seconds()))
 		assert.True(t, timeToRecover > 0)
 
 		// indicate the time rate limiting kicked in
@@ -644,7 +644,7 @@ func TestSession_SetAndGetLastWarnLevel(t *testing.T) {
 	assert.Equal(t, level, s.Warning())
 }
 
-func TestSession_IncrementWarningWithRateLimitScaling(t *testing.T) {
+func TestSession_ScaleWarningAndRateLimit(t *testing.T) {
 	t.Run("scale up", func(t *testing.T) {
 		classParams := [5]wire.RateClass{
 			{},
@@ -683,64 +683,64 @@ func TestSession_IncrementWarningWithRateLimitScaling(t *testing.T) {
 			}
 		}()
 
-		assert.Equal(t, int32(5000), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5100), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4000), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5100), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5190), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4200), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5200), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5280), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4400), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5300), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5370), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4600), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5400), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5460), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4800), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5500), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5550), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5000), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5600), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5640), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5200), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5700), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5730), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5400), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5800), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5820), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5600), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(5900), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5910), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5800), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(100, 3)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].LimitLevel)
+		assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5100), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5190), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4200), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5200), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5280), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4400), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5300), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5370), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4600), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5400), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5460), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4800), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5500), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5550), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5600), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5640), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5200), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5700), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5730), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5400), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5800), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5820), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5600), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(5900), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5910), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5800), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(100, 3)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].LimitLevel)
 
 		cancel()
 		wg.Wait()
@@ -785,67 +785,118 @@ func TestSession_IncrementWarningWithRateLimitScaling(t *testing.T) {
 		}()
 
 		for i := 0; i < 10; i++ {
-			sess.IncrementWarning(100, 3)
+			sess.ScaleWarningAndRateLimit(100, 3)
 		}
 
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(6000), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5900), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5910), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5800), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5800), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5820), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5600), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5700), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5730), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5400), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5600), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5640), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5200), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5500), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5550), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(5000), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5400), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5460), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4800), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5300), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5370), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4600), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5200), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5280), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4400), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5100), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5190), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4200), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5000), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5100), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4000), sess.rateByClassID[2].LimitLevel)
-
-		sess.IncrementWarning(-100, 3)
-		assert.Equal(t, int32(5000), sess.rateByClassID[2].AlertLevel)
-		assert.Equal(t, int32(5100), sess.rateByClassID[2].ClearLevel)
-		assert.Equal(t, int32(4000), sess.rateByClassID[2].LimitLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5900), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5910), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5800), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5800), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5820), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5600), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5700), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5730), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5400), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5600), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5640), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5200), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5500), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5550), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(5000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5400), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5460), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4800), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5300), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5370), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4600), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5200), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5280), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4400), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5100), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5190), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4200), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(-100, 3)
+		assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
+
+		cancel()
+		wg.Wait()
+	})
+
+	t.Run("increment 100%", func(t *testing.T) {
+		classParams := [5]wire.RateClass{
+			{},
+			{},
+			{
+				ID:              3,
+				WindowSize:      20,
+				ClearLevel:      5100,
+				AlertLevel:      5000,
+				LimitLevel:      4000,
+				DisconnectLevel: 3000,
+				MaxLevel:        6000,
+			},
+			{},
+			{},
+		}
+		rateClasses := wire.NewRateLimitClasses(classParams)
+
+		now := time.Now()
+
+		sess := NewSession()
+		sess.SetRateClasses(now, rateClasses)
+
+		var wg sync.WaitGroup
+		wg.Add(1)
+
+		ctx, cancel := context.WithCancel(t.Context())
+		go func() {
+			defer wg.Done()
+			for {
+				select {
+				case <-ctx.Done():
+					return
+				case <-sess.WarningCh():
+				}
+			}
+		}()
+
+		assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
+
+		sess.ScaleWarningAndRateLimit(1000, 3)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].AlertLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].ClearLevel)
+		assert.Equal(t, int32(6000), sess.rateLimitStates[2].LimitLevel)
 
 		cancel()
 		wg.Wait()