server.go 19 KB

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