server.go 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674
  1. package oscar
  2. import (
  3. "bytes"
  4. "context"
  5. "errors"
  6. "fmt"
  7. "io"
  8. "log/slog"
  9. "net"
  10. "net/netip"
  11. "sync"
  12. "time"
  13. "github.com/google/uuid"
  14. "github.com/patrickmn/go-cache"
  15. "golang.org/x/time/rate"
  16. "github.com/mk6i/retro-aim-server/config"
  17. "github.com/mk6i/retro-aim-server/server/oscar/middleware"
  18. "github.com/mk6i/retro-aim-server/state"
  19. "github.com/mk6i/retro-aim-server/wire"
  20. )
  21. func NewServer(
  22. authService AuthService,
  23. buddyListRegistry BuddyListRegistry,
  24. chatSessionManager *state.InMemoryChatSessionManager,
  25. departureNotifier DepartureNotifier,
  26. logger *slog.Logger,
  27. onlineNotifier OnlineNotifier,
  28. SNACHandler func(ctx context.Context, serverType uint16, sess *state.Session, inFrame wire.SNACFrame, r io.Reader, rw ResponseWriter, listener config.Listener) error,
  29. rateLimitUpdater RateLimitUpdater,
  30. limits wire.SNACRateLimits,
  31. limiter *IPRateLimiter,
  32. listenerCfg []config.Listener,
  33. recalcWarning func(ctx context.Context, sess *state.Session) error,
  34. lowerWarnLevel func(ctx context.Context, sess *state.Session),
  35. ) *Server {
  36. oscarSvc := oscarServer{
  37. AuthService: authService,
  38. BuddyListRegistry: buddyListRegistry,
  39. ChatSessionManager: chatSessionManager,
  40. DepartureNotifier: departureNotifier,
  41. Logger: logger,
  42. OnlineNotifier: onlineNotifier,
  43. SNACHandler: SNACHandler,
  44. RateLimitUpdater: rateLimitUpdater,
  45. SNACRateLimits: limits,
  46. IPRateLimiter: limiter,
  47. recalcWarning: recalcWarning,
  48. lowerWarnLevel: lowerWarnLevel,
  49. }
  50. ctx, cancel := context.WithCancel(context.Background())
  51. return &Server{
  52. closed: make(chan struct{}),
  53. conns: make(map[net.Conn]struct{}),
  54. handler: oscarSvc.routeConnection,
  55. listenerCfg: listenerCfg,
  56. logger: logger,
  57. shutdownCancel: cancel,
  58. shutdownCtx: ctx,
  59. }
  60. }
  61. type Server struct {
  62. logger *slog.Logger
  63. listenerCfg []config.Listener
  64. listeners []net.Listener
  65. connMu sync.Mutex
  66. conns map[net.Conn]struct{}
  67. connWg sync.WaitGroup
  68. listenWg sync.WaitGroup
  69. shutdownCtx context.Context
  70. shutdownCancel context.CancelFunc
  71. closed chan struct{}
  72. handler func(ctx context.Context, conn net.Conn, listener config.Listener) error
  73. }
  74. func (s *Server) ListenAndServe() error {
  75. for _, listenCfg := range s.listenerCfg {
  76. ln, err := net.Listen("tcp", listenCfg.BOSListenAddress)
  77. if err != nil {
  78. s.cleanupListeners()
  79. s.shutdownCancel()
  80. return fmt.Errorf("failed to listen on %s: %w", listenCfg.BOSListenAddress, err)
  81. }
  82. args := []any{
  83. "listen_address", listenCfg.BOSListenAddress,
  84. "advertised_host_plain", listenCfg.BOSAdvertisedHostPlain,
  85. }
  86. if listenCfg.HasSSL {
  87. args = append(args, "advertised_host_ssl", listenCfg.BOSAdvertisedHostSSL)
  88. }
  89. s.logger.Info("starting server", args...)
  90. s.listeners = append(s.listeners, ln)
  91. s.listenWg.Add(1)
  92. go s.acceptLoop(ln, listenCfg)
  93. }
  94. <-s.closed // block until Shutdown is called
  95. return nil
  96. }
  97. func (s *Server) Shutdown(ctx context.Context) error {
  98. s.logger.Debug("Initiating graceful shutdown...")
  99. s.shutdownCancel()
  100. s.cleanupListeners()
  101. // Wait for handlers to complete
  102. done := make(chan struct{})
  103. go func() {
  104. s.connWg.Wait()
  105. s.listenWg.Wait()
  106. close(done)
  107. }()
  108. select {
  109. case <-done:
  110. s.logger.Info("shutdown complete")
  111. case <-ctx.Done():
  112. s.logger.Info("shutdown complete, but connections didn't close cleanly")
  113. }
  114. close(s.closed)
  115. return nil
  116. }
  117. func (s *Server) acceptLoop(ln net.Listener, listener config.Listener) {
  118. defer s.listenWg.Done()
  119. for {
  120. conn, err := ln.Accept()
  121. if err != nil {
  122. if errors.Is(err, net.ErrClosed) {
  123. return
  124. }
  125. s.logger.Error("accept error", "err", err.Error())
  126. continue
  127. }
  128. // track connection
  129. s.connMu.Lock()
  130. s.conns[conn] = struct{}{}
  131. s.connMu.Unlock()
  132. s.connWg.Add(1)
  133. go s.handleConnection(s.shutdownCtx, conn, listener)
  134. }
  135. }
  136. func (s *Server) handleConnection(ctx context.Context, conn net.Conn, listener config.Listener) {
  137. defer func() {
  138. // untrack connections
  139. s.connMu.Lock()
  140. delete(s.conns, conn)
  141. s.connMu.Unlock()
  142. _ = conn.Close()
  143. s.connWg.Done()
  144. }()
  145. if err := s.handler(ctx, conn, listener); err != nil {
  146. s.logger.InfoContext(ctx, "user session failed", "err", err.Error())
  147. }
  148. }
  149. func (s *Server) cleanupListeners() {
  150. for _, ln := range s.listeners {
  151. _ = ln.Close()
  152. }
  153. s.listeners = nil
  154. }
  155. type oscarServer struct {
  156. AuthService
  157. BuddyListRegistry
  158. ChatSessionManager
  159. DepartureNotifier
  160. Logger *slog.Logger
  161. OnlineNotifier
  162. SNACHandler func(ctx context.Context, serverType uint16, sess *state.Session, inFrame wire.SNACFrame, r io.Reader, rw ResponseWriter, listener config.Listener) error
  163. RateLimitUpdater
  164. wire.SNACRateLimits
  165. *IPRateLimiter
  166. recalcWarning func(ctx context.Context, sess *state.Session) error
  167. lowerWarnLevel func(ctx context.Context, sess *state.Session)
  168. }
  169. func (s oscarServer) routeConnection(ctx context.Context, conn net.Conn, listener config.Listener) error {
  170. ip, _, err := net.SplitHostPort(conn.RemoteAddr().String())
  171. if err != nil {
  172. s.Logger.Error("failed to parse remote address", "err", err.Error())
  173. return err
  174. }
  175. flapc := wire.NewFlapClient(100, conn, conn)
  176. if err := flapc.SendSignonFrame(nil); err != nil {
  177. return err
  178. }
  179. flap, err := flapc.ReceiveSignonFrame()
  180. if err != nil {
  181. return err
  182. }
  183. if flap.HasTag(wire.OServiceTLVTagsLoginCookie) {
  184. return s.connectToOSCARService(ctx, flap, flapc, conn, listener)
  185. }
  186. return s.authenticate(ctx, flap, ip, conn, flapc, listener.BOSAdvertisedHostPlain)
  187. }
  188. func (s oscarServer) connectToOSCARService(
  189. ctx context.Context,
  190. flap wire.FLAPSignonFrame,
  191. flapc *wire.FlapClient,
  192. conn net.Conn,
  193. listener config.Listener,
  194. ) error {
  195. authCookie, ok := flap.Bytes(wire.OServiceTLVTagsLoginCookie)
  196. if !ok {
  197. return errors.New("unable to get session id from payload")
  198. }
  199. cookie, err := s.CrackCookie(authCookie)
  200. if err != nil {
  201. return err
  202. }
  203. s.Logger.Debug("connecting to service", "service", wire.FoodGroupName(cookie.Service))
  204. var sess *state.Session
  205. switch cookie.Service {
  206. case wire.BOS:
  207. sess, err = s.AuthService.RegisterBOSSession(ctx, cookie)
  208. if err != nil {
  209. return err
  210. }
  211. if sess == nil {
  212. return errors.New("session not found")
  213. }
  214. if err := s.BuddyListRegistry.RegisterBuddyList(ctx, sess.IdentScreenName()); err != nil {
  215. return fmt.Errorf("unable to init buddy list: %w", err)
  216. }
  217. defer func() {
  218. sess.Close()
  219. ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
  220. defer cancel()
  221. if err := s.DepartureNotifier.BroadcastBuddyDeparted(ctx, sess); err != nil {
  222. s.Logger.ErrorContext(ctx, "error sending buddy departure notifications", "err", err.Error())
  223. }
  224. // buddy list must be cleared before session is closed, otherwise
  225. // there will be a race condition that could cause the buddy list
  226. // be prematurely deleted.
  227. if err := s.BuddyListRegistry.UnregisterBuddyList(ctx, sess.IdentScreenName()); err != nil {
  228. s.Logger.ErrorContext(ctx, "error removing buddy list entry", "err", err.Error())
  229. }
  230. s.ChatSessionManager.RemoveUserFromAllChats(sess.IdentScreenName())
  231. s.Signout(ctx, sess)
  232. }()
  233. remoteAddr, ok := ctx.Value("ip").(string)
  234. if ok {
  235. ip, err := netip.ParseAddrPort(remoteAddr)
  236. if err != nil {
  237. return errors.New("unable to parse ip addr")
  238. }
  239. sess.SetRemoteAddr(&ip)
  240. }
  241. if err := s.recalcWarning(ctx, sess); err != nil {
  242. return fmt.Errorf("failed to recalculate warning level: %w", err)
  243. }
  244. go s.lowerWarnLevel(ctx, sess)
  245. go s.receiveSessMessages(ctx, sess, flapc)
  246. case wire.Chat:
  247. sess, err = s.AuthService.RegisterChatSession(ctx, cookie)
  248. if err != nil {
  249. return err
  250. }
  251. if sess == nil {
  252. return errors.New("session not found")
  253. }
  254. defer func() {
  255. ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
  256. defer cancel()
  257. s.SignoutChat(ctx, sess)
  258. }()
  259. go s.receiveSessMessages(ctx, sess, flapc)
  260. default:
  261. sess, err = s.AuthService.RetrieveBOSSession(ctx, cookie)
  262. if err != nil {
  263. return err
  264. }
  265. if sess == nil {
  266. return errors.New("session not found")
  267. }
  268. }
  269. ctx = context.WithValue(ctx, "screenName", sess.IdentScreenName())
  270. msg := s.OnlineNotifier.HostOnline(cookie.Service)
  271. if err := flapc.SendSNAC(msg.Frame, msg.Body); err != nil {
  272. return err
  273. }
  274. return s.dispatchIncomingMessages(ctx, cookie.Service, sess, flapc, conn, listener)
  275. }
  276. func (s oscarServer) receiveSessMessages(ctx context.Context, sess *state.Session, flapc *wire.FlapClient) {
  277. for {
  278. select {
  279. case <-sess.Closed():
  280. return
  281. case <-ctx.Done():
  282. return
  283. case m := <-sess.ReceiveMessage():
  284. // forward a notification sent from another client to this client
  285. if err := flapc.SendSNAC(m.Frame, m.Body); err != nil {
  286. middleware.LogRequestError(ctx, s.Logger, m.Frame, err)
  287. } else {
  288. middleware.LogRequest(ctx, s.Logger, m.Frame, m.Body)
  289. }
  290. }
  291. }
  292. }
  293. func (s oscarServer) authenticate(
  294. ctx context.Context,
  295. flap wire.FLAPSignonFrame,
  296. ip string,
  297. conn net.Conn,
  298. flapc *wire.FlapClient,
  299. advertisedHost string,
  300. ) error {
  301. if ok, isBUCP := s.Allow(ip); !ok {
  302. s.Logger.Error("user rate limited at login", "remote", ip)
  303. tlv := wire.TLVRestBlock{
  304. TLVList: []wire.TLV{
  305. wire.NewTLVBE(wire.LoginTLVTagsErrorSubcode, wire.LoginErrRateLimitExceeded),
  306. },
  307. }
  308. // gives wrong response if you quickly switch between BUCP/FLAP clients
  309. if isBUCP {
  310. return flapc.SendSNAC(
  311. wire.SNACFrame{
  312. FoodGroup: wire.BUCP,
  313. SubGroup: wire.BUCPLoginResponse,
  314. },
  315. wire.SNAC_0x17_0x03_BUCPLoginResponse{
  316. TLVRestBlock: tlv,
  317. },
  318. )
  319. } else {
  320. return flapc.NewSignoff(tlv)
  321. }
  322. }
  323. // auth must complete within the next 30 seconds
  324. if err := conn.SetDeadline(time.Now().Add(30 * time.Second)); err != nil {
  325. return fmt.Errorf("failed to set deadline: %w", err)
  326. }
  327. // decide whether the client is using BUCP or FLAP authentication based on
  328. // the presence of the screen name TLV. this block used to check for the
  329. // presence of the roasted password TLV, however that proved an unreliable
  330. // indicator of FLAP-auth because older ICQ clients appear to omit the
  331. // roasted password TLV when the password is not stored client-side.
  332. if _, hasScreenName := flap.Uint16BE(wire.LoginTLVTagsScreenName); hasScreenName {
  333. return s.processFLAPAuth(ctx, flap, flapc, advertisedHost)
  334. }
  335. s.SetBUCP(ip)
  336. return s.processBUCPAuth(ctx, flapc, advertisedHost)
  337. }
  338. func (s oscarServer) processFLAPAuth(
  339. ctx context.Context,
  340. signonFrame wire.FLAPSignonFrame,
  341. flapc *wire.FlapClient,
  342. advertisedHost string,
  343. ) error {
  344. tlv, err := s.AuthService.FLAPLogin(ctx, signonFrame, state.NewStubUser, advertisedHost)
  345. if err != nil {
  346. return err
  347. }
  348. return flapc.NewSignoff(tlv)
  349. }
  350. func (s oscarServer) processBUCPAuth(ctx context.Context, flapc *wire.FlapClient, advertisedHost string) error {
  351. frames := 0
  352. for {
  353. frame, err := flapc.ReceiveFLAP()
  354. if err != nil {
  355. return err
  356. }
  357. if frames > 10 {
  358. // a lot of frames received, the client is misbehaving
  359. return fmt.Errorf("too many auth flap packets received")
  360. }
  361. frames++
  362. switch frame.FrameType {
  363. case wire.FLAPFrameSignoff:
  364. s.Logger.Debug("signed off mid-login")
  365. return io.EOF // client disconnected
  366. case wire.FLAPFrameKeepAlive:
  367. s.Logger.Debug("received flap keepalive frame")
  368. case wire.FLAPFrameData:
  369. buf := bytes.NewReader(frame.Payload)
  370. fr := wire.SNACFrame{}
  371. if err := wire.UnmarshalBE(&fr, buf); err != nil {
  372. return err
  373. }
  374. switch {
  375. case fr.FoodGroup == wire.BUCP && fr.SubGroup == wire.BUCPChallengeRequest:
  376. challengeRequest := wire.SNAC_0x17_0x06_BUCPChallengeRequest{}
  377. if err := wire.UnmarshalBE(&challengeRequest, buf); err != nil {
  378. return err
  379. }
  380. outSNAC, err := s.BUCPChallenge(ctx, challengeRequest, uuid.New)
  381. if err != nil {
  382. return err
  383. }
  384. if err := flapc.SendSNAC(outSNAC.Frame, outSNAC.Body); err != nil {
  385. return err
  386. }
  387. if outSNAC.Frame.SubGroup == wire.BUCPLoginResponse {
  388. screenName, _ := challengeRequest.String(wire.LoginTLVTagsScreenName)
  389. s.Logger.Debug("failed BUCP challenge: user does not exist", "screen_name", screenName)
  390. return nil // account does not exist
  391. }
  392. case fr.FoodGroup == wire.BUCP && fr.SubGroup == wire.BUCPLoginRequest:
  393. loginRequest := wire.SNAC_0x17_0x02_BUCPLoginRequest{}
  394. if err := wire.UnmarshalBE(&loginRequest, buf); err != nil {
  395. return err
  396. }
  397. outSNAC, err := s.BUCPLogin(ctx, loginRequest, state.NewStubUser, advertisedHost)
  398. if err != nil {
  399. return err
  400. }
  401. return flapc.SendSNAC(outSNAC.Frame, outSNAC.Body)
  402. default:
  403. s.Logger.Debug("unexpected SNAC received during login",
  404. "foodgroup", wire.FoodGroupName(fr.FoodGroup),
  405. "subgroup", wire.SubGroupName(fr.FoodGroup, fr.SubGroup))
  406. return io.EOF
  407. }
  408. default:
  409. s.Logger.Debug("unexpected frame type received during login", "type", frame.FrameType)
  410. return io.EOF
  411. }
  412. }
  413. }
  414. func sendInvalidSNACErr(frameIn wire.SNACFrame, rw ResponseWriter) error {
  415. frameOut := wire.SNACFrame{
  416. FoodGroup: frameIn.FoodGroup,
  417. SubGroup: 0x01, // error subgroup for all SNACs
  418. RequestID: frameIn.RequestID,
  419. }
  420. bodyOut := wire.SNACError{
  421. Code: wire.ErrorCodeInvalidSnac,
  422. }
  423. return rw.SendSNAC(frameOut, bodyOut)
  424. }
  425. // dispatchIncomingMessages receives incoming messages and sends them to the
  426. // appropriate message handler. Messages from the client are sent to the
  427. // router. Messages relayed from the user session are forwarded to the client.
  428. // This function ensures that the same sequence number is incremented for both
  429. // types of messages. The function terminates upon receiving a connection error
  430. // or when the session closes.
  431. func (s oscarServer) dispatchIncomingMessages(
  432. ctx context.Context,
  433. fg uint16,
  434. sess *state.Session,
  435. flapc *wire.FlapClient,
  436. r io.ReadCloser,
  437. listener config.Listener,
  438. ) error {
  439. defer func() {
  440. s.Logger.InfoContext(ctx, "user disconnected")
  441. }()
  442. // buffered so that the go routine has room to exit
  443. msgCh := make(chan wire.FLAPFrame, 1)
  444. errCh := make(chan error, 1)
  445. // consume flap frames
  446. go func() {
  447. defer close(msgCh)
  448. defer close(errCh)
  449. for {
  450. frame := wire.FLAPFrame{}
  451. if err := wire.UnmarshalBE(&frame, r); err != nil {
  452. errCh <- err
  453. return
  454. }
  455. msgCh <- frame
  456. }
  457. }()
  458. for {
  459. select {
  460. case flap, ok := <-msgCh:
  461. if !ok {
  462. return nil
  463. }
  464. switch flap.FrameType {
  465. case wire.FLAPFrameData:
  466. flapBuf := bytes.NewBuffer(flap.Payload)
  467. inFrame := wire.SNACFrame{}
  468. if err := wire.UnmarshalBE(&inFrame, flapBuf); err != nil {
  469. return err
  470. }
  471. rateClassID, ok := s.SNACRateLimits.RateClassLookup(inFrame.FoodGroup, inFrame.SubGroup)
  472. if ok {
  473. if status := sess.EvaluateRateLimit(time.Now(), rateClassID); status == wire.RateLimitStatusLimited {
  474. s.Logger.DebugContext(ctx, "rate limit exceeded, dropping SNAC",
  475. "foodgroup", wire.FoodGroupName(inFrame.FoodGroup),
  476. "subgroup", wire.SubGroupName(inFrame.FoodGroup, inFrame.SubGroup))
  477. break
  478. }
  479. } else {
  480. s.Logger.ErrorContext(ctx, "rate limit not found, allowing request through")
  481. }
  482. // route a client request to the appropriate service handler. the
  483. // handler may write a response to the client connection.
  484. if err := s.SNACHandler(ctx, fg, sess, inFrame, flapBuf, flapc, listener); err != nil {
  485. middleware.LogRequestError(ctx, s.Logger, inFrame, err)
  486. if errors.Is(err, ErrRouteNotFound) {
  487. if err1 := sendInvalidSNACErr(inFrame, flapc); err1 != nil {
  488. return errors.Join(err1, err)
  489. }
  490. break
  491. }
  492. return err
  493. }
  494. case wire.FLAPFrameSignon:
  495. return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", flap)
  496. case wire.FLAPFrameError:
  497. return fmt.Errorf("got FLAPFrameError. flap: %v", flap)
  498. case wire.FLAPFrameSignoff:
  499. s.Logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", flap)
  500. return nil
  501. case wire.FLAPFrameKeepAlive:
  502. s.Logger.DebugContext(ctx, "keepalive heartbeat")
  503. default:
  504. return fmt.Errorf("got unknown FLAP frame type. flap: %v", flap)
  505. }
  506. case <-time.After(1 * time.Second):
  507. updates := s.RateLimitUpdater.RateLimitUpdates(ctx, sess, time.Now())
  508. for _, update := range updates {
  509. if err := flapc.SendSNAC(update.Frame, update.Body); err != nil {
  510. middleware.LogRequestError(ctx, s.Logger, update.Frame, err)
  511. return err
  512. }
  513. }
  514. case <-sess.Closed():
  515. // add logoff reason to clients that support multi-conn
  516. if sess.MultiConnFlag() == wire.MultiConnFlagsOldClient {
  517. if err := flapc.OldSignoff(); err != nil {
  518. return fmt.Errorf("unable to gracefully disconnect user. %w", err)
  519. }
  520. } else {
  521. block := wire.TLVRestBlock{}
  522. // error code indicating user signed in a different location
  523. block.Append(wire.NewTLVBE(0x0009, wire.OServiceDiscErrNewLogin))
  524. // "more info" button
  525. block.Append(wire.NewTLVBE(0x000b, "https://github.com/mk6i/retro-aim-server"))
  526. if err := flapc.NewSignoff(block); err != nil {
  527. return fmt.Errorf("unable to gracefully disconnect user. %w", err)
  528. }
  529. }
  530. return nil
  531. case <-ctx.Done():
  532. if sess.MultiConnFlag() == wire.MultiConnFlagsOldClient {
  533. if err := flapc.OldSignoff(); err != nil {
  534. return fmt.Errorf("unable to gracefully disconnect user. %w", err)
  535. }
  536. } else {
  537. if err := flapc.NewSignoff(wire.TLVRestBlock{}); err != nil {
  538. return fmt.Errorf("unable to gracefully disconnect user. %w", err)
  539. }
  540. }
  541. return nil
  542. case err := <-errCh:
  543. if !errors.Is(io.EOF, err) {
  544. s.Logger.ErrorContext(ctx, "client disconnected with error", "err", err)
  545. }
  546. return nil
  547. }
  548. }
  549. }
  550. // IPRateLimiter enforces a per-IP rate limit using a token bucket algorithm.
  551. // It caches individual rate limiters by IP address and supports tagging requests
  552. // as originating from the BUCP or FLAP auth.
  553. //
  554. // The limiter uses an in-memory cache with TTL expiration, so rate limits reset
  555. // after the TTL if no activity is observed for a given IP.
  556. type IPRateLimiter struct {
  557. cache *cache.Cache // In-memory cache mapping IPs to rate limiters with optional BUCP tag
  558. rate rate.Limit // Requests allowed per second
  559. burst int // Maximum burst size allowed
  560. }
  561. type rateLimitEntry struct {
  562. isBUCP bool
  563. limiter *rate.Limiter
  564. }
  565. // NewIPRateLimiter initializes a new IPRateLimiter with the specified rate,
  566. // burst size, and TTL for each IP's limiter. Entries expire after 2×TTL.
  567. func NewIPRateLimiter(rate rate.Limit, burst int, ttl time.Duration) *IPRateLimiter {
  568. return &IPRateLimiter{
  569. cache: cache.New(ttl, 2*ttl),
  570. rate: rate,
  571. burst: burst,
  572. }
  573. }
  574. // SetBUCP marks the rate limiter for the given IP as originating from BUCP auth
  575. // (default FLAP auth).
  576. func (l *IPRateLimiter) SetBUCP(ip string) {
  577. limiter, found := l.cache.Get(ip)
  578. if !found {
  579. limiter = &rateLimitEntry{
  580. isBUCP: true,
  581. limiter: rate.NewLimiter(l.rate, l.burst),
  582. }
  583. l.cache.Set(ip, limiter, cache.DefaultExpiration)
  584. }
  585. limiter.(*rateLimitEntry).isBUCP = true
  586. }
  587. // Allow checks if a request from the given IP is allowed under its rate limit.
  588. // It returns whether the request is allowed and whether the connection uses
  589. // BUCP auth.
  590. func (l *IPRateLimiter) Allow(ip string) (allowed bool, isBUCP bool) {
  591. limiter, found := l.cache.Get(ip)
  592. if !found {
  593. limiter = &rateLimitEntry{
  594. limiter: rate.NewLimiter(l.rate, l.burst),
  595. }
  596. l.cache.Set(ip, limiter, cache.DefaultExpiration)
  597. }
  598. entry := limiter.(*rateLimitEntry)
  599. return entry.limiter.Allow(), entry.isBUCP
  600. }