server.go 18 KB

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