Browse Source

refactor dispatchIncomingMessages

Mike 2 years ago
parent
commit
c84494e910
7 changed files with 404 additions and 176 deletions
  1. 1 1
      server/chat.go
  2. 78 110
      server/connection.go
  3. 182 0
      server/connection_test.go
  4. 22 4
      server/logging.go
  5. 88 31
      server/router.go
  6. 21 18
      server/session.go
  7. 12 12
      server/session_test.go

+ 1 - 1
server/chat.go

@@ -91,7 +91,7 @@ func (s ChatService) ChannelMsgToHostHandler(ctx context.Context, sess *Session,
 	return ret, nil
 	return ret, nil
 }
 }
 
 
-func SetOnlineChatUsers(ctx context.Context, sess *Session, sm SessionManager) {
+func SetOnlineChatUsers(ctx context.Context, sess *Session, sm ChatRoom) {
 	snacPayloadOut := oscar.SNAC_0x0E_0x03_ChatUsersJoined{}
 	snacPayloadOut := oscar.SNAC_0x0E_0x03_ChatUsersJoined{}
 	sessions := sm.Participants()
 	sessions := sm.Participants()
 
 

+ 78 - 110
server/connection.go

@@ -4,7 +4,6 @@ import (
 	"bytes"
 	"bytes"
 	"context"
 	"context"
 	"errors"
 	"errors"
-	"fmt"
 	"io"
 	"io"
 	"log"
 	"log"
 	"log/slog"
 	"log/slog"
@@ -16,120 +15,99 @@ import (
 )
 )
 
 
 var (
 var (
-	CapChat, _ = uuid.MustParse("748F2420-6287-11D1-8222-444553540000").MarshalBinary()
+	CapChat, _             = uuid.MustParse("748F2420-6287-11D1-8222-444553540000").MarshalBinary()
+	ErrUnsupportedSubGroup = errors.New("unimplemented subgroup, your client version may be unsupported")
 )
 )
 
 
-var (
-	ErrUnsupportedFoodGroup = errors.New("unimplemented food group, your client version may be unsupported")
-	ErrUnsupportedSubGroup  = errors.New("unimplemented subgroup, your client version may be unsupported")
+type (
+	XMessage struct {
+		snacFrame oscar.SnacFrame
+		snacOut   any
+	}
+	incomingMessage struct {
+		flap    oscar.FlapFrame
+		payload *bytes.Buffer
+	}
+	alertHandler     func(ctx context.Context, msg XMessage, w io.Writer, u *uint32) error
+	clientReqHandler func(ctx context.Context, r io.Reader, w io.Writer, u *uint32) error
 )
 )
 
 
-type IncomingMessage struct {
-	flap oscar.FlapFrame
-	snac oscar.SnacFrame
-	buf  io.Reader
-}
-
-type XMessage struct {
-	snacFrame oscar.SnacFrame
-	snacOut   any
-}
-
-func readIncomingRequests(ctx context.Context, logger *slog.Logger, rw io.Reader, msgCh chan IncomingMessage, errCh chan error) {
+func consumeFLAPFrames(r io.Reader, msgCh chan incomingMessage, errCh chan error) {
 	defer close(msgCh)
 	defer close(msgCh)
 	defer close(errCh)
 	defer close(errCh)
 
 
 	for {
 	for {
-		flap := oscar.FlapFrame{}
-		if err := oscar.Unmarshal(&flap, rw); err != nil {
+		in := incomingMessage{}
+		if err := oscar.Unmarshal(&in.flap, r); err != nil {
 			errCh <- err
 			errCh <- err
 			return
 			return
 		}
 		}
 
 
-		switch flap.FrameType {
-		case oscar.FlapFrameSignon:
-			errCh <- errors.New("shouldn't get FlapFrameSignon")
-			return
-		case oscar.FlapFrameData:
-			b := make([]byte, flap.PayloadLength)
-			if _, err := rw.Read(b); err != nil {
+		if in.flap.FrameType == oscar.FlapFrameData {
+			buf := make([]byte, in.flap.PayloadLength)
+			if _, err := r.Read(buf); err != nil {
 				errCh <- err
 				errCh <- err
 				return
 				return
 			}
 			}
-
-			snac := oscar.SnacFrame{}
-			buf := bytes.NewBuffer(b)
-			if err := oscar.Unmarshal(&snac, buf); err != nil {
-				errCh <- err
-				return
-			}
-
-			msgCh <- IncomingMessage{
-				flap: flap,
-				snac: snac,
-				buf:  buf,
-			}
-		case oscar.FlapFrameError:
-			errCh <- fmt.Errorf("got FlapFrameError: %v", flap)
-			return
-		case oscar.FlapFrameSignoff:
-			errCh <- ErrSignedOff
-			return
-		case oscar.FlapFrameKeepAlive:
-			logger.DebugContext(ctx, "keepalive heartbeat")
-		default:
-			errCh <- fmt.Errorf("unknown frame type: %v", flap)
-			return
+			in.payload = bytes.NewBuffer(buf)
 		}
 		}
-	}
-}
 
 
-func Signout(ctx context.Context, logger *slog.Logger, sess *Session, sm SessionManager, fm *FeedbagStore) {
-	if err := BroadcastDeparture(ctx, sess, sm, fm); err != nil {
-		logger.ErrorContext(ctx, "error notifying departure", "err", err.Error())
+		msgCh <- in
 	}
 	}
-	sm.Remove(sess)
 }
 }
 
 
-type RouteSig func(ctx context.Context, rwc io.ReadWriter, u *uint32, snac oscar.SnacFrame, buf io.Reader) error
-
-func ReadBos(ctx context.Context, cfg Config, sess *Session, seq uint32, rwc io.ReadWriter, logger *slog.Logger, fn RouteSig) {
+func dispatchIncomingMessages(ctx context.Context, sess *Session, seq uint32, rw io.ReadWriter, logger *slog.Logger, fn clientReqHandler, alertHandler alertHandler) {
 	// buffered so that the go routine has room to exit
 	// buffered so that the go routine has room to exit
-	msgCh := make(chan IncomingMessage, 1)
-	errCh := make(chan error, 1)
-	go readIncomingRequests(ctx, logger, rwc, msgCh, errCh)
-
-	rl := RouteLogger{
-		Logger: logger,
-	}
+	msgCh := make(chan incomingMessage, 1)
+	readErrCh := make(chan error, 1)
+	go consumeFLAPFrames(rw, msgCh, readErrCh)
 
 
 	for {
 	for {
 		select {
 		select {
 		case m := <-msgCh:
 		case m := <-msgCh:
-			if err := fn(ctx, rwc, &seq, m.snac, m.buf); err != nil {
-				if errors.Is(err, ErrUnsupportedSubGroup) || errors.Is(err, ErrUnsupportedFoodGroup) {
-					if err1 := sendInvalidSNACErr(m.snac, rwc, &seq); err1 != nil {
-						err = errors.Join(err1, err)
-					}
-					if cfg.FailFast {
-						panic(err.Error())
-					}
+			switch m.flap.FrameType {
+			case oscar.FlapFrameData:
+				// route a client request to the appropriate service handler. the
+				// handler may write a response to the client connection.
+				if err := fn(ctx, m.payload, rw, &seq); err != nil {
+					return
 				}
 				}
-				logRequestError(ctx, logger, m.snac, err)
+			case oscar.FlapFrameSignon:
+				logger.ErrorContext(ctx, "shouldn't get FlapFrameSignon", "flap", m.flap)
+			case oscar.FlapFrameError:
+				logger.ErrorContext(ctx, "got FlapFrameError", "flap", m.flap)
+				return
+			case oscar.FlapFrameSignoff:
+				logger.InfoContext(ctx, "got FlapFrameSignoff", "flap", m.flap)
+				return
+			case oscar.FlapFrameKeepAlive:
+				logger.DebugContext(ctx, "keepalive heartbeat")
+			default:
+				logger.ErrorContext(ctx, "got unknown FLAP frame type", "flap", m.flap)
 				return
 				return
 			}
 			}
 		case m := <-sess.RecvMessage():
 		case m := <-sess.RecvMessage():
-			if err := writeOutSNAC(oscar.SnacFrame{}, m.snacFrame, m.snacOut, &seq, rwc); err != nil {
+			// forward a notification sent from another client to this client
+			if err := alertHandler(ctx, m, rw, &seq); err != nil {
 				logRequestError(ctx, logger, m.snacFrame, err)
 				logRequestError(ctx, logger, m.snacFrame, err)
 				return
 				return
 			}
 			}
-			rl.logRequest(ctx, m.snacFrame, m.snacOut)
+			logRequest(ctx, logger, m.snacFrame, m.snacOut)
 		case <-sess.Closed():
 		case <-sess.Closed():
-			if err := gracefulDisconnect(seq, rwc); err != nil {
+			// gracefully disconnect so that the client does not try to
+			// reconnect when the connection closes.
+			flap := oscar.FlapFrame{
+				StartMarker:   42,
+				FrameType:     oscar.FlapFrameSignoff,
+				Sequence:      uint16(seq),
+				PayloadLength: uint16(0),
+			}
+			if err := oscar.Marshal(flap, rw); err != nil {
 				logger.ErrorContext(ctx, "unable to gracefully disconnect user", "err", err)
 				logger.ErrorContext(ctx, "unable to gracefully disconnect user", "err", err)
 			}
 			}
 			return
 			return
-		case err := <-errCh:
+		case err := <-readErrCh:
+			// handle a read error
 			switch {
 			switch {
 			case errors.Is(io.EOF, err):
 			case errors.Is(io.EOF, err):
 				fallthrough
 				fallthrough
@@ -143,26 +121,8 @@ func ReadBos(ctx context.Context, cfg Config, sess *Session, seq uint32, rwc io.
 	}
 	}
 }
 }
 
 
-func logRequestError(ctx context.Context, logger *slog.Logger, inFrame oscar.SnacFrame, err error) {
-	logger.LogAttrs(ctx, slog.LevelError, "client disconnected with error",
-		slog.Group("request",
-			slog.String("food_group", oscar.FoodGroupStr(inFrame.FoodGroup)),
-			slog.String("sub_group", oscar.SubGroupStr(inFrame.FoodGroup, inFrame.SubGroup)),
-		),
-		slog.String("err", err.Error()),
-	)
-}
-
-func gracefulDisconnect(seq uint32, rwc io.ReadWriter) error {
-	return oscar.Marshal(oscar.FlapFrame{
-		StartMarker: 42,
-		FrameType:   oscar.FlapFrameSignoff,
-		Sequence:    uint16(seq),
-	}, rwc)
-}
-
-func HandleChatConnection(ctx context.Context, cfg Config, cr *ChatRegistry, conn net.Conn, router ChatServiceRouter, logger *slog.Logger) {
-	cookie, seq, err := VerifyChatLogin(conn)
+func HandleChatConnection(ctx context.Context, cr *ChatRegistry, rw io.ReadWriter, router ChatServiceRouter, logger *slog.Logger) {
+	cookie, seq, err := VerifyChatLogin(rw)
 	if err != nil {
 	if err != nil {
 		logger.ErrorContext(ctx, "user disconnected with error", "err", err.Error())
 		logger.ErrorContext(ctx, "user disconnected with error", "err", err.Error())
 		return
 		return
@@ -186,19 +146,21 @@ func HandleChatConnection(ctx context.Context, cfg Config, cr *ChatRegistry, con
 		AlertUserLeft(ctx, chatSess, room)
 		AlertUserLeft(ctx, chatSess, room)
 		room.Remove(chatSess)
 		room.Remove(chatSess)
 		cr.MaybeRemoveRoom(room.Cookie)
 		cr.MaybeRemoveRoom(room.Cookie)
-		conn.Close()
 	}()
 	}()
 
 
 	ctx = context.WithValue(ctx, "screenName", chatSess.ScreenName)
 	ctx = context.WithValue(ctx, "screenName", chatSess.ScreenName)
 
 
-	if err := router.WriteOServiceHostOnline(conn, &seq); err != nil {
+	if err := router.WriteOServiceHostOnline(rw, &seq); err != nil {
 		logger.ErrorContext(ctx, "error WriteOServiceHostOnline")
 		logger.ErrorContext(ctx, "error WriteOServiceHostOnline")
 	}
 	}
 
 
-	fn := func(ctx context.Context, rwc io.ReadWriter, u *uint32, snac oscar.SnacFrame, buf io.Reader) error {
-		return router.Route(ctx, chatSess, rwc, u, snac, buf, room)
+	fnClientReqHandler := func(ctx context.Context, r io.Reader, w io.Writer, seq *uint32) error {
+		return router.Route(ctx, chatSess, r, w, seq, room)
 	}
 	}
-	ReadBos(ctx, cfg, chatSess, seq, conn, logger, fn)
+	fnAlertHandler := func(ctx context.Context, msg XMessage, w io.Writer, seq *uint32) error {
+		return writeOutSNAC(oscar.SnacFrame{}, msg.snacFrame, msg.snacOut, seq, w)
+	}
+	dispatchIncomingMessages(ctx, chatSess, seq, rw, logger, fnClientReqHandler, fnAlertHandler)
 }
 }
 
 
 func HandleAuthConnection(cfg Config, sm *InMemorySessionManager, fm *FeedbagStore, conn net.Conn) {
 func HandleAuthConnection(cfg Config, sm *InMemorySessionManager, fm *FeedbagStore, conn net.Conn) {
@@ -223,7 +185,7 @@ func HandleAuthConnection(cfg Config, sm *InMemorySessionManager, fm *FeedbagSto
 	}
 	}
 }
 }
 
 
-func HandleBOSConnection(ctx context.Context, cfg Config, conn net.Conn, router BOSServiceRouter, logger *slog.Logger) {
+func HandleBOSConnection(ctx context.Context, conn net.Conn, router BOSServiceRouter, logger *slog.Logger) {
 	sess, seq, err := router.VerifyLogin(conn)
 	sess, seq, err := router.VerifyLogin(conn)
 	if err != nil {
 	if err != nil {
 		logger.ErrorContext(ctx, "user disconnected with error", "err", err.Error())
 		logger.ErrorContext(ctx, "user disconnected with error", "err", err.Error())
@@ -244,10 +206,13 @@ func HandleBOSConnection(ctx context.Context, cfg Config, conn net.Conn, router
 		logger.ErrorContext(ctx, "error WriteOServiceHostOnline")
 		logger.ErrorContext(ctx, "error WriteOServiceHostOnline")
 	}
 	}
 
 
-	fn := func(ctx context.Context, rwc io.ReadWriter, u *uint32, snac oscar.SnacFrame, buf io.Reader) error {
-		return router.Route(ctx, sess, rwc, u, snac, buf)
+	fnClientReqHandler := func(ctx context.Context, r io.Reader, w io.Writer, seq *uint32) error {
+		return router.Route(ctx, sess, r, w, seq)
+	}
+	fnAlertHandler := func(ctx context.Context, msg XMessage, w io.Writer, seq *uint32) error {
+		return writeOutSNAC(oscar.SnacFrame{}, msg.snacFrame, msg.snacOut, seq, w)
 	}
 	}
-	ReadBos(ctx, cfg, sess, seq, conn, logger, fn)
+	dispatchIncomingMessages(ctx, sess, seq, conn, logger, fnClientReqHandler, fnAlertHandler)
 }
 }
 
 
 func ListenChat(cfg Config, router ChatServiceRouter, cr *ChatRegistry, logger *slog.Logger) {
 func ListenChat(cfg Config, router ChatServiceRouter, cr *ChatRegistry, logger *slog.Logger) {
@@ -270,7 +235,10 @@ func ListenChat(cfg Config, router ChatServiceRouter, cr *ChatRegistry, logger *
 		ctx := context.Background()
 		ctx := context.Background()
 		ctx = context.WithValue(ctx, "ip", conn.RemoteAddr().String())
 		ctx = context.WithValue(ctx, "ip", conn.RemoteAddr().String())
 		logger.DebugContext(ctx, "accepted connection")
 		logger.DebugContext(ctx, "accepted connection")
-		go HandleChatConnection(ctx, cfg, cr, conn, router, logger)
+		go func() {
+			HandleChatConnection(ctx, cr, conn, router, logger)
+			conn.Close()
+		}()
 	}
 	}
 }
 }
 
 
@@ -294,7 +262,7 @@ func ListenBOS(cfg Config, router BOSServiceRouter, logger *slog.Logger) {
 		ctx := context.Background()
 		ctx := context.Background()
 		ctx = context.WithValue(ctx, "ip", conn.RemoteAddr().String())
 		ctx = context.WithValue(ctx, "ip", conn.RemoteAddr().String())
 		logger.DebugContext(ctx, "accepted connection")
 		logger.DebugContext(ctx, "accepted connection")
-		go HandleBOSConnection(ctx, cfg, conn, router, logger)
+		go HandleBOSConnection(ctx, conn, router, logger)
 	}
 	}
 }
 }
 
 

+ 182 - 0
server/connection_test.go

@@ -0,0 +1,182 @@
+package server
+
+import (
+	"bufio"
+	"bytes"
+	"context"
+	"github.com/mkaminski/goaim/oscar"
+	"github.com/stretchr/testify/assert"
+	"io"
+	"sync"
+	"testing"
+)
+
+func TestHandleChatConnection_Notification(t *testing.T) {
+
+	ctx := context.Background()
+	cfg := Config{}
+	cr := NewChatRegistry()
+	logger := NewLogger(cfg)
+
+	room := ChatRoom{
+		Name:           "test chat room!",
+		SessionManager: NewSessionManager(logger),
+	}
+	bobSess := room.NewSessionWithSN("bob-sess-id", "bob")
+	cr.Register(room)
+
+	msgIn := []XMessage{
+		{
+			snacFrame: oscar.SnacFrame{
+				FoodGroup: oscar.CHAT,
+				SubGroup:  oscar.ChatUsersJoined,
+			},
+			snacOut: oscar.SNAC_0x0E_0x03_ChatUsersJoined{
+				Users: []oscar.TLVUserInfo{
+					bobSess.GetTLVUserInfo(),
+				},
+			},
+		},
+		{
+			snacFrame: oscar.SnacFrame{
+				FoodGroup: oscar.CHAT,
+				SubGroup:  oscar.ChatUsersLeft,
+			},
+			snacOut: oscar.SNAC_0x0E_0x03_ChatUsersJoined{
+				Users: []oscar.TLVUserInfo{},
+			},
+		},
+	}
+
+	routeSig := func(ctx context.Context, buf io.Reader, w io.Writer, u *uint32) error {
+		return nil
+	}
+
+	wg := sync.WaitGroup{}
+	wg.Add(len(msgIn))
+
+	var msgOut []XMessage
+	alertHandler := func(ctx context.Context, msg XMessage, w io.Writer, u *uint32) error {
+		msgOut = append(msgOut, msg)
+		wg.Done()
+		return nil
+	}
+
+	go func() {
+		wg.Wait()
+		bobSess.Close()
+	}()
+
+	pr, _ := io.Pipe()
+	rw := bufio.NewReadWriter(bufio.NewReader(pr), bufio.NewWriter(&bytes.Buffer{}))
+
+	for _, msg := range msgIn {
+		room.SendToScreenName(ctx, "bob", msg)
+	}
+
+	dispatchIncomingMessages(ctx, bobSess, uint32(0), rw, logger, routeSig, alertHandler)
+
+	assert.Equal(t, msgIn, msgOut)
+}
+
+func TestHandleChatConnection_ClientRequestFLAP(t *testing.T) {
+
+	ctx := context.Background()
+	cfg := Config{}
+	cr := NewChatRegistry()
+	logger := NewLogger(cfg)
+
+	room := ChatRoom{
+		Name:           "test chat room!",
+		SessionManager: NewSessionManager(logger),
+	}
+	bobSess := room.NewSessionWithSN("bob-sess-id", "bob")
+	cr.Register(room)
+
+	payloads := [][]byte{
+		{'a', 'b', 'c', 'd'},
+		{'e', 'f', 'g', 'h'},
+	}
+
+	pr, pw := io.Pipe()
+	_, pw2 := io.Pipe()
+	go func() {
+		for _, buf := range payloads {
+			flap := oscar.FlapFrame{
+				StartMarker:   42,
+				FrameType:     oscar.FlapFrameData,
+				PayloadLength: uint16(len(buf)),
+			}
+			assert.NoError(t, oscar.Marshal(flap, pw))
+			assert.NoError(t, oscar.Marshal(buf, pw))
+		}
+	}()
+
+	var msgOut [][]byte
+	wg := sync.WaitGroup{}
+	wg.Add(len(payloads))
+
+	routeSig := func(ctx context.Context, buf io.Reader, w io.Writer, u *uint32) error {
+		var err error
+		b, err := io.ReadAll(buf)
+		msgOut = append(msgOut, b)
+		wg.Done()
+		return err
+	}
+	alertHandler := func(ctx context.Context, msg XMessage, w io.Writer, u *uint32) error {
+		return nil
+	}
+
+	rw := bufio.NewReadWriter(bufio.NewReader(pr), bufio.NewWriter(pw2))
+
+	go func() {
+		wg.Wait()
+		pw.Close()
+	}()
+
+	dispatchIncomingMessages(ctx, bobSess, uint32(0), rw, logger, routeSig, alertHandler)
+
+	assert.Equal(t, payloads, msgOut)
+}
+
+func TestHandleChatConnection_SessionClosed(t *testing.T) {
+
+	ctx := context.Background()
+	cfg := Config{}
+	cr := NewChatRegistry()
+	logger := NewLogger(cfg)
+
+	room := ChatRoom{
+		Name:           "test chat room!",
+		SessionManager: NewSessionManager(logger),
+	}
+	sess := room.NewSessionWithSN("bob-sess-id", "bob")
+	cr.Register(room)
+
+	routeSig := func(ctx context.Context, buf io.Reader, w io.Writer, u *uint32) error {
+		t.Fatal("not expecting any output")
+		return nil
+	}
+	alertHandler := func(ctx context.Context, msg XMessage, w io.Writer, u *uint32) error {
+		t.Fatal("not expecting any alerts")
+		return nil
+	}
+
+	pr1, _ := io.Pipe()
+	pr2, pw2 := io.Pipe()
+
+	in := struct {
+		io.Reader
+		io.Writer
+	}{
+		Reader: pr1,
+		Writer: pw2,
+	}
+	sess.Close()
+
+	go dispatchIncomingMessages(ctx, sess, 0, in, logger, routeSig, alertHandler)
+
+	flap := oscar.FlapFrame{}
+	assert.NoError(t, oscar.Unmarshal(&flap, pr2))
+	assert.Equal(t, oscar.FlapFrameSignoff, flap.FrameType)
+}

+ 22 - 4
server/logging.go

@@ -90,13 +90,31 @@ func (rt RouteLogger) logRequestAndResponse(ctx context.Context, inFrame oscar.S
 	}
 	}
 }
 }
 
 
+func (rt RouteLogger) logRequestError(ctx context.Context, inFrame oscar.SnacFrame, err error) {
+	logRequestError(ctx, rt.Logger, inFrame, err)
+}
+
+func logRequestError(ctx context.Context, logger *slog.Logger, inFrame oscar.SnacFrame, err error) {
+	logger.LogAttrs(ctx, slog.LevelError, "client request error",
+		slog.Group("request",
+			slog.String("food_group", oscar.FoodGroupStr(inFrame.FoodGroup)),
+			slog.String("sub_group", oscar.SubGroupStr(inFrame.FoodGroup, inFrame.SubGroup)),
+		),
+		slog.String("err", err.Error()),
+	)
+}
+
 func (rt RouteLogger) logRequest(ctx context.Context, inFrame oscar.SnacFrame, inSNAC any) {
 func (rt RouteLogger) logRequest(ctx context.Context, inFrame oscar.SnacFrame, inSNAC any) {
+	logRequest(ctx, rt.Logger, inFrame, inSNAC)
+}
+
+func logRequest(ctx context.Context, logger *slog.Logger, inFrame oscar.SnacFrame, inSNAC any) {
 	const msg = "client request"
 	const msg = "client request"
 	switch {
 	switch {
-	case rt.Logger.Enabled(ctx, LevelTrace):
-		rt.Logger.LogAttrs(ctx, LevelTrace, msg, SNACLogGroupWithPayload("request", inFrame, inSNAC))
-	case rt.Logger.Enabled(ctx, slog.LevelDebug):
-		rt.Logger.LogAttrs(ctx, slog.LevelDebug, msg, slog.Group("request", SNACLogGroup("request", inFrame)))
+	case logger.Enabled(ctx, LevelTrace):
+		logger.LogAttrs(ctx, LevelTrace, msg, SNACLogGroupWithPayload("request", inFrame, inSNAC))
+	case logger.Enabled(ctx, slog.LevelDebug):
+		logger.LogAttrs(ctx, slog.LevelDebug, msg, slog.Group("request", SNACLogGroup("request", inFrame)))
 	}
 	}
 }
 }
 
 

+ 88 - 31
server/router.go

@@ -5,10 +5,11 @@ import (
 	"context"
 	"context"
 	"errors"
 	"errors"
 	"fmt"
 	"fmt"
-	"github.com/mkaminski/goaim/oscar"
 	"io"
 	"io"
 	"log/slog"
 	"log/slog"
 	"net"
 	"net"
+
+	"github.com/mkaminski/goaim/oscar"
 )
 )
 
 
 func NewBOSServiceRouter(logger *slog.Logger, cfg Config, fm FeedbagManager, sm SessionManager, cr *ChatRegistry, pm ProfileManager) BOSServiceRouter {
 func NewBOSServiceRouter(logger *slog.Logger, cfg Config, fm FeedbagManager, sm SessionManager, cr *ChatRegistry, pm ProfileManager) BOSServiceRouter {
@@ -22,6 +23,10 @@ func NewBOSServiceRouter(logger *slog.Logger, cfg Config, fm FeedbagManager, sm
 		OServiceBOSRouter: NewOServiceRouterForBOS(logger, cfg, fm, sm, cr),
 		OServiceBOSRouter: NewOServiceRouterForBOS(logger, cfg, fm, sm, cr),
 		sm:                sm,
 		sm:                sm,
 		fm:                fm,
 		fm:                fm,
+		cfg:               cfg,
+		RouteLogger: RouteLogger{
+			Logger: logger,
+		},
 	}
 	}
 }
 }
 
 
