| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743 |
- package oscar
- import (
- "bytes"
- "context"
- "errors"
- "fmt"
- "io"
- "log/slog"
- "net"
- "net/netip"
- "sync"
- "time"
- "github.com/google/uuid"
- "github.com/patrickmn/go-cache"
- "golang.org/x/time/rate"
- "github.com/mk6i/open-oscar-server/config"
- "github.com/mk6i/open-oscar-server/server/oscar/middleware"
- "github.com/mk6i/open-oscar-server/state"
- "github.com/mk6i/open-oscar-server/wire"
- )
- func NewServer(
- authService AuthService,
- buddyListRegistry BuddyListRegistry,
- chatSessionManager *state.InMemoryChatSessionManager,
- departureNotifier DepartureNotifier,
- logger *slog.Logger,
- onlineNotifier OnlineNotifier,
- SNACHandler func(ctx context.Context, serverType uint16, instance *state.SessionInstance, inFrame wire.SNACFrame, r io.Reader, rw ResponseWriter, listener config.Listener) error,
- rateLimitUpdater RateLimitUpdater,
- limits wire.SNACRateLimits,
- limiter *IPRateLimiter,
- listenerCfg []config.Listener,
- recalcWarning func(ctx context.Context, instance *state.SessionInstance) error,
- lowerWarnLevel func(ctx context.Context, instance *state.SessionInstance),
- ) *Server {
- oscarSvc := oscarServer{
- authService: authService,
- buddyListRegistry: buddyListRegistry,
- chatSessionManager: chatSessionManager,
- departureNotifier: departureNotifier,
- logger: logger,
- onlineNotifier: onlineNotifier,
- snacHandler: SNACHandler,
- rateLimitUpdater: rateLimitUpdater,
- rateLimits: limits,
- ipRateLimiter: limiter,
- recalcWarning: recalcWarning,
- lowerWarnLevel: lowerWarnLevel,
- }
- ctx, cancel := context.WithCancel(context.Background())
- return &Server{
- closed: make(chan struct{}),
- conns: make(map[net.Conn]struct{}),
- handler: oscarSvc.routeConnection,
- listenerCfg: listenerCfg,
- logger: logger,
- shutdownCancel: cancel,
- shutdownCtx: ctx,
- }
- }
- type Server struct {
- logger *slog.Logger
- listenerCfg []config.Listener
- listeners []net.Listener
- connMu sync.Mutex
- conns map[net.Conn]struct{}
- connWg sync.WaitGroup
- listenWg sync.WaitGroup
- shutdownCtx context.Context
- shutdownCancel context.CancelFunc
- closed chan struct{}
- handler func(ctx context.Context, conn net.Conn, listener config.Listener) error
- }
- func (s *Server) ListenAndServe() error {
- for _, listenCfg := range s.listenerCfg {
- ln, err := net.Listen("tcp", listenCfg.BOSListenAddress)
- if err != nil {
- s.cleanupListeners()
- s.shutdownCancel()
- return fmt.Errorf("failed to listen on %s: %w", listenCfg.BOSListenAddress, err)
- }
- args := []any{
- "listen_address", listenCfg.BOSListenAddress,
- "advertised_host_plain", listenCfg.BOSAdvertisedHostPlain,
- }
- if listenCfg.HasSSL {
- args = append(args, "advertised_host_ssl", listenCfg.BOSAdvertisedHostSSL)
- }
- s.logger.Info("starting server", args...)
- s.listeners = append(s.listeners, ln)
- s.listenWg.Add(1)
- go s.acceptLoop(ln, listenCfg)
- }
- <-s.closed // block until Shutdown is called
- return nil
- }
- func (s *Server) Shutdown(ctx context.Context) error {
- s.logger.Debug("Initiating graceful shutdown...")
- s.shutdownCancel()
- s.cleanupListeners()
- // Wait for handlers to complete
- done := make(chan struct{})
- go func() {
- s.connWg.Wait()
- s.listenWg.Wait()
- close(done)
- }()
- select {
- case <-done:
- s.logger.Info("shutdown complete")
- case <-ctx.Done():
- s.logger.Info("shutdown complete, but connections didn't close cleanly")
- }
- close(s.closed)
- return nil
- }
- func (s *Server) acceptLoop(ln net.Listener, listener config.Listener) {
- defer s.listenWg.Done()
- for {
- conn, err := ln.Accept()
- if err != nil {
- if errors.Is(err, net.ErrClosed) {
- return
- }
- s.logger.Error("accept error", "err", err.Error())
- continue
- }
- // track connection
- s.connMu.Lock()
- s.conns[conn] = struct{}{}
- s.connMu.Unlock()
- s.connWg.Add(1)
- go s.handleConnection(s.shutdownCtx, conn, listener)
- }
- }
- func (s *Server) handleConnection(ctx context.Context, conn net.Conn, listener config.Listener) {
- defer func() {
- // untrack connections
- s.connMu.Lock()
- delete(s.conns, conn)
- s.connMu.Unlock()
- _ = conn.Close()
- s.connWg.Done()
- }()
- ctx = middleware.WithIP(ctx, conn.RemoteAddr().String())
- if err := s.handler(ctx, conn, listener); err != nil {
- s.logger.InfoContext(ctx, "user session failed", "err", err.Error())
- }
- }
- func (s *Server) cleanupListeners() {
- for _, ln := range s.listeners {
- _ = ln.Close()
- }
- s.listeners = nil
- }
- type oscarServer struct {
- authService AuthService
- buddyListRegistry BuddyListRegistry
- chatSessionManager ChatSessionManager
- departureNotifier DepartureNotifier
- logger *slog.Logger
- onlineNotifier OnlineNotifier
- snacHandler func(ctx context.Context, serverType uint16, instance *state.SessionInstance, inFrame wire.SNACFrame, r io.Reader, rw ResponseWriter, listener config.Listener) error
- rateLimitUpdater RateLimitUpdater
- rateLimits wire.SNACRateLimits
- ipRateLimiter *IPRateLimiter
- recalcWarning func(ctx context.Context, instance *state.SessionInstance) error
- lowerWarnLevel func(ctx context.Context, instance *state.SessionInstance)
- }
- func (s oscarServer) routeConnection(ctx context.Context, conn net.Conn, listener config.Listener) error {
- ip, _, err := net.SplitHostPort(conn.RemoteAddr().String())
- if err != nil {
- s.logger.Error("failed to parse remote address", "err", err.Error())
- return err
- }
- flapc := wire.NewFlapClient(100, conn, conn)
- // send flap signon with server capabilities
- if err := flapc.SendSignonFrame(nil); err != nil {
- return err
- }
- flap, err := flapc.ReceiveSignonFrame()
- if err != nil {
- return err
- }
- if flap.HasTag(wire.OServiceTLVTagsLoginCookie) {
- return s.connectToOSCARService(ctx, flap, flapc, conn, listener)
- }
- return s.authenticate(ctx, flap, ip, conn, flapc, listener.BOSAdvertisedHostPlain)
- }
- func (s oscarServer) connectToOSCARService(
- ctx context.Context,
- flap wire.FLAPSignonFrame,
- flapc *wire.FlapClient,
- conn net.Conn,
- listener config.Listener,
- ) error {
- authCookie, ok := flap.Bytes(wire.OServiceTLVTagsLoginCookie)
- if !ok {
- return errors.New("unable to get session id from payload")
- }
- cookie, err := s.authService.CrackCookie(authCookie)
- if err != nil {
- return err
- }
- s.logger.Debug("connecting to service", "service", wire.FoodGroupName(cookie.Service))
- var instance *state.SessionInstance
- switch cookie.Service {
- case wire.BOS:
- sessCfg := func(sess *state.Session) {
- sess.OnSessionClose(func() {
- if !shuttingDown(ctx) {
- if err := s.departureNotifier.BroadcastBuddyDeparted(ctx, sess.IdentScreenName()); err != nil {
- s.logger.ErrorContext(ctx, "error sending buddy departure notifications", "err", err.Error())
- }
- }
- ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
- defer cancel()
- // buddy list must be cleared before session is closed, otherwise
- // there will be a race condition that could cause the buddy list
- // be prematurely deleted.
- if err := s.buddyListRegistry.UnregisterBuddyList(ctx, instance.IdentScreenName()); err != nil {
- s.logger.ErrorContext(ctx, "error removing buddy list entry", "err", err.Error())
- }
- s.chatSessionManager.RemoveUserFromAllChats(instance.IdentScreenName())
- s.authService.Signout(ctx, sess)
- })
- }
- instance, err = s.authService.RegisterBOSSession(ctx, cookie, sessCfg)
- if err != nil {
- if errors.Is(err, state.ErrMaxConcurrentSessionsReached) {
- s.logger.Debug("session registration failed", "err", err.Error())
- block := wire.TLVRestBlock{}
- // error code indicating the signon is blocked. i can't find a
- // more appropriate error code to indicate the maximum session limit is reached
- block.Append(wire.NewTLVBE(0x0008, uint8(0x18)))
- if err := flapc.NewSignoff(block); err != nil {
- return fmt.Errorf("unable to gracefully disconnect user. %w", err)
- }
- return nil
- }
- return err
- }
- if instance == nil {
- return errors.New("session not found")
- }
- defer func() {
- instance.CloseInstance()
- }()
- if err = instance.Session().RunOnce(func() error {
- // make buddy list visible to other users
- if err := s.buddyListRegistry.RegisterBuddyList(ctx, instance.IdentScreenName()); err != nil {
- return fmt.Errorf("unable to init buddy list: %w", err)
- }
- // restore warning level from last session
- if err := s.recalcWarning(ctx, instance); err != nil {
- return fmt.Errorf("failed to recalculate warning level: %w", err)
- }
- // periodically decay warning level
- go s.lowerWarnLevel(ctx, instance)
- return nil
- }); err != nil {
- return err
- }
- // Update user visibility when an instance closes, as the user's overall status may change.
- // Example: With 1 away and 1 non-away instance, the user appears available. If the non-away
- // instance closes, the user should appear away.
- instance.OnClose(func() {
- if shuttingDown(ctx) {
- return
- }
- if instance.Session().Invisible() {
- if err := s.departureNotifier.BroadcastBuddyDeparted(ctx, instance.IdentScreenName()); err != nil {
- s.logger.ErrorContext(ctx, "error sending buddy departure notifications", "err", err.Error())
- }
- } else {
- if err := s.departureNotifier.BroadcastBuddyArrived(ctx, instance.IdentScreenName(), instance.Session().TLVUserInfo()); err != nil {
- s.logger.ErrorContext(ctx, "error sending buddy arrival notifications", "err", err.Error())
- }
- }
- })
- if remoteAddr, ok := middleware.IPFromContext(ctx); ok {
- ip, err := netip.ParseAddrPort(remoteAddr)
- if err != nil {
- return errors.New("unable to parse ip addr")
- }
- instance.SetRemoteAddr(&ip)
- }
- go s.receiveSessMessages(ctx, instance, flapc)
- case wire.Chat:
- sessCfg := func(sess *state.Session) {
- sess.OnSessionClose(func() {
- ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
- defer cancel()
- s.authService.SignoutChat(ctx, sess)
- })
- }
- instance, err = s.authService.RegisterChatSession(ctx, cookie, sessCfg)
- if err != nil {
- return err
- }
- if instance == nil {
- return errors.New("session not found")
- }
- defer func() {
- instance.CloseInstance()
- }()
- go s.receiveSessMessages(ctx, instance, flapc)
- default:
- instance, err = s.authService.RetrieveBOSSession(ctx, cookie)
- if err != nil {
- return err
- }
- if instance == nil {
- return errors.New("session not found")
- }
- }
- ctx = middleware.WithScreenName(ctx, instance.IdentScreenName())
- msg := s.onlineNotifier.HostOnline(cookie.Service)
- if err := flapc.SendSNAC(msg.Frame, msg.Body); err != nil {
- return err
- }
- return s.dispatchIncomingMessages(ctx, cookie.Service, instance, flapc, conn, listener)
- }
- func shuttingDown(ctx context.Context) bool {
- select {
- case <-ctx.Done():
- // server is shutting down, don't send buddy notifications
- return true
- default:
- }
- return false
- }
- func (s oscarServer) receiveSessMessages(ctx context.Context, instance *state.SessionInstance, flapc *wire.FlapClient) {
- for {
- select {
- case <-instance.Closed():
- return
- case <-ctx.Done():
- return
- case m := <-instance.ReceiveMessage():
- // forward a notification sent from another client to this client
- if err := flapc.SendSNAC(m.Frame, m.Body); err != nil {
- middleware.LogRequestError(ctx, s.logger, m.Frame, err)
- } else {
- middleware.LogRequest(ctx, s.logger, m.Frame, m.Body)
- }
- }
- }
- }
- func (s oscarServer) authenticate(
- ctx context.Context,
- flap wire.FLAPSignonFrame,
- ip string,
- conn net.Conn,
- flapc *wire.FlapClient,
- advertisedHost string,
- ) error {
- if ok, isBUCP := s.ipRateLimiter.Allow(ip); !ok {
- s.logger.InfoContext(ctx, "user rate limited at login, dropping connection")
- tlv := wire.TLVRestBlock{
- TLVList: []wire.TLV{
- wire.NewTLVBE(wire.LoginTLVTagsErrorSubcode, wire.LoginErrRateLimitExceeded),
- },
- }
- // gives wrong response if you quickly switch between BUCP/FLAP clients
- if isBUCP {
- return flapc.SendSNAC(
- wire.SNACFrame{
- FoodGroup: wire.BUCP,
- SubGroup: wire.BUCPLoginResponse,
- },
- wire.SNAC_0x17_0x03_BUCPLoginResponse{
- TLVRestBlock: tlv,
- },
- )
- } else {
- return flapc.NewSignoff(tlv)
- }
- }
- // auth must complete within the next 30 seconds
- if err := conn.SetDeadline(time.Now().Add(30 * time.Second)); err != nil {
- return fmt.Errorf("failed to set deadline: %w", err)
- }
- // decide whether the client is using BUCP or FLAP authentication based on
- // the presence of the screen name TLV. this block used to check for the
- // presence of the roasted password TLV, however that proved an unreliable
- // indicator of FLAP-auth because older ICQ clients appear to omit the
- // roasted password TLV when the password is not stored client-side.
- if _, hasScreenName := flap.Uint16BE(wire.LoginTLVTagsScreenName); hasScreenName {
- return s.processFLAPAuth(ctx, flap, flapc, advertisedHost)
- }
- s.ipRateLimiter.SetBUCP(ip)
- return s.processBUCPAuth(ctx, flapc, advertisedHost)
- }
- func (s oscarServer) processFLAPAuth(
- ctx context.Context,
- signonFrame wire.FLAPSignonFrame,
- flapc *wire.FlapClient,
- advertisedHost string,
- ) error {
- tlv, err := s.authService.FLAPLogin(ctx, signonFrame, advertisedHost)
- if err != nil {
- return err
- }
- return flapc.NewSignoff(tlv)
- }
- func (s oscarServer) processBUCPAuth(ctx context.Context, flapc *wire.FlapClient, advertisedHost string) error {
- frames := 0
- for {
- frame, err := flapc.ReceiveFLAP()
- if err != nil {
- return err
- }
- if frames > 10 {
- // a lot of frames received, the client is misbehaving
- return fmt.Errorf("too many auth flap packets received")
- }
- frames++
- switch frame.FrameType {
- case wire.FLAPFrameSignoff:
- s.logger.Debug("signed off mid-login")
- return io.EOF // client disconnected
- case wire.FLAPFrameKeepAlive:
- s.logger.Debug("received flap keepalive frame")
- case wire.FLAPFrameData:
- buf := bytes.NewReader(frame.Payload)
- fr := wire.SNACFrame{}
- if err := wire.UnmarshalBE(&fr, buf); err != nil {
- return err
- }
- switch {
- case fr.FoodGroup == wire.BUCP && fr.SubGroup == wire.BUCPChallengeRequest:
- challengeRequest := wire.SNAC_0x17_0x06_BUCPChallengeRequest{}
- if err := wire.UnmarshalBE(&challengeRequest, buf); err != nil {
- return err
- }
- outSNAC, err := s.authService.BUCPChallenge(ctx, challengeRequest, uuid.New)
- if err != nil {
- return err
- }
- outSNAC.Frame.RequestID = fr.RequestID
- if err := flapc.SendSNAC(outSNAC.Frame, outSNAC.Body); err != nil {
- return err
- }
- if outSNAC.Frame.SubGroup == wire.BUCPLoginResponse {
- screenName, _ := challengeRequest.String(wire.LoginTLVTagsScreenName)
- s.logger.Debug("failed BUCP challenge: user does not exist", "screen_name", screenName)
- return nil // account does not exist
- }
- case fr.FoodGroup == wire.BUCP && fr.SubGroup == wire.BUCPLoginRequest:
- loginRequest := wire.SNAC_0x17_0x02_BUCPLoginRequest{}
- if err := wire.UnmarshalBE(&loginRequest, buf); err != nil {
- return err
- }
- outSNAC, err := s.authService.BUCPLogin(ctx, loginRequest, advertisedHost)
- if err != nil {
- return err
- }
- outSNAC.Frame.RequestID = fr.RequestID
- // Clients expect login response as SNAC on FLAP
- // channel 2 followed by a FLAP signoff frame to properly close the auth
- // connection
- if err := flapc.SendSNAC(outSNAC.Frame, outSNAC.Body); err != nil {
- return err
- }
- return flapc.NewSignoff(wire.TLVRestBlock{})
- default:
- s.logger.Debug("unexpected SNAC received during login",
- "foodgroup", wire.FoodGroupName(fr.FoodGroup),
- "subgroup", wire.SubGroupName(fr.FoodGroup, fr.SubGroup))
- return io.EOF
- }
- default:
- s.logger.Debug("unexpected frame type received during login", "type", frame.FrameType)
- return io.EOF
- }
- }
- }
- func sendInvalidSNACErr(frameIn wire.SNACFrame, rw ResponseWriter) error {
- frameOut := wire.SNACFrame{
- FoodGroup: frameIn.FoodGroup,
- SubGroup: 0x01, // error subgroup for all SNACs
- RequestID: frameIn.RequestID,
- }
- bodyOut := wire.SNACError{
- Code: wire.ErrorCodeInvalidSnac,
- }
- return rw.SendSNAC(frameOut, bodyOut)
- }
- // dispatchIncomingMessages receives incoming messages and sends them to the
- // appropriate message handler. Messages from the client are sent to the
- // router. Messages relayed from the user session are forwarded to the client.
- // This function ensures that the same sequence number is incremented for both
- // types of messages. The function terminates upon receiving a connection error
- // or when the session closes.
- func (s oscarServer) dispatchIncomingMessages(
- ctx context.Context,
- fg uint16,
- instance *state.SessionInstance,
- flapc *wire.FlapClient,
- r io.ReadCloser,
- listener config.Listener,
- ) error {
- defer func() {
- s.logger.InfoContext(ctx, "user disconnected")
- }()
- // buffered so that the go routine has room to exit
- msgCh := make(chan wire.FLAPFrame, 1)
- errCh := make(chan error, 1)
- // consume flap frames
- go func() {
- defer close(msgCh)
- defer close(errCh)
- for {
- frame := wire.FLAPFrame{}
- if err := wire.UnmarshalBE(&frame, r); err != nil {
- errCh <- err
- return
- }
- msgCh <- frame
- }
- }()
- for {
- select {
- case flap, ok := <-msgCh:
- if !ok {
- return nil
- }
- switch flap.FrameType {
- case wire.FLAPFrameData:
- flapBuf := bytes.NewBuffer(flap.Payload)
- inFrame := wire.SNACFrame{}
- if err := wire.UnmarshalBE(&inFrame, flapBuf); err != nil {
- return err
- }
- rateClassID, ok := s.rateLimits.RateClassLookup(inFrame.FoodGroup, inFrame.SubGroup)
- if ok {
- if status := instance.Session().EvaluateRateLimit(time.Now(), rateClassID); status == wire.RateLimitStatusLimited {
- s.logger.DebugContext(ctx, "rate limit exceeded, dropping SNAC",
- "foodgroup", wire.FoodGroupName(inFrame.FoodGroup),
- "subgroup", wire.SubGroupName(inFrame.FoodGroup, inFrame.SubGroup))
- break
- }
- } else {
- s.logger.ErrorContext(ctx, "rate limit not found, allowing request through")
- }
- // route a client request to the appropriate service handler. the
- // handler may write a response to the client connection.
- if err := s.snacHandler(ctx, fg, instance, inFrame, flapBuf, flapc, listener); err != nil {
- middleware.LogRequestError(ctx, s.logger, inFrame, err)
- if errors.Is(err, ErrRouteNotFound) {
- if err1 := sendInvalidSNACErr(inFrame, flapc); err1 != nil {
- return errors.Join(err1, err)
- }
- break
- }
- return err
- }
- case wire.FLAPFrameSignon:
- return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", flap)
- case wire.FLAPFrameError:
- return fmt.Errorf("got FLAPFrameError. flap: %v", flap)
- case wire.FLAPFrameSignoff:
- s.logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", flap)
- return nil
- case wire.FLAPFrameKeepAlive:
- s.logger.DebugContext(ctx, "keepalive heartbeat")
- default:
- return fmt.Errorf("got unknown FLAP frame type. flap: %v", flap)
- }
- case <-time.After(1 * time.Second):
- updates := s.rateLimitUpdater.RateLimitUpdates(ctx, instance, time.Now())
- for _, update := range updates {
- if err := flapc.SendSNAC(update.Frame, update.Body); err != nil {
- middleware.LogRequestError(ctx, s.logger, update.Frame, err)
- return err
- }
- }
- case <-instance.Closed():
- // add logoff reason to clients that support multi-conn
- if instance.MultiConnFlag() == wire.MultiConnFlagsOldClient {
- if err := flapc.OldSignoff(); err != nil {
- return fmt.Errorf("unable to gracefully disconnect user. %w", err)
- }
- } else {
- block := wire.TLVRestBlock{}
- // error code indicating user signed in a different location
- block.Append(wire.NewTLVBE(0x0009, wire.OServiceDiscErrNewLogin))
- // "more info" button
- block.Append(wire.NewTLVBE(0x000b, "https://github.com/mk6i/open-oscar-server"))
- if err := flapc.NewSignoff(block); err != nil {
- return fmt.Errorf("unable to gracefully disconnect user. %w", err)
- }
- }
- return nil
- case <-ctx.Done():
- if instance.MultiConnFlag() == wire.MultiConnFlagsOldClient {
- if err := flapc.OldSignoff(); err != nil {
- return fmt.Errorf("unable to gracefully disconnect user. %w", err)
- }
- } else {
- if err := flapc.NewSignoff(wire.TLVRestBlock{}); err != nil {
- return fmt.Errorf("unable to gracefully disconnect user. %w", err)
- }
- }
- return nil
- case err := <-errCh:
- if !errors.Is(err, io.EOF) {
- s.logger.ErrorContext(ctx, "client disconnected with error", "err", err)
- }
- return nil
- }
- }
- }
- // IPRateLimiter enforces a per-IP rate limit using a token bucket algorithm.
- // It caches individual rate limiters by IP address and supports tagging requests
- // as originating from the BUCP or FLAP auth.
- //
- // The limiter uses an in-memory cache with TTL expiration, so rate limits reset
- // after the TTL if no activity is observed for a given IP.
- type IPRateLimiter struct {
- cache *cache.Cache // In-memory cache mapping IPs to rate limiters with optional BUCP tag
- rate rate.Limit // Requests allowed per second
- burst int // Maximum burst size allowed
- }
- type rateLimitEntry struct {
- isBUCP bool
- limiter *rate.Limiter
- }
- // NewIPRateLimiter initializes a new IPRateLimiter with the specified rate,
- // burst size, and TTL for each IP's limiter. Entries expire after 2×TTL.
- func NewIPRateLimiter(rate rate.Limit, burst int, ttl time.Duration) *IPRateLimiter {
- return &IPRateLimiter{
- cache: cache.New(ttl, 2*ttl),
- rate: rate,
- burst: burst,
- }
- }
- // SetBUCP marks the rate limiter for the given IP as originating from BUCP auth
- // (default FLAP auth).
- func (l *IPRateLimiter) SetBUCP(ip string) {
- limiter, found := l.cache.Get(ip)
- if !found {
- limiter = &rateLimitEntry{
- isBUCP: true,
- limiter: rate.NewLimiter(l.rate, l.burst),
- }
- l.cache.Set(ip, limiter, cache.DefaultExpiration)
- }
- limiter.(*rateLimitEntry).isBUCP = true
- }
- // Allow checks if a request from the given IP is allowed under its rate limit.
- // It returns whether the request is allowed and whether the connection uses
- // BUCP auth.
- func (l *IPRateLimiter) Allow(ip string) (allowed bool, isBUCP bool) {
- limiter, found := l.cache.Get(ip)
- if !found {
- limiter = &rateLimitEntry{
- limiter: rate.NewLimiter(l.rate, l.burst),
- }
- l.cache.Set(ip, limiter, cache.DefaultExpiration)
- }
- entry := limiter.(*rateLimitEntry)
- return entry.limiter.Allow(), entry.isBUCP
- }
|