session_test.go 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904
  1. package state
  2. import (
  3. "context"
  4. "math"
  5. "net/netip"
  6. "sync"
  7. "testing"
  8. "time"
  9. "github.com/mk6i/retro-aim-server/wire"
  10. "github.com/stretchr/testify/assert"
  11. )
  12. func TestSession_SetAndGetAwayMessage(t *testing.T) {
  13. s := NewSession()
  14. assert.Empty(t, s.AwayMessage())
  15. msg := "here's my message"
  16. s.SetAwayMessage(msg)
  17. assert.Equal(t, msg, s.AwayMessage())
  18. }
  19. func TestSession_IncrementAndGetWarning(t *testing.T) {
  20. s := NewSession()
  21. var wg sync.WaitGroup
  22. wg.Add(1)
  23. go func() {
  24. defer wg.Done()
  25. s.ScaleWarningAndRateLimit(1, 1)
  26. s.ScaleWarningAndRateLimit(2, 1)
  27. s.ScaleWarningAndRateLimit(3, 1)
  28. }()
  29. assert.Equal(t, uint16(1), <-s.WarningCh())
  30. assert.Equal(t, uint16(3), <-s.WarningCh())
  31. assert.Equal(t, uint16(6), <-s.WarningCh())
  32. wg.Wait()
  33. }
  34. func TestSession_SetAndGetInvisible(t *testing.T) {
  35. s := NewSession()
  36. assert.False(t, s.Invisible())
  37. s.SetUserStatusBitmask(wire.OServiceUserStatusInvisible)
  38. assert.True(t, s.Invisible())
  39. }
  40. func TestSession_SetAndGetScreenName(t *testing.T) {
  41. s := NewSession()
  42. assert.Empty(t, s.IdentScreenName())
  43. sn := NewIdentScreenName("user-screen-name")
  44. s.SetIdentScreenName(sn)
  45. assert.Equal(t, sn, s.IdentScreenName())
  46. }
  47. func TestSession_SetAndGetChatRoomCookie(t *testing.T) {
  48. s := NewSession()
  49. assert.Empty(t, s.ChatRoomCookie())
  50. sn := "the-chat-cookie"
  51. s.SetChatRoomCookie(sn)
  52. assert.Equal(t, sn, s.ChatRoomCookie())
  53. }
  54. func TestSession_SetAndGetUIN(t *testing.T) {
  55. s := NewSession()
  56. assert.Empty(t, s.UIN())
  57. uin := uint32(100003)
  58. s.SetUIN(uin)
  59. assert.Equal(t, uin, s.UIN())
  60. }
  61. func TestSession_SetAndGetClientID(t *testing.T) {
  62. s := NewSession()
  63. assert.Empty(t, s.ClientID())
  64. clientID := "AIM Client ID"
  65. s.SetClientID(clientID)
  66. assert.Equal(t, clientID, s.ClientID())
  67. }
  68. func TestSession_SetAndGetRemoteAddr(t *testing.T) {
  69. s := NewSession()
  70. assert.Empty(t, s.RemoteAddr())
  71. remoteAddr, _ := netip.ParseAddrPort("1.2.3.4:1234")
  72. s.SetRemoteAddr(&remoteAddr)
  73. assert.Equal(t, &remoteAddr, s.RemoteAddr())
  74. }
  75. func TestSession_TLVUserInfo(t *testing.T) {
  76. tests := []struct {
  77. name string
  78. givenSessionFn func() *Session
  79. want wire.TLVUserInfo
  80. }{
  81. {
  82. name: "user is active and visible",
  83. givenSessionFn: func() *Session {
  84. s := NewSession()
  85. s.SetSignonTime(time.Unix(1, 0))
  86. s.SetIdentScreenName(NewIdentScreenName("xXAIMUSERXx"))
  87. s.SetDisplayScreenName("xXAIMUSERXx")
  88. s.ScaleWarningAndRateLimit(10, 1)
  89. s.SetUserInfoFlag(wire.OServiceUserFlagOSCARFree)
  90. return s
  91. },
  92. want: wire.TLVUserInfo{
  93. ScreenName: "xXAIMUSERXx",
  94. WarningLevel: 10,
  95. TLVBlock: wire.TLVBlock{
  96. TLVList: wire.TLVList{
  97. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  98. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x0010)),
  99. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  100. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  101. },
  102. },
  103. },
  104. },
  105. {
  106. name: "user is on ICQ",
  107. givenSessionFn: func() *Session {
  108. s := NewSession()
  109. s.SetSignonTime(time.Unix(1, 0))
  110. s.SetIdentScreenName(NewIdentScreenName("1000003"))
  111. s.SetDisplayScreenName("1000003")
  112. s.SetUserInfoFlag(wire.OServiceUserFlagICQ)
  113. return s
  114. },
  115. want: wire.TLVUserInfo{
  116. ScreenName: "1000003",
  117. TLVBlock: wire.TLVBlock{
  118. TLVList: wire.TLVList{
  119. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  120. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, wire.OServiceUserFlagOSCARFree|wire.OServiceUserFlagICQ),
  121. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  122. wire.NewTLVBE(wire.OServiceUserInfoICQDC, wire.ICQDCInfo{}),
  123. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  124. },
  125. },
  126. },
  127. },
  128. {
  129. name: "user has away message set",
  130. givenSessionFn: func() *Session {
  131. s := NewSession()
  132. s.SetSignonTime(time.Unix(1, 0))
  133. s.SetAwayMessage("here's my away message")
  134. return s
  135. },
  136. want: wire.TLVUserInfo{
  137. TLVBlock: wire.TLVBlock{
  138. TLVList: wire.TLVList{
  139. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  140. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x30)),
  141. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  142. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  143. },
  144. },
  145. },
  146. },
  147. {
  148. name: "user is invisible",
  149. givenSessionFn: func() *Session {
  150. s := NewSession()
  151. s.SetSignonTime(time.Unix(1, 0))
  152. s.SetUserStatusBitmask(wire.OServiceUserStatusInvisible)
  153. return s
  154. },
  155. want: wire.TLVUserInfo{
  156. TLVBlock: wire.TLVBlock{
  157. TLVList: wire.TLVList{
  158. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  159. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x0010)),
  160. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0100)),
  161. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  162. },
  163. },
  164. },
  165. },
  166. {
  167. name: "user is idle",
  168. givenSessionFn: func() *Session {
  169. s := NewSession()
  170. // sign on at t=0m
  171. timeBegin := time.Unix(0, 0)
  172. s.SetSignonTime(timeBegin)
  173. // set idle for 1m at t=+5m (ergo user idled @ t=+4m)
  174. timeIdle := timeBegin.Add(5 * time.Minute)
  175. s.nowFn = func() time.Time { return timeIdle }
  176. s.SetIdle(1 * time.Minute)
  177. // now it's t=+10m, ergo idle time should be t10-t4=6m
  178. timeNow := timeBegin.Add(10 * time.Minute)
  179. s.nowFn = func() time.Time { return timeNow }
  180. return s
  181. },
  182. want: wire.TLVUserInfo{
  183. TLVBlock: wire.TLVBlock{
  184. TLVList: wire.TLVList{
  185. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(0)),
  186. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x0010)),
  187. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  188. wire.NewTLVBE(wire.OServiceUserInfoIdleTime, uint16(6)),
  189. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  190. },
  191. },
  192. },
  193. },
  194. {
  195. name: "user goes idle then returns",
  196. givenSessionFn: func() *Session {
  197. s := NewSession()
  198. s.SetSignonTime(time.Unix(1, 0))
  199. s.SetIdle(1 * time.Second)
  200. s.UnsetIdle()
  201. return s
  202. },
  203. want: wire.TLVUserInfo{
  204. TLVBlock: wire.TLVBlock{
  205. TLVList: wire.TLVList{
  206. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  207. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x0010)),
  208. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  209. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  210. },
  211. },
  212. },
  213. },
  214. {
  215. name: "user has capabilities",
  216. givenSessionFn: func() *Session {
  217. s := NewSession()
  218. s.SetSignonTime(time.Unix(1, 0))
  219. s.SetCaps([][16]byte{
  220. {
  221. // chat: "748F2420-6287-11D1-8222-444553540000"
  222. 0x74, 0x8f, 0x24, 0x20, 0x62, 0x87, 0x11, 0xd1,
  223. 0x82, 0x22, 0x44, 0x45, 0x53, 0x54, 0x00, 0x00,
  224. },
  225. {
  226. // chat2: "748F2420-6287-11D1-8222-444553540000"
  227. 0x75, 0x8f, 0x24, 0x20, 0x62, 0x87, 0x11, 0xd1,
  228. 0x82, 0x22, 0x44, 0x45, 0x53, 0x54, 0x00, 0x01,
  229. },
  230. })
  231. return s
  232. },
  233. want: wire.TLVUserInfo{
  234. TLVBlock: wire.TLVBlock{
  235. TLVList: wire.TLVList{
  236. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  237. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x0010)),
  238. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  239. wire.NewTLVBE(wire.OServiceUserInfoOscarCaps, []byte{
  240. // chat: "748F2420-6287-11D1-8222-444553540000"
  241. 0x74, 0x8f, 0x24, 0x20, 0x62, 0x87, 0x11, 0xd1,
  242. 0x82, 0x22, 0x44, 0x45, 0x53, 0x54, 0x00, 0x00,
  243. // chat: "748F2420-6287-11D1-8222-444553540000"
  244. 0x75, 0x8f, 0x24, 0x20, 0x62, 0x87, 0x11, 0xd1,
  245. 0x82, 0x22, 0x44, 0x45, 0x53, 0x54, 0x00, 0x01,
  246. }),
  247. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  248. },
  249. },
  250. },
  251. },
  252. {
  253. name: "user has buddy icon",
  254. givenSessionFn: func() *Session {
  255. s := NewSession()
  256. s.SetSignonTime(time.Unix(1, 0))
  257. return s
  258. },
  259. want: wire.TLVUserInfo{
  260. WarningLevel: 0,
  261. TLVBlock: wire.TLVBlock{
  262. TLVList: wire.TLVList{
  263. wire.NewTLVBE(wire.OServiceUserInfoSignonTOD, uint32(1)),
  264. wire.NewTLVBE(wire.OServiceUserInfoUserFlags, uint16(0x0010)),
  265. wire.NewTLVBE(wire.OServiceUserInfoStatus, uint32(0x0000)),
  266. wire.NewTLVBE(wire.OServiceUserInfoMySubscriptions, uint32(0)),
  267. },
  268. },
  269. },
  270. },
  271. }
  272. for _, tt := range tests {
  273. t.Run(tt.name, func(t *testing.T) {
  274. s := tt.givenSessionFn()
  275. assert.Equal(t, tt.want, s.TLVUserInfo())
  276. })
  277. }
  278. }
  279. func TestSession_SendAndRecvMessage_ExpectSessSendOK(t *testing.T) {
  280. s := NewSession()
  281. msg := wire.SNACMessage{
  282. Frame: wire.SNACFrame{
  283. FoodGroup: wire.ICBM,
  284. },
  285. }
  286. var wg sync.WaitGroup
  287. wg.Add(1)
  288. go func() {
  289. defer wg.Done()
  290. defer s.Close()
  291. status := s.RelayMessage(msg)
  292. assert.Equal(t, SessSendOK, status)
  293. }()
  294. loop:
  295. for {
  296. select {
  297. case m := <-s.ReceiveMessage():
  298. assert.Equal(t, msg, m)
  299. case <-s.Closed():
  300. break loop
  301. }
  302. }
  303. wg.Wait()
  304. }
  305. func TestSession_SendMessage_SessSendClosed(t *testing.T) {
  306. s := Session{
  307. msgCh: make(chan wire.SNACMessage, 1),
  308. stopCh: make(chan struct{}),
  309. }
  310. s.Close()
  311. if res := s.RelayMessage(wire.SNACMessage{}); res != SessSendClosed {
  312. t.Fatalf("expected SessSendClosed, got %+v", res)
  313. }
  314. }
  315. func TestSession_SendMessage_SessQueueFull(t *testing.T) {
  316. bufSize := 10
  317. s := Session{
  318. msgCh: make(chan wire.SNACMessage, bufSize),
  319. stopCh: make(chan struct{}),
  320. }
  321. for i := 0; i < bufSize; i++ {
  322. assert.Equal(t, SessSendOK, s.RelayMessage(wire.SNACMessage{}))
  323. }
  324. assert.Equal(t, SessQueueFull, s.RelayMessage(wire.SNACMessage{}))
  325. }
  326. func TestSession_Close_Twice(t *testing.T) {
  327. s := Session{
  328. stopCh: make(chan struct{}),
  329. }
  330. s.Close()
  331. s.Close() // make sure close is idempotent
  332. if !s.closed {
  333. t.Fatal("expected session to be closed")
  334. }
  335. select {
  336. case <-s.Closed():
  337. case <-time.After(1 * time.Second):
  338. t.Fatalf("channel is not closed")
  339. }
  340. }
  341. func TestSession_Close(t *testing.T) {
  342. s := NewSession()
  343. select {
  344. case <-s.Closed():
  345. assert.Fail(t, "channel is closed")
  346. default:
  347. // channel is open by default
  348. }
  349. s.Close()
  350. <-s.Closed()
  351. }
  352. func TestSession_EvaluateRateLimit_ObserveRateChanges(t *testing.T) {
  353. classParams := [5]wire.RateClass{
  354. {
  355. ID: 1,
  356. WindowSize: 80,
  357. ClearLevel: 2500,
  358. AlertLevel: 2000,
  359. LimitLevel: 1500,
  360. DisconnectLevel: 800,
  361. MaxLevel: 6000,
  362. },
  363. {
  364. ID: 2,
  365. WindowSize: 80,
  366. ClearLevel: 3000,
  367. AlertLevel: 2000,
  368. LimitLevel: 1500,
  369. DisconnectLevel: 1000,
  370. MaxLevel: 6000,
  371. },
  372. {
  373. ID: 3,
  374. WindowSize: 20,
  375. ClearLevel: 5100,
  376. AlertLevel: 5000,
  377. LimitLevel: 4000,
  378. DisconnectLevel: 3000,
  379. MaxLevel: 6000,
  380. },
  381. {
  382. ID: 4,
  383. WindowSize: 20,
  384. ClearLevel: 5500,
  385. AlertLevel: 5300,
  386. LimitLevel: 4200,
  387. DisconnectLevel: 3000,
  388. MaxLevel: 8000,
  389. },
  390. {
  391. ID: 5,
  392. WindowSize: 10,
  393. ClearLevel: 5500,
  394. AlertLevel: 5300,
  395. LimitLevel: 4200,
  396. DisconnectLevel: 3000,
  397. MaxLevel: 8000,
  398. },
  399. }
  400. rateClasses := wire.NewRateLimitClasses(classParams)
  401. t.Run("we can action every 5 seconds indefinitely without getting rate limited", func(t *testing.T) {
  402. now := time.Now()
  403. sess := NewSession()
  404. sess.SetRateClasses(now, rateClasses)
  405. rateClass := rateClasses.Get(3)
  406. sess.SubscribeRateLimits([]wire.RateLimitClassID{rateClass.ID})
  407. for i := 0; i < 100; i++ {
  408. now = now.Add(5 * time.Second)
  409. have := sess.EvaluateRateLimit(now, rateClass.ID)
  410. assert.Equal(t, wire.RateLimitStatusClear, have)
  411. }
  412. })
  413. t.Run("reach disconnect threshold", func(t *testing.T) {
  414. now := time.Now()
  415. sess := NewSession()
  416. sess.SetRateClasses(now, rateClasses)
  417. rateClass := rateClasses.Get(3)
  418. sess.SubscribeRateLimits([]wire.RateLimitClassID{rateClass.ID})
  419. // record some event in the rate limiter
  420. want := []wire.RateLimitStatus{
  421. wire.RateLimitStatusClear,
  422. wire.RateLimitStatusClear,
  423. wire.RateLimitStatusClear,
  424. wire.RateLimitStatusClear,
  425. wire.RateLimitStatusAlert,
  426. wire.RateLimitStatusAlert,
  427. wire.RateLimitStatusAlert,
  428. wire.RateLimitStatusAlert,
  429. wire.RateLimitStatusAlert,
  430. wire.RateLimitStatusLimited,
  431. wire.RateLimitStatusLimited,
  432. wire.RateLimitStatusLimited,
  433. wire.RateLimitStatusLimited,
  434. wire.RateLimitStatusLimited,
  435. wire.RateLimitStatusLimited,
  436. wire.RateLimitStatusLimited,
  437. wire.RateLimitStatusLimited,
  438. wire.RateLimitStatusDisconnect,
  439. }
  440. for i := 0; i < len(want); i++ {
  441. now = now.Add(1 * time.Second)
  442. have := sess.EvaluateRateLimit(now, rateClass.ID)
  443. assert.Equal(t, want[i], have)
  444. }
  445. select {
  446. case <-sess.Closed():
  447. default:
  448. t.Error("expected session to be closed")
  449. }
  450. })
  451. t.Run("reach rate limit threshold, wait for clear threshold", func(t *testing.T) {
  452. now := time.Now()
  453. sess := NewSession()
  454. sess.SetRateClasses(now, rateClasses)
  455. rateClass := rateClasses.Get(3)
  456. sess.SubscribeRateLimits([]wire.RateLimitClassID{rateClass.ID})
  457. // first reach the rate limit threshold
  458. want := []wire.RateLimitStatus{
  459. wire.RateLimitStatusClear,
  460. wire.RateLimitStatusClear,
  461. wire.RateLimitStatusClear,
  462. wire.RateLimitStatusClear,
  463. wire.RateLimitStatusAlert,
  464. wire.RateLimitStatusAlert,
  465. wire.RateLimitStatusAlert,
  466. wire.RateLimitStatusAlert,
  467. wire.RateLimitStatusAlert,
  468. wire.RateLimitStatusLimited,
  469. }
  470. for i := 0; i < len(want); i++ {
  471. now = now.Add(1 * time.Second)
  472. have := sess.EvaluateRateLimit(now, rateClass.ID)
  473. assert.Equal(t, want[i], have)
  474. if i > 0 && want[i-1] != want[i] {
  475. classChanges, rateChanges := sess.ObserveRateChanges(now)
  476. assert.Empty(t, classChanges)
  477. if assert.NotEmpty(t, rateChanges) {
  478. rateDelta := rateChanges[0]
  479. assert.Equal(t, rateClass, rateDelta.RateClass)
  480. assert.Equal(t, want[i], rateDelta.CurrentStatus)
  481. assert.True(t, rateDelta.Subscribed)
  482. if want[i] == wire.RateLimitStatusLimited {
  483. assert.True(t, rateDelta.LimitedNow)
  484. }
  485. }
  486. }
  487. }
  488. // this is a rearranged moving average formula that determines how many
  489. // milliseconds it will take to reach the clear threshold
  490. timeToRecover := int(math.Ceil((time.Duration(rateClass.ClearLevel*rateClass.WindowSize-sess.rateLimitStates[rateClass.ID-1].CurrentLevel*(rateClass.WindowSize-1)) * time.Millisecond).Seconds()))
  491. assert.True(t, timeToRecover > 0)
  492. // indicate the time rate limiting kicked in
  493. timeLimited := now
  494. for i := 0; i < timeToRecover; i++ {
  495. now = now.Add(1 * time.Second)
  496. classDelta, stateDelta := sess.ObserveRateChanges(now)
  497. assert.Empty(t, classDelta)
  498. if i == timeToRecover-1 {
  499. // assert that the clear threshold has been met.
  500. assert.ElementsMatch(t, stateDelta, []RateClassState{
  501. {
  502. RateClass: rateClass,
  503. CurrentLevel: 5140,
  504. CurrentStatus: wire.RateLimitStatusClear,
  505. LastTime: timeLimited,
  506. Subscribed: true,
  507. LimitedNow: false,
  508. }})
  509. } else {
  510. // assert that no changed have been observed, it's still rate-limited
  511. assert.Nil(t, stateDelta)
  512. }
  513. }
  514. })
  515. t.Run("observe a rate class change", func(t *testing.T) {
  516. now := time.Now()
  517. sess := NewSession()
  518. sess.SetRateClasses(now, rateClasses)
  519. rateClass := rateClasses.Get(3)
  520. sess.SubscribeRateLimits([]wire.RateLimitClassID{rateClass.ID})
  521. now = now.Add(1 * time.Second)
  522. classDelta, stateDelta := sess.ObserveRateChanges(now)
  523. assert.Empty(t, classDelta)
  524. assert.Empty(t, stateDelta)
  525. paramsCopy := classParams
  526. paramsCopy[rateClass.ID-1].LimitLevel++
  527. newRateClasses := wire.NewRateLimitClasses(paramsCopy)
  528. now = now.Add(1 * time.Second)
  529. sess.SetRateClasses(now, newRateClasses)
  530. now = now.Add(1 * time.Second)
  531. classDelta, stateDelta = sess.ObserveRateChanges(now)
  532. assert.Equal(t, classDelta[0].RateClass, newRateClasses.Get(rateClass.ID))
  533. assert.Empty(t, stateDelta)
  534. })
  535. t.Run("as a bot, I can action every second indefinitely without getting rate limited", func(t *testing.T) {
  536. now := time.Now()
  537. sess := NewSession()
  538. sess.SetUserInfoFlag(wire.OServiceUserFlagBot)
  539. sess.SetRateClasses(now, rateClasses)
  540. for i := 0; i < 100; i++ {
  541. now = now.Add(1 * time.Second)
  542. have := sess.EvaluateRateLimit(now, wire.RateLimitClassID(1))
  543. assert.Equal(t, wire.RateLimitStatusClear, have)
  544. }
  545. })
  546. }
  547. func TestSession_SetAndGetFoodGroupVersions(t *testing.T) {
  548. versions := [wire.MDir + 1]uint16{}
  549. versions[wire.Feedbag] = 1
  550. versions[wire.OService] = 2
  551. s := NewSession()
  552. s.SetFoodGroupVersions(versions)
  553. assert.Equal(t, versions, s.FoodGroupVersions())
  554. }
  555. func TestSession_SetAndGetTypingEventsEnabled(t *testing.T) {
  556. s := NewSession()
  557. assert.False(t, s.TypingEventsEnabled())
  558. s.SetTypingEventsEnabled(true)
  559. assert.True(t, s.TypingEventsEnabled())
  560. s.SetTypingEventsEnabled(false)
  561. assert.False(t, s.TypingEventsEnabled())
  562. }
  563. func TestSession_SetAndGetMultiConnFlag(t *testing.T) {
  564. s := NewSession()
  565. assert.Zero(t, s.MultiConnFlag())
  566. s.SetMultiConnFlag(wire.MultiConnFlagsOldClient)
  567. assert.Equal(t, wire.MultiConnFlagsOldClient, s.MultiConnFlag())
  568. s.SetMultiConnFlag(wire.MultiConnFlagsRecentClient)
  569. assert.Equal(t, wire.MultiConnFlagsRecentClient, s.MultiConnFlag())
  570. s.SetMultiConnFlag(wire.MultiConnFlagsSingleClient)
  571. assert.Equal(t, wire.MultiConnFlagsSingleClient, s.MultiConnFlag())
  572. }
  573. func TestSession_SetAndGetLastWarnLevel(t *testing.T) {
  574. s := NewSession()
  575. assert.Zero(t, s.Warning())
  576. level := uint16(500)
  577. s.SetWarning(level)
  578. assert.Equal(t, level, s.Warning())
  579. }
  580. func TestSession_ScaleWarningAndRateLimit(t *testing.T) {
  581. t.Run("scale up", func(t *testing.T) {
  582. classParams := [5]wire.RateClass{
  583. {},
  584. {},
  585. {
  586. ID: 3,
  587. WindowSize: 20,
  588. ClearLevel: 5100,
  589. AlertLevel: 5000,
  590. LimitLevel: 4000,
  591. DisconnectLevel: 3000,
  592. MaxLevel: 6000,
  593. },
  594. {},
  595. {},
  596. }
  597. rateClasses := wire.NewRateLimitClasses(classParams)
  598. now := time.Now()
  599. sess := NewSession()
  600. sess.SetRateClasses(now, rateClasses)
  601. var wg sync.WaitGroup
  602. wg.Add(1)
  603. ctx, cancel := context.WithCancel(t.Context())
  604. go func() {
  605. defer wg.Done()
  606. for {
  607. select {
  608. case <-ctx.Done():
  609. return
  610. case <-sess.WarningCh():
  611. }
  612. }
  613. }()
  614. assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
  615. assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
  616. assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
  617. sess.ScaleWarningAndRateLimit(100, 3)
  618. assert.Equal(t, int32(5085), sess.rateLimitStates[2].AlertLevel)
  619. assert.Equal(t, int32(5175), sess.rateLimitStates[2].ClearLevel)
  620. assert.Equal(t, int32(4185), sess.rateLimitStates[2].LimitLevel)
  621. sess.ScaleWarningAndRateLimit(100, 3)
  622. assert.Equal(t, int32(5170), sess.rateLimitStates[2].AlertLevel)
  623. assert.Equal(t, int32(5250), sess.rateLimitStates[2].ClearLevel)
  624. assert.Equal(t, int32(4370), sess.rateLimitStates[2].LimitLevel)
  625. sess.ScaleWarningAndRateLimit(100, 3)
  626. assert.Equal(t, int32(5255), sess.rateLimitStates[2].AlertLevel)
  627. assert.Equal(t, int32(5325), sess.rateLimitStates[2].ClearLevel)
  628. assert.Equal(t, int32(4555), sess.rateLimitStates[2].LimitLevel)
  629. sess.ScaleWarningAndRateLimit(100, 3)
  630. assert.Equal(t, int32(5340), sess.rateLimitStates[2].AlertLevel)
  631. assert.Equal(t, int32(5400), sess.rateLimitStates[2].ClearLevel)
  632. assert.Equal(t, int32(4740), sess.rateLimitStates[2].LimitLevel)
  633. sess.ScaleWarningAndRateLimit(100, 3)
  634. assert.Equal(t, int32(5425), sess.rateLimitStates[2].AlertLevel)
  635. assert.Equal(t, int32(5475), sess.rateLimitStates[2].ClearLevel)
  636. assert.Equal(t, int32(4925), sess.rateLimitStates[2].LimitLevel)
  637. sess.ScaleWarningAndRateLimit(100, 3)
  638. assert.Equal(t, int32(5510), sess.rateLimitStates[2].AlertLevel)
  639. assert.Equal(t, int32(5550), sess.rateLimitStates[2].ClearLevel)
  640. assert.Equal(t, int32(5110), sess.rateLimitStates[2].LimitLevel)
  641. sess.ScaleWarningAndRateLimit(100, 3)
  642. assert.Equal(t, int32(5595), sess.rateLimitStates[2].AlertLevel)
  643. assert.Equal(t, int32(5625), sess.rateLimitStates[2].ClearLevel)
  644. assert.Equal(t, int32(5295), sess.rateLimitStates[2].LimitLevel)
  645. sess.ScaleWarningAndRateLimit(100, 3)
  646. assert.Equal(t, int32(5680), sess.rateLimitStates[2].AlertLevel)
  647. assert.Equal(t, int32(5700), sess.rateLimitStates[2].ClearLevel)
  648. assert.Equal(t, int32(5480), sess.rateLimitStates[2].LimitLevel)
  649. sess.ScaleWarningAndRateLimit(100, 3)
  650. assert.Equal(t, int32(5765), sess.rateLimitStates[2].AlertLevel)
  651. assert.Equal(t, int32(5775), sess.rateLimitStates[2].ClearLevel)
  652. assert.Equal(t, int32(5665), sess.rateLimitStates[2].LimitLevel)
  653. sess.ScaleWarningAndRateLimit(100, 3)
  654. assert.Equal(t, int32(5850), sess.rateLimitStates[2].AlertLevel)
  655. assert.Equal(t, int32(5850), sess.rateLimitStates[2].ClearLevel)
  656. assert.Equal(t, int32(5850), sess.rateLimitStates[2].LimitLevel)
  657. sess.ScaleWarningAndRateLimit(100, 3)
  658. assert.Equal(t, int32(5850), sess.rateLimitStates[2].AlertLevel)
  659. assert.Equal(t, int32(5850), sess.rateLimitStates[2].ClearLevel)
  660. assert.Equal(t, int32(5850), sess.rateLimitStates[2].LimitLevel)
  661. cancel()
  662. wg.Wait()
  663. })
  664. t.Run("scale down", func(t *testing.T) {
  665. currentClassParams := [5]wire.RateClass{
  666. {},
  667. {},
  668. {
  669. ID: 3,
  670. WindowSize: 20,
  671. ClearLevel: 5100,
  672. AlertLevel: 5000,
  673. LimitLevel: 4000,
  674. DisconnectLevel: 3000,
  675. MaxLevel: 6000,
  676. },
  677. {},
  678. {},
  679. }
  680. rateClasses := wire.NewRateLimitClasses(currentClassParams)
  681. now := time.Now()
  682. sess := NewSession()
  683. sess.SetRateClasses(now, rateClasses)
  684. var wg sync.WaitGroup
  685. wg.Add(1)
  686. ctx, cancel := context.WithCancel(t.Context())
  687. go func() {
  688. defer wg.Done()
  689. for {
  690. select {
  691. case <-ctx.Done():
  692. return
  693. case <-sess.WarningCh():
  694. }
  695. }
  696. }()
  697. for i := 0; i < 10; i++ {
  698. sess.ScaleWarningAndRateLimit(100, 3)
  699. }
  700. assert.Equal(t, int32(5850), sess.rateLimitStates[2].AlertLevel)
  701. assert.Equal(t, int32(5850), sess.rateLimitStates[2].ClearLevel)
  702. assert.Equal(t, int32(5850), sess.rateLimitStates[2].LimitLevel)
  703. sess.ScaleWarningAndRateLimit(-100, 3)
  704. assert.Equal(t, int32(5765), sess.rateLimitStates[2].AlertLevel)
  705. assert.Equal(t, int32(5775), sess.rateLimitStates[2].ClearLevel)
  706. assert.Equal(t, int32(5665), sess.rateLimitStates[2].LimitLevel)
  707. sess.ScaleWarningAndRateLimit(-100, 3)
  708. assert.Equal(t, int32(5680), sess.rateLimitStates[2].AlertLevel)
  709. assert.Equal(t, int32(5700), sess.rateLimitStates[2].ClearLevel)
  710. assert.Equal(t, int32(5480), sess.rateLimitStates[2].LimitLevel)
  711. sess.ScaleWarningAndRateLimit(-100, 3)
  712. assert.Equal(t, int32(5595), sess.rateLimitStates[2].AlertLevel)
  713. assert.Equal(t, int32(5625), sess.rateLimitStates[2].ClearLevel)
  714. assert.Equal(t, int32(5295), sess.rateLimitStates[2].LimitLevel)
  715. sess.ScaleWarningAndRateLimit(-100, 3)
  716. assert.Equal(t, int32(5510), sess.rateLimitStates[2].AlertLevel)
  717. assert.Equal(t, int32(5550), sess.rateLimitStates[2].ClearLevel)
  718. assert.Equal(t, int32(5110), sess.rateLimitStates[2].LimitLevel)
  719. sess.ScaleWarningAndRateLimit(-100, 3)
  720. assert.Equal(t, int32(5425), sess.rateLimitStates[2].AlertLevel)
  721. assert.Equal(t, int32(5475), sess.rateLimitStates[2].ClearLevel)
  722. assert.Equal(t, int32(4925), sess.rateLimitStates[2].LimitLevel)
  723. sess.ScaleWarningAndRateLimit(-100, 3)
  724. assert.Equal(t, int32(5340), sess.rateLimitStates[2].AlertLevel)
  725. assert.Equal(t, int32(5400), sess.rateLimitStates[2].ClearLevel)
  726. assert.Equal(t, int32(4740), sess.rateLimitStates[2].LimitLevel)
  727. sess.ScaleWarningAndRateLimit(-100, 3)
  728. assert.Equal(t, int32(5255), sess.rateLimitStates[2].AlertLevel)
  729. assert.Equal(t, int32(5325), sess.rateLimitStates[2].ClearLevel)
  730. assert.Equal(t, int32(4555), sess.rateLimitStates[2].LimitLevel)
  731. sess.ScaleWarningAndRateLimit(-100, 3)
  732. assert.Equal(t, int32(5170), sess.rateLimitStates[2].AlertLevel)
  733. assert.Equal(t, int32(5250), sess.rateLimitStates[2].ClearLevel)
  734. assert.Equal(t, int32(4370), sess.rateLimitStates[2].LimitLevel)
  735. sess.ScaleWarningAndRateLimit(-100, 3)
  736. assert.Equal(t, int32(5085), sess.rateLimitStates[2].AlertLevel)
  737. assert.Equal(t, int32(5175), sess.rateLimitStates[2].ClearLevel)
  738. assert.Equal(t, int32(4185), sess.rateLimitStates[2].LimitLevel)
  739. sess.ScaleWarningAndRateLimit(-100, 3)
  740. assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
  741. assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
  742. assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
  743. sess.ScaleWarningAndRateLimit(-100, 3)
  744. assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
  745. assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
  746. assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
  747. cancel()
  748. wg.Wait()
  749. })
  750. t.Run("increment 100%", func(t *testing.T) {
  751. classParams := [5]wire.RateClass{
  752. {},
  753. {},
  754. {
  755. ID: 3,
  756. WindowSize: 20,
  757. ClearLevel: 5100,
  758. AlertLevel: 5000,
  759. LimitLevel: 4000,
  760. DisconnectLevel: 3000,
  761. MaxLevel: 6000,
  762. },
  763. {},
  764. {},
  765. }
  766. rateClasses := wire.NewRateLimitClasses(classParams)
  767. now := time.Now()
  768. sess := NewSession()
  769. sess.SetRateClasses(now, rateClasses)
  770. var wg sync.WaitGroup
  771. wg.Add(1)
  772. ctx, cancel := context.WithCancel(t.Context())
  773. go func() {
  774. defer wg.Done()
  775. for {
  776. select {
  777. case <-ctx.Done():
  778. return
  779. case <-sess.WarningCh():
  780. }
  781. }
  782. }()
  783. assert.Equal(t, int32(5000), sess.rateLimitStates[2].AlertLevel)
  784. assert.Equal(t, int32(5100), sess.rateLimitStates[2].ClearLevel)
  785. assert.Equal(t, int32(4000), sess.rateLimitStates[2].LimitLevel)
  786. sess.ScaleWarningAndRateLimit(1000, 3)
  787. assert.Equal(t, int32(5850), sess.rateLimitStates[2].AlertLevel)
  788. assert.Equal(t, int32(5850), sess.rateLimitStates[2].ClearLevel)
  789. assert.Equal(t, int32(5850), sess.rateLimitStates[2].LimitLevel)
  790. cancel()
  791. wg.Wait()
  792. })
  793. }