@@ -29,6 +34,10 @@ func NewChatServiceRouter(logger *slog.Logger, cfg Config, fm FeedbagManager, sm
 	return ChatServiceRouter{
 	return ChatServiceRouter{
 		OServiceChatRouter: NewOServiceRouterForChat(logger, cfg, fm, sm),
 		OServiceChatRouter: NewOServiceRouterForChat(logger, cfg, fm, sm),
 		ChatRouter:         NewChatRouter(logger),
 		ChatRouter:         NewChatRouter(logger),
+		cfg:                cfg,
+		RouteLogger: RouteLogger{
+			Logger: logger,
+		},
 	}
 	}
 }
 }
 
 
@@ -40,31 +49,55 @@ type BOSServiceRouter struct {
 	ICBMRouter
 	ICBMRouter
 	LocateRouter
 	LocateRouter
 	OServiceBOSRouter
 	OServiceBOSRouter
-	sm SessionManager
-	fm FeedbagManager
+	sm  SessionManager
+	fm  FeedbagManager
+	cfg Config
+	RouteLogger
 }
 }
 
 
-func (rt *BOSServiceRouter) Route(ctx context.Context, sess *Session, w io.Writer, sequence *uint32, snac oscar.SnacFrame, buf io.Reader) error {
-	switch snac.FoodGroup {
-	case oscar.OSERVICE:
-		return rt.RouteOService(ctx, sess, snac, buf, w, sequence)
-	case oscar.LOCATE:
-		return rt.RouteLocate(ctx, sess, snac, buf, w, sequence)
-	case oscar.BUDDY:
-		return rt.RouteBuddy(ctx, snac, buf, w, sequence)
-	case oscar.ICBM:
-		return rt.RouteICBM(ctx, sess, snac, buf, w, sequence)
-	case oscar.CHAT_NAV:
-		return rt.RouteChatNav(ctx, sess, snac, buf, w, sequence)
-	case oscar.FEEDBAG:
-		return rt.RouteFeedbag(ctx, sess, snac, buf, w, sequence)
-	case oscar.BUCP:
-		return routeBUCP(ctx)
-	case oscar.ALERT:
-		return rt.RouteAlert(ctx, snac)
-	default:
-		return ErrUnsupportedFoodGroup
+func (rt *BOSServiceRouter) Route(ctx context.Context, sess *Session, r io.Reader, w io.Writer, sequence *uint32) error {
+	snac := oscar.SnacFrame{}
+	if err := oscar.Unmarshal(&snac, r); err != nil {
+		return err
 	}
 	}
+
+	err := func() error {
+		switch snac.FoodGroup {
+		case oscar.OSERVICE:
+			return rt.RouteOService(ctx, sess, snac, r, w, sequence)
+		case oscar.LOCATE:
+			return rt.RouteLocate(ctx, sess, snac, r, w, sequence)
+		case oscar.BUDDY:
+			return rt.RouteBuddy(ctx, snac, r, w, sequence)
+		case oscar.ICBM:
+			return rt.RouteICBM(ctx, sess, snac, r, w, sequence)
+		case oscar.CHAT_NAV:
+			return rt.RouteChatNav(ctx, sess, snac, r, w, sequence)
+		case oscar.FEEDBAG:
+			return rt.RouteFeedbag(ctx, sess, snac, r, w, sequence)
+		case oscar.BUCP:
+			return routeBUCP(ctx)
+		case oscar.ALERT:
+			return rt.RouteAlert(ctx, snac)
+		default:
+			return ErrUnsupportedSubGroup
+		}
+	}()
+
+	if err != nil {
+		rt.logRequestError(ctx, snac, err)
+		if errors.Is(err, ErrUnsupportedSubGroup) {
+			if err1 := sendInvalidSNACErr(snac, w, sequence); err1 != nil {
+				err = errors.Join(err1, err)
+			}
+			if rt.cfg.FailFast {
+				panic(err.Error())
+			}
+			return nil
+		}
+	}
+
+	return err
 }
 }
 
 
 func (rt *BOSServiceRouter) Signout(ctx context.Context, logger *slog.Logger, sess *Session) {
 func (rt *BOSServiceRouter) Signout(ctx context.Context, logger *slog.Logger, sess *Session) {
@@ -147,15 +180,39 @@ func sendInvalidSNACErr(snac oscar.SnacFrame, w io.Writer, sequence *uint32) err
 type ChatServiceRouter struct {
 type ChatServiceRouter struct {
 	ChatRouter
 	ChatRouter
 	OServiceChatRouter
 	OServiceChatRouter
+	cfg Config
+	RouteLogger
 }
 }
 
 
-func (rt *ChatServiceRouter) Route(ctx context.Context, sess *Session, w io.Writer, sequence *uint32, snac oscar.SnacFrame, buf io.Reader, room ChatRoom) error {
-	switch snac.FoodGroup {
-	case oscar.OSERVICE:
-		return rt.RouteOService(ctx, sess, room, snac, buf, w, sequence)
-	case oscar.CHAT:
-		return rt.RouteChat(ctx, sess, room, snac, buf, w, sequence)
-	default:
-		return ErrUnsupportedFoodGroup
+func (rt *ChatServiceRouter) Route(ctx context.Context, sess *Session, r io.Reader, w io.Writer, sequence *uint32, room ChatRoom) error {
+	snac := oscar.SnacFrame{}
+	if err := oscar.Unmarshal(&snac, r); err != nil {
+		return err
 	}
 	}
+
+	err := func() error {
+		switch snac.FoodGroup {
+		case oscar.OSERVICE:
+			return rt.RouteOService(ctx, sess, room, snac, r, w, sequence)
+		case oscar.CHAT:
+			return rt.RouteChat(ctx, sess, room, snac, r, w, sequence)
+		default:
+			return ErrUnsupportedSubGroup
+		}
+	}()
+
+	if err != nil {
+		rt.logRequestError(ctx, snac, err)
+		if errors.Is(err, ErrUnsupportedSubGroup) {
+			if err1 := sendInvalidSNACErr(snac, w, sequence); err1 != nil {
+				err = errors.Join(err1, err)
+			}
+			if rt.cfg.FailFast {
+				panic(err.Error())
+			}
+			return nil
+		}
+	}
+
+	return err
 }
 }

+ 21 - 18
server/session.go

@@ -4,10 +4,11 @@ import (
 	"context"
 	"context"
 	"errors"
 	"errors"
 	"fmt"
 	"fmt"
-	"github.com/mkaminski/goaim/oscar"
 	"log/slog"
 	"log/slog"
 	"sync"
 	"sync"
 	"time"
 	"time"
+
+	"github.com/mkaminski/goaim/oscar"
 )
 )
 
 
 var (
 var (
@@ -22,12 +23,11 @@ const (
 	SessSendOK SessSendStatus = iota
 	SessSendOK SessSendStatus = iota
 	// SessSendClosed indicates send did not complete because session is closed
 	// SessSendClosed indicates send did not complete because session is closed
 	SessSendClosed
 	SessSendClosed
-	// SessSendTimeout indicates send timed out due to blocked recipient
-	SessSendTimeout
+	// SessQueueFull indicates send failed due to full queue -- client is likely
+	// dead
+	SessQueueFull
 )
 )
 
 
-const sendTimeout = 10 * time.Second
-
 type Session struct {
 type Session struct {
 	ID          string
 	ID          string
 	ScreenName  string
 	ScreenName  string
@@ -40,7 +40,6 @@ type Session struct {
 	invisible   bool
 	invisible   bool
 	idle        bool
 	idle        bool
 	idleTime    time.Time
 	idleTime    time.Time
-	sendTimeout time.Duration
 	closed      bool
 	closed      bool
 }
 }
 
 
@@ -156,13 +155,18 @@ func (s *Session) RecvMessage() chan XMessage {
 }
 }
 
 
 func (s *Session) SendMessage(msg XMessage) SessSendStatus {
 func (s *Session) SendMessage(msg XMessage) SessSendStatus {
+	s.Mutex.Lock()
+	if s.closed {
+		return SessSendClosed
+	}
+	s.Mutex.Unlock()
 	select {
 	select {
 	case s.msgCh <- msg:
 	case s.msgCh <- msg:
 		return SessSendOK
 		return SessSendOK
 	case <-s.stopCh:
 	case <-s.stopCh:
 		return SessSendClosed
 		return SessSendClosed
-	case <-time.After(s.sendTimeout):
-		return SessSendTimeout
+	default:
+		return SessQueueFull
 	}
 	}
 }
 }
 
 
@@ -197,7 +201,7 @@ func (s *InMemorySessionManager) Broadcast(ctx context.Context, msg XMessage) {
 	s.mapMutex.RLock()
 	s.mapMutex.RLock()
 	defer s.mapMutex.RUnlock()
 	defer s.mapMutex.RUnlock()
 	for _, sess := range s.store {
 	for _, sess := range s.store {
-		go s.maybeSendMessage(ctx, msg, sess)
+		s.maybeSendMessage(ctx, msg, sess)
 	}
 	}
 }
 }
 
 
@@ -205,8 +209,8 @@ func (s *InMemorySessionManager) maybeSendMessage(ctx context.Context, msg XMess
 	switch sess.SendMessage(msg) {
 	switch sess.SendMessage(msg) {
 	case SessSendClosed:
 	case SessSendClosed:
 		s.logger.WarnContext(ctx, "can't send notification because the user's session is closed", "recipient", sess.ScreenName, "message", msg)
 		s.logger.WarnContext(ctx, "can't send notification because the user's session is closed", "recipient", sess.ScreenName, "message", msg)
-	case SessSendTimeout:
-		s.logger.WarnContext(ctx, "can't send notification because of send timeout", "recipient", sess.ScreenName, "message", msg)
+	case SessQueueFull:
+		s.logger.WarnContext(ctx, "can't send notification because queue is full", "recipient", sess.ScreenName, "message", msg)
 		sess.Close()
 		sess.Close()
 	}
 	}
 }
 }
