admin.go 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224
  1. package oscar
  2. import (
  3. "bytes"
  4. "context"
  5. "errors"
  6. "fmt"
  7. "io"
  8. "log/slog"
  9. "net"
  10. "sync"
  11. "time"
  12. "github.com/mk6i/retro-aim-server/config"
  13. "github.com/mk6i/retro-aim-server/server/oscar/middleware"
  14. "github.com/mk6i/retro-aim-server/state"
  15. "github.com/mk6i/retro-aim-server/wire"
  16. )
  17. // AdminServer provides client connection lifecycle management for the BOS
  18. // service.
  19. type AdminServer struct {
  20. AuthService
  21. Handler
  22. ListenAddr string
  23. Logger *slog.Logger
  24. OnlineNotifier
  25. config.Config
  26. RateLimitUpdater
  27. wire.SNACRateLimits
  28. }
  29. // Start starts a TCP server and listens for connections. The initial
  30. // authentication handshake sequences are handled by this method. The remaining
  31. // requests are relayed to BOSRouter.
  32. func (rt AdminServer) Start(ctx context.Context) error {
  33. listener, err := net.Listen("tcp", rt.ListenAddr)
  34. if err != nil {
  35. return fmt.Errorf("unable to start admin server: %w", err)
  36. }
  37. go func() {
  38. <-ctx.Done()
  39. listener.Close()
  40. }()
  41. rt.Logger.Info("starting server", "listen_host", rt.ListenAddr, "oscar_host", rt.Config.OSCARHost)
  42. wg := sync.WaitGroup{}
  43. for {
  44. conn, err := listener.Accept()
  45. if err != nil {
  46. if errors.Is(err, net.ErrClosed) {
  47. break
  48. }
  49. rt.Logger.Error("accept failed", "err", err.Error())
  50. continue
  51. }
  52. wg.Add(1)
  53. go func() {
  54. defer wg.Done()
  55. connCtx := context.WithValue(ctx, "ip", conn.RemoteAddr().String())
  56. rt.Logger.DebugContext(connCtx, "accepted connection")
  57. if err := rt.handleNewConnection(connCtx, conn); err != nil {
  58. rt.Logger.Info("user session failed", "err", err.Error())
  59. }
  60. }()
  61. }
  62. if !waitForShutdown(&wg) {
  63. rt.Logger.Error("shutdown complete, but connections didn't close cleanly")
  64. } else {
  65. rt.Logger.Info("shutdown complete")
  66. }
  67. return nil
  68. }
  69. func (rt AdminServer) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser) error {
  70. flapc := wire.NewFlapClient(100, rwc, rwc)
  71. if err := flapc.SendSignonFrame(nil); err != nil {
  72. return err
  73. }
  74. flap, err := flapc.ReceiveSignonFrame()
  75. if err != nil {
  76. return err
  77. }
  78. authCookie, ok := flap.Bytes(wire.OServiceTLVTagsLoginCookie)
  79. if !ok {
  80. return errors.New("unable to get session id from payload")
  81. }
  82. sess, err := rt.RetrieveBOSSession(ctx, authCookie)
  83. if err != nil {
  84. return err
  85. }
  86. if sess == nil {
  87. return errors.New("session not found")
  88. }
  89. defer func() {
  90. rwc.Close()
  91. }()
  92. ctx = context.WithValue(ctx, "screenName", sess.IdentScreenName())
  93. msg := rt.OnlineNotifier.HostOnline()
  94. if err := flapc.SendSNAC(msg.Frame, msg.Body); err != nil {
  95. return err
  96. }
  97. return dispatchIncomingMessagesSimple(ctx, sess, flapc, rwc, rt.Logger, rt.Handler, rt.RateLimitUpdater, rt.SNACRateLimits)
  98. }
  99. func dispatchIncomingMessagesSimple(
  100. ctx context.Context,
  101. sess *state.Session,
  102. flapc *wire.FlapClient,
  103. r io.Reader,
  104. logger *slog.Logger,
  105. router Handler,
  106. rateLimitUpdater RateLimitUpdater,
  107. snacRateLimits wire.SNACRateLimits,
  108. ) error {
  109. defer func() {
  110. logger.InfoContext(ctx, "user disconnected")
  111. }()
  112. // buffered so that the go routine has room to exit
  113. msgCh := make(chan wire.FLAPFrame, 1)
  114. errCh := make(chan error, 1)
  115. // consume flap frames
  116. go func() {
  117. defer close(msgCh)
  118. defer close(errCh)
  119. for {
  120. frame := wire.FLAPFrame{}
  121. if err := wire.UnmarshalBE(&frame, r); err != nil {
  122. errCh <- err
  123. return
  124. }
  125. msgCh <- frame
  126. }
  127. }()
  128. for {
  129. select {
  130. case flap, ok := <-msgCh:
  131. if !ok {
  132. return nil
  133. }
  134. switch flap.FrameType {
  135. case wire.FLAPFrameData:
  136. flapBuf := bytes.NewBuffer(flap.Payload)
  137. inFrame := wire.SNACFrame{}
  138. if err := wire.UnmarshalBE(&inFrame, flapBuf); err != nil {
  139. return err
  140. }
  141. rateClassID, ok := snacRateLimits.RateClassLookup(inFrame.FoodGroup, inFrame.SubGroup)
  142. if ok {
  143. if status := sess.EvaluateRateLimit(time.Now(), rateClassID); status == wire.RateLimitStatusLimited {
  144. logger.DebugContext(ctx, "rate limit exceeded, dropping SNAC",
  145. "foodgroup", wire.FoodGroupName(inFrame.FoodGroup),
  146. "subgroup", wire.SubGroupName(inFrame.FoodGroup, inFrame.SubGroup))
  147. break
  148. }
  149. } else {
  150. logger.ErrorContext(ctx, "rate limit not found, allowing request through")
  151. }
  152. // route a client request to the appropriate service handler. the
  153. // handler may write a response to the client connection.
  154. if err := router.Handle(ctx, sess, inFrame, flapBuf, flapc); err != nil {
  155. middleware.LogRequestError(ctx, logger, inFrame, err)
  156. if errors.Is(err, ErrRouteNotFound) {
  157. if err1 := sendInvalidSNACErr(inFrame, flapc); err1 != nil {
  158. return errors.Join(err1, err)
  159. }
  160. break
  161. }
  162. return err
  163. }
  164. case wire.FLAPFrameSignon:
  165. return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", flap)
  166. case wire.FLAPFrameError:
  167. return fmt.Errorf("got FLAPFrameError. flap: %v", flap)
  168. case wire.FLAPFrameSignoff:
  169. logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", flap)
  170. return nil
  171. case wire.FLAPFrameKeepAlive:
  172. logger.DebugContext(ctx, "keepalive heartbeat")
  173. default:
  174. return fmt.Errorf("got unknown FLAP frame type. flap: %v", flap)
  175. }
  176. case <-time.After(1 * time.Second):
  177. msgs := rateLimitUpdater.RateLimitUpdates(ctx, sess, time.Now())
  178. for _, rate := range msgs {
  179. if err := flapc.SendSNAC(rate.Frame, rate.Body); err != nil {
  180. middleware.LogRequestError(ctx, logger, rate.Frame, err)
  181. return err
  182. }
  183. }
  184. case <-ctx.Done():
  185. // application is shutting down
  186. if err := flapc.Disconnect(); err != nil {
  187. return fmt.Errorf("unable to gracefully disconnect user. %w", err)
  188. }
  189. return nil
  190. case err := <-errCh:
  191. if !errors.Is(io.EOF, err) {
  192. logger.ErrorContext(ctx, "client disconnected with error", "err", err)
  193. }
  194. return nil
  195. }
  196. }
  197. }