session_test.go 27 KB

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