@@ -234,7 +238,7 @@ func (s *InMemorySessionManager) BroadcastExcept(ctx context.Context, except *Se
 		if sess == except {
 		if sess == except {
 			continue
 			continue
 		}
 		}
-		go s.maybeSendMessage(ctx, msg, sess)
+		s.maybeSendMessage(ctx, msg, sess)
 	}
 	}
 }
 }
 
 
@@ -276,21 +280,20 @@ func (s *InMemorySessionManager) SendToScreenName(ctx context.Context, screenNam
 		s.logger.WarnContext(ctx, "can't send notification because user is not online", "recipient", screenName, "message", msg)
 		s.logger.WarnContext(ctx, "can't send notification because user is not online", "recipient", screenName, "message", msg)
 		return
 		return
 	}
 	}
-	go s.maybeSendMessage(ctx, msg, sess)
+	s.maybeSendMessage(ctx, msg, sess)
 }
 }
 
 
 func (s *InMemorySessionManager) BroadcastToScreenNames(ctx context.Context, screenNames []string, msg XMessage) {
 func (s *InMemorySessionManager) BroadcastToScreenNames(ctx context.Context, screenNames []string, msg XMessage) {
 	for _, sess := range s.retrieveByScreenNames(screenNames) {
 	for _, sess := range s.retrieveByScreenNames(screenNames) {
-		go s.maybeSendMessage(ctx, msg, sess)
+		s.maybeSendMessage(ctx, msg, sess)
 	}
 	}
 }
 }
 
 
 func makeSession() *Session {
 func makeSession() *Session {
 	return &Session{
 	return &Session{
-		msgCh:       make(chan XMessage, 1),
-		stopCh:      make(chan struct{}),
-		sendTimeout: sendTimeout,
-		SignonTime:  time.Now(),
+		msgCh:      make(chan XMessage, 1000),
+		stopCh:     make(chan struct{}),
+		SignonTime: time.Now(),
 	}
 	}
 }
 }
 
 

