rate_limit_test.go 7.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293
  1. package wire
  2. import (
  3. "testing"
  4. "time"
  5. "github.com/stretchr/testify/assert"
  6. "github.com/stretchr/testify/require"
  7. )
  8. func TestNewRateLimitClasses(t *testing.T) {
  9. input := [5]RateClass{
  10. {ID: 1, WindowSize: 10},
  11. {ID: 2, WindowSize: 20},
  12. {ID: 3, WindowSize: 30},
  13. {ID: 4, WindowSize: 40},
  14. {ID: 5, WindowSize: 50},
  15. }
  16. classes := NewRateLimitClasses(input)
  17. assert.Equal(t, input, classes.All())
  18. }
  19. func TestCheckRateLimit(t *testing.T) {
  20. baseTime := time.Date(2025, 1, 1, 12, 0, 0, 0, time.UTC)
  21. testCases := []struct {
  22. name string
  23. lastTime time.Time
  24. currentTime time.Time
  25. rateClass RateClass
  26. currentAvg int32
  27. limitedNow bool
  28. wantStatus RateLimitStatus
  29. wantNewAvg int32
  30. }{
  31. {
  32. name: "Already limited, but newAvg >= ClearLevel => Clear",
  33. lastTime: baseTime,
  34. currentTime: baseTime.Add(50 * time.Millisecond),
  35. rateClass: RateClass{
  36. WindowSize: 4,
  37. MaxLevel: 1000,
  38. ClearLevel: 10,
  39. DisconnectLevel: 2,
  40. LimitLevel: 5,
  41. AlertLevel: 8,
  42. },
  43. currentAvg: 9, // currentAvg close to ClearLevel
  44. limitedNow: true, // already limited
  45. // newAvg = (9*(4-1) + 50) / 4 = (27 + 50) / 4 = 77 / 4 = 19
  46. // 19 >= ClearLevel(10) => RateLimitStatusClear
  47. wantStatus: RateLimitStatusClear,
  48. wantNewAvg: 19,
  49. },
  50. {
  51. name: "Already limited, but newAvg < ClearLevel => Remain Limited",
  52. lastTime: baseTime,
  53. currentTime: baseTime.Add(30 * time.Millisecond),
  54. rateClass: RateClass{
  55. WindowSize: 4,
  56. MaxLevel: 1000,
  57. ClearLevel: 50,
  58. DisconnectLevel: 10,
  59. LimitLevel: 20,
  60. AlertLevel: 30,
  61. },
  62. currentAvg: 10,
  63. limitedNow: true,
  64. // newAvg = (10*(4-1) + 30) / 4 = (30 + 30) / 4 = 60 / 4 = 15
  65. // 15 < ClearLevel(50) => remain limited
  66. wantStatus: RateLimitStatusLimited,
  67. wantNewAvg: 15,
  68. },
  69. {
  70. name: "Not Limited Now, New average < DisconnectLevel => Disconnect",
  71. lastTime: baseTime,
  72. currentTime: baseTime.Add(10 * time.Millisecond),
  73. rateClass: RateClass{
  74. WindowSize: 4,
  75. MaxLevel: 1000,
  76. ClearLevel: 100,
  77. DisconnectLevel: 5,
  78. LimitLevel: 20,
  79. AlertLevel: 40,
  80. },
  81. currentAvg: 1,
  82. limitedNow: false,
  83. // newAvg = (1*(4-1) + 10) / 4 = (3 + 10) / 4 = 13 / 4 = 3
  84. // 3 < DisconnectLevel(5) => RateLimitStatusDisconnect
  85. wantStatus: RateLimitStatusDisconnect,
  86. wantNewAvg: 3,
  87. },
  88. {
  89. name: "Limited Now, New average < DisconnectLevel => Disconnect",
  90. lastTime: baseTime,
  91. currentTime: baseTime.Add(10 * time.Millisecond),
  92. rateClass: RateClass{
  93. WindowSize: 4,
  94. MaxLevel: 1000,
  95. ClearLevel: 100,
  96. DisconnectLevel: 5,
  97. LimitLevel: 20,
  98. AlertLevel: 40,
  99. },
  100. currentAvg: 1,
  101. limitedNow: true,
  102. // newAvg = (1*(4-1) + 10) / 4 = (3 + 10) / 4 = 13 / 4 = 3
  103. // 3 < DisconnectLevel(5) => RateLimitStatusDisconnect
  104. wantStatus: RateLimitStatusDisconnect,
  105. wantNewAvg: 3,
  106. },
  107. {
  108. name: "New average < LimitLevel => Limited",
  109. lastTime: baseTime,
  110. currentTime: baseTime.Add(20 * time.Millisecond),
  111. rateClass: RateClass{
  112. WindowSize: 4,
  113. MaxLevel: 1000,
  114. ClearLevel: 100,
  115. DisconnectLevel: 5,
  116. LimitLevel: 40,
  117. AlertLevel: 60,
  118. },
  119. currentAvg: 10,
  120. limitedNow: false,
  121. // newAvg = (10*(4-1) + 20) / 4 = (30 + 20) / 4 = 50 / 4 = 12
  122. // 12 < LimitLevel(40) => RateLimitStatusLimited
  123. wantStatus: RateLimitStatusLimited,
  124. wantNewAvg: 12,
  125. },
  126. {
  127. name: "New average < AlertLevel => Alert",
  128. lastTime: baseTime,
  129. currentTime: baseTime.Add(30 * time.Millisecond),
  130. rateClass: RateClass{
  131. WindowSize: 4,
  132. MaxLevel: 1000,
  133. ClearLevel: 100,
  134. DisconnectLevel: 5,
  135. LimitLevel: 20,
  136. AlertLevel: 40,
  137. },
  138. currentAvg: 20,
  139. limitedNow: false,
  140. // newAvg = (20*(4-1) + 30) / 4 = (60 + 30) / 4 = 90 / 4 = 22
  141. // 22 >= 20 => not "Limited"; 22 < 40 => "Alert"
  142. wantStatus: RateLimitStatusAlert,
  143. wantNewAvg: 22,
  144. },
  145. {
  146. name: "New average >= AlertLevel => Clear",
  147. lastTime: baseTime,
  148. currentTime: baseTime.Add(50 * time.Millisecond),
  149. rateClass: RateClass{
  150. WindowSize: 4,
  151. MaxLevel: 1000,
  152. ClearLevel: 100,
  153. DisconnectLevel: 5,
  154. LimitLevel: 20,
  155. AlertLevel: 40,
  156. },
  157. // Choose 39 so the resulting newAvg is 41, which is >= AlertLevel.
  158. currentAvg: 39,
  159. limitedNow: false,
  160. // newAvg = (39*(4-1) + 50) / 4 = (117 + 50) / 4 = 167 / 4 = 41
  161. // 41 >= AlertLevel(40) => RateLimitStatusClear
  162. wantStatus: RateLimitStatusClear,
  163. wantNewAvg: 41,
  164. },
  165. {
  166. name: "Clamp newAvg to MaxLevel if exceeded",
  167. lastTime: baseTime,
  168. currentTime: baseTime.Add(9999 * time.Millisecond),
  169. rateClass: RateClass{
  170. WindowSize: 4,
  171. MaxLevel: 100,
  172. ClearLevel: 80,
  173. DisconnectLevel: 20,
  174. LimitLevel: 40,
  175. AlertLevel: 60,
  176. },
  177. currentAvg: 95,
  178. limitedNow: false,
  179. // Without clamping, newAvg would be huge:
  180. // newAvg = (95*(4-1) + 9999) / 4 = (285 + 9999)/4 = 10284/4 = 2571
  181. // Clamped to 100 => 100 >= AlertLevel(60) => RateLimitStatusClear
  182. wantStatus: RateLimitStatusClear,
  183. wantNewAvg: 100,
  184. },
  185. }
  186. for _, tc := range testCases {
  187. t.Run(tc.name, func(t *testing.T) {
  188. gotStatus, gotNewAvg := CheckRateLimit(
  189. tc.lastTime,
  190. tc.currentTime,
  191. tc.rateClass,
  192. tc.currentAvg,
  193. tc.limitedNow,
  194. )
  195. assert.Equal(t, tc.wantStatus, gotStatus)
  196. assert.Equal(t, tc.wantNewAvg, gotNewAvg)
  197. })
  198. }
  199. }
  200. func TestRateLimitClasses_Get(t *testing.T) {
  201. classes := DefaultRateLimitClasses()
  202. // Test Get() returns correct class for each ID
  203. for i := 1; i <= 5; i++ {
  204. id := RateLimitClassID(i)
  205. class := classes.Get(id)
  206. assert.Equal(t, id, class.ID)
  207. assert.Equal(t, classes.All()[i-1], class)
  208. }
  209. }
  210. func TestRateLimitClasses_All(t *testing.T) {
  211. classes := DefaultRateLimitClasses()
  212. // Test All() returns exactly 5 classes with correct IDs
  213. all := classes.All()
  214. assert.Len(t, all, 5)
  215. for i, class := range all {
  216. expectedID := RateLimitClassID(i + 1)
  217. assert.Equal(t, expectedID, class.ID, "class ID mismatch at index %d", i)
  218. }
  219. }
  220. func TestSNACRateLimits_RateClassLookup(t *testing.T) {
  221. limits := DefaultSNACRateLimits()
  222. testCases := []struct {
  223. foodGroup uint16
  224. subGroup uint16
  225. expected RateLimitClassID
  226. found bool
  227. }{
  228. {Chat, ChatUsersJoined, 1, true},
  229. {Chat, ChatChannelMsgToHost, 2, true},
  230. {0xFFFF, 0x0001, 0, false},
  231. {Chat, 0xFFFF, 0, false},
  232. }
  233. for _, tc := range testCases {
  234. classID, ok := limits.RateClassLookup(tc.foodGroup, tc.subGroup)
  235. assert.Equal(t, tc.found, ok)
  236. assert.Equal(t, tc.expected, classID)
  237. }
  238. }
  239. func TestSNACRateLimits_All(t *testing.T) {
  240. limits := DefaultSNACRateLimits()
  241. seen := map[uint16]map[uint16]RateLimitClassID{}
  242. for entry := range limits.All() {
  243. if _, ok := seen[entry.FoodGroup]; !ok {
  244. seen[entry.FoodGroup] = map[uint16]RateLimitClassID{}
  245. }
  246. seen[entry.FoodGroup][entry.SubGroup] = entry.RateLimitClass
  247. }
  248. // Spot-check a few values
  249. require.Contains(t, seen, ICBM)
  250. assert.Equal(t, RateLimitClassID(3), seen[ICBM][ICBMChannelMsgToHost])
  251. assert.Equal(t, RateLimitClassID(1), seen[ICBM][ICBMChannelMsgToClient])
  252. require.Contains(t, seen, Locate)
  253. assert.Equal(t, RateLimitClassID(4), seen[Locate][LocateSetDirInfo])
  254. assert.Equal(t, RateLimitClassID(3), seen[Locate][LocateUserInfoQuery])
  255. }
  256. func TestSNACRateLimits_All_YieldStopsEarly(t *testing.T) {
  257. limits := DefaultSNACRateLimits()
  258. count := 0
  259. limits.All()(func(entry struct {
  260. FoodGroup uint16
  261. SubGroup uint16
  262. RateLimitClass RateLimitClassID
  263. }) bool {
  264. count++
  265. // stop iteration after first item to trigger `if !yield(...) { return }`
  266. return false
  267. })
  268. // Should only yield one entry
  269. assert.Equal(t, 1, count)
  270. }