+ 12 - 12
server/session_test.go

@@ -1,15 +1,15 @@
 package server
 package server
 
 
 import (
 import (
+	"github.com/stretchr/testify/assert"
 	"testing"
 	"testing"
 	"time"
 	"time"
 )
 )
 
 
 func TestSession_SendMessage_SessSendOK(t *testing.T) {
 func TestSession_SendMessage_SessSendOK(t *testing.T) {
 	s := Session{
 	s := Session{
-		msgCh:       make(chan XMessage, 1),
-		stopCh:      make(chan struct{}),
-		sendTimeout: sendTimeout,
+		msgCh:  make(chan XMessage, 1),
+		stopCh: make(chan struct{}),
 	}
 	}
 	if res := s.SendMessage(XMessage{}); res != SessSendOK {
 	if res := s.SendMessage(XMessage{}); res != SessSendOK {
 		t.Fatalf("expected SessSendOK, got %+v", res)
 		t.Fatalf("expected SessSendOK, got %+v", res)
@@ -18,9 +18,8 @@ func TestSession_SendMessage_SessSendOK(t *testing.T) {
 
 
 func TestSession_SendMessage_SessSendClosed(t *testing.T) {
 func TestSession_SendMessage_SessSendClosed(t *testing.T) {
 	s := Session{
 	s := Session{
-		msgCh:       make(chan XMessage, 1),
-		stopCh:      make(chan struct{}),
-		sendTimeout: sendTimeout,
+		msgCh:  make(chan XMessage, 1),
+		stopCh: make(chan struct{}),
 	}
 	}
 	s.Close()
 	s.Close()
 	if res := s.SendMessage(XMessage{}); res != SessSendClosed {
 	if res := s.SendMessage(XMessage{}); res != SessSendClosed {
@@ -28,15 +27,16 @@ func TestSession_SendMessage_SessSendClosed(t *testing.T) {
 	}
 	}
 }
 }
 
 
-func TestSession_SendMessage_SessSendTimeout(t *testing.T) {
+func TestSession_SendMessage_SessQueueFull(t *testing.T) {
+	bufSize := 10
 	s := Session{
 	s := Session{
-		msgCh:       make(chan XMessage),
-		stopCh:      make(chan struct{}),
-		sendTimeout: 0,
+		msgCh:  make(chan XMessage, bufSize),
+		stopCh: make(chan struct{}),
 	}
 	}
-	if res := s.SendMessage(XMessage{}); res != SessSendTimeout {
-		t.Fatalf("expected SessSendTimeout got %+v", res)
+	for i := 0; i < bufSize; i++ {
+		assert.Equal(t, SessSendOK, s.SendMessage(XMessage{}))
 	}
 	}
+	assert.Equal(t, SessQueueFull, s.SendMessage(XMessage{}))
 }
 }
 
 
 func TestSession_Close_Twice(t *testing.T) {
 func TestSession_Close_Twice(t *testing.T) {