Bladeren bron

use len_prefix for FLAP frame buffer

Mike 2 jaren geleden
bovenliggende
commit
495f33e166

+ 31 - 14
server/oscar/admin.go

@@ -1,6 +1,7 @@
 package oscar
 
 import (
+	"bytes"
 	"context"
 	"errors"
 	"fmt"
@@ -96,30 +97,46 @@ func (rt AdminServer) handleNewConnection(ctx context.Context, rwc io.ReadWriteC
 }
 
 func dispatchIncomingMessagesSimple(ctx context.Context, sess *state.Session, flapc *wire.FlapClient, r io.Reader, logger *slog.Logger, router Handler, config config.Config) error {
-	// buffered so that the go routine has room to exit
-	msgCh := make(chan incomingMessage, 1)
-	readErrCh := make(chan error, 1)
-	go consumeFLAPFrames(r, msgCh, readErrCh)
-
 	defer func() {
 		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.Unmarshal(&frame, r); err != nil {
+				errCh <- err
+				return
+			}
+			msgCh <- frame
+		}
+	}()
+
 	for {
 		select {
-		case m, ok := <-msgCh:
+		case flap, ok := <-msgCh:
 			if !ok {
 				return nil
 			}
-			switch m.flap.FrameType {
+			switch flap.FrameType {
 			case wire.FLAPFrameData:
+				flapBuf := bytes.NewBuffer(flap.Payload)
+
 				inFrame := wire.SNACFrame{}
-				if err := wire.Unmarshal(&inFrame, m.payload); err != nil {
+				if err := wire.Unmarshal(&inFrame, flapBuf); err != nil {
 					return err
 				}
 				// route a client request to the appropriate service handler. the
 				// handler may write a response to the client connection.
-				if err := router.Handle(ctx, sess, inFrame, m.payload, flapc); err != nil {
+				if err := router.Handle(ctx, sess, inFrame, flapBuf, flapc); err != nil {
 					middleware.LogRequestError(ctx, logger, inFrame, err)
 					if errors.Is(err, ErrRouteNotFound) {
 						if err1 := sendInvalidSNACErr(inFrame, flapc); err1 != nil {
@@ -133,18 +150,18 @@ func dispatchIncomingMessagesSimple(ctx context.Context, sess *state.Session, fl
 					return err
 				}
 			case wire.FLAPFrameSignon:
-				return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", m.flap)
+				return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", flap)
 			case wire.FLAPFrameError:
-				return fmt.Errorf("got FLAPFrameError. flap: %v", m.flap)
+				return fmt.Errorf("got FLAPFrameError. flap: %v", flap)
 			case wire.FLAPFrameSignoff:
-				logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", m.flap)
+				logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", flap)
 				return nil
 			case wire.FLAPFrameKeepAlive:
 				logger.DebugContext(ctx, "keepalive heartbeat")
 			default:
-				return fmt.Errorf("got unknown FLAP frame type. flap: %v", m.flap)
+				return fmt.Errorf("got unknown FLAP frame type. flap: %v", flap)
 			}
-		case err := <-readErrCh:
+		case err := <-errCh:
 			if !errors.Is(io.EOF, err) {
 				logger.ErrorContext(ctx, "client disconnected with error", "err", err)
 			}

+ 5 - 9
server/oscar/auth_test.go

@@ -20,25 +20,21 @@ func TestBUCPAuthService_handleNewConnection(t *testing.T) {
 		// < receive FLAPSignonFrame
 		flap := wire.FLAPFrame{}
 		assert.NoError(t, wire.Unmarshal(&flap, serverReader))
-		buf, err := flap.ReadBody(serverReader)
-		assert.NoError(t, err)
 		flapSignonFrame := wire.FLAPSignonFrame{}
-		assert.NoError(t, wire.Unmarshal(&flapSignonFrame, buf))
+		assert.NoError(t, wire.Unmarshal(&flapSignonFrame, bytes.NewBuffer(flap.Payload)))
 
 		// > send FLAPSignonFrame
 		flapSignonFrame = wire.FLAPSignonFrame{
 			FLAPVersion: 1,
 		}
-		buf = &bytes.Buffer{}
+		buf := &bytes.Buffer{}
 		assert.NoError(t, wire.Marshal(flapSignonFrame, buf))
 		flap = wire.FLAPFrame{
-			StartMarker:   42,
-			FrameType:     wire.FLAPFrameSignon,
-			PayloadLength: uint16(buf.Len()),
+			StartMarker: 42,
+			FrameType:   wire.FLAPFrameSignon,
+			Payload:     buf.Bytes(),
 		}
 		assert.NoError(t, wire.Marshal(flap, serverWriter))
-		_, err = serverWriter.Write(buf.Bytes())
-		assert.NoError(t, err)
 
 		// > send SNAC_0x17_0x06_BUCPChallengeRequest
 		flapc := wire.NewFlapClient(0, serverReader, serverWriter)

+ 8 - 16
server/oscar/bos_test.go

@@ -38,39 +38,31 @@ func TestBOSService_handleNewConnection(t *testing.T) {
 		// < receive FLAPSignonFrame
 		flap := wire.FLAPFrame{}
 		assert.NoError(t, wire.Unmarshal(&flap, serverReader))
-		buf, err := flap.ReadBody(serverReader)
-		assert.NoError(t, err)
 		flapSignonFrame := wire.FLAPSignonFrame{}
-		assert.NoError(t, wire.Unmarshal(&flapSignonFrame, buf))
+		assert.NoError(t, wire.Unmarshal(&flapSignonFrame, bytes.NewBuffer(flap.Payload)))
 
 		// > send FLAPSignonFrame
 		flapSignonFrame = wire.FLAPSignonFrame{
 			FLAPVersion: 1,
 		}
 		flapSignonFrame.Append(wire.NewTLV(wire.OServiceTLVTagsLoginCookie, []byte("the-cookie")))
-		buf = &bytes.Buffer{}
+		buf := &bytes.Buffer{}
 		assert.NoError(t, wire.Marshal(flapSignonFrame, buf))
 		flap = wire.FLAPFrame{
-			StartMarker:   42,
-			FrameType:     wire.FLAPFrameSignon,
-			PayloadLength: uint16(buf.Len()),
+			StartMarker: 42,
+			FrameType:   wire.FLAPFrameSignon,
+			Payload:     buf.Bytes(),
 		}
 		assert.NoError(t, wire.Marshal(flap, serverWriter))
-		_, err = serverWriter.Write(buf.Bytes())
-		assert.NoError(t, err)
+
+		flapc := wire.NewFlapClient(0, serverReader, serverWriter)
 
 		// < receive SNAC_0x01_0x03_OServiceHostOnline
-		flap = wire.FLAPFrame{}
-		assert.NoError(t, wire.Unmarshal(&flap, serverReader))
-		buf, err = flap.ReadBody(serverReader)
-		assert.NoError(t, err)
 		frame := wire.SNACFrame{}
-		assert.NoError(t, wire.Unmarshal(&frame, buf))
 		body := wire.SNAC_0x01_0x03_OServiceHostOnline{}
-		assert.NoError(t, wire.Unmarshal(&body, buf))
+		assert.NoError(t, flapc.ReceiveSNAC(&frame, &body))
 
 		// send the first request that should get relayed to BOSRouter.Handle
-		flapc := wire.NewFlapClient(0, nil, serverWriter)
 		frame = wire.SNACFrame{
 			FoodGroup: wire.OService,
 			SubGroup:  wire.OServiceClientOnline,

+ 8 - 16
server/oscar/chat_test.go

@@ -24,39 +24,31 @@ func TestChatService_handleNewConnection(t *testing.T) {
 		// < receive FLAPSignonFrame
 		flap := wire.FLAPFrame{}
 		assert.NoError(t, wire.Unmarshal(&flap, serverReader))
-		buf, err := flap.ReadBody(serverReader)
-		assert.NoError(t, err)
 		flapSignonFrame := wire.FLAPSignonFrame{}
-		assert.NoError(t, wire.Unmarshal(&flapSignonFrame, buf))
+		assert.NoError(t, wire.Unmarshal(&flapSignonFrame, bytes.NewBuffer(flap.Payload)))
 
 		// > send FLAPSignonFrame
 		flapSignonFrame = wire.FLAPSignonFrame{
 			FLAPVersion: 1,
 		}
 		flapSignonFrame.Append(wire.NewTLV(wire.OServiceTLVTagsLoginCookie, []byte(`the-chat-login-cookie`)))
-		buf = &bytes.Buffer{}
+		buf := &bytes.Buffer{}
 		assert.NoError(t, wire.Marshal(flapSignonFrame, buf))
 		flap = wire.FLAPFrame{
-			StartMarker:   42,
-			FrameType:     wire.FLAPFrameSignon,
-			PayloadLength: uint16(buf.Len()),
+			StartMarker: 42,
+			FrameType:   wire.FLAPFrameSignon,
+			Payload:     buf.Bytes(),
 		}
 		assert.NoError(t, wire.Marshal(flap, serverWriter))
-		_, err = serverWriter.Write(buf.Bytes())
-		assert.NoError(t, err)
+
+		flapc := wire.NewFlapClient(0, serverReader, serverWriter)
 
 		// < receive SNAC_0x01_0x03_OServiceHostOnline
-		flap = wire.FLAPFrame{}
-		assert.NoError(t, wire.Unmarshal(&flap, serverReader))
-		buf, err = flap.ReadBody(serverReader)
-		assert.NoError(t, err)
 		frame := wire.SNACFrame{}
-		assert.NoError(t, wire.Unmarshal(&frame, buf))
 		body := wire.SNAC_0x01_0x03_OServiceHostOnline{}
-		assert.NoError(t, wire.Unmarshal(&body, buf))
+		assert.NoError(t, flapc.ReceiveSNAC(&frame, &body))
 
 		// send the first request that should get relayed to BOSRouter.Handle
-		flapc := wire.NewFlapClient(0, nil, serverWriter)
 		frame = wire.SNACFrame{
 			FoodGroup: wire.Chat,
 			SubGroup:  wire.ChatNavNavInfo,

+ 30 - 43
server/oscar/connection.go

@@ -14,11 +14,6 @@ import (
 	"github.com/mk6i/retro-aim-server/wire"
 )
 
-type incomingMessage struct {
-	flap    wire.FLAPFrame
-	payload *bytes.Buffer
-}
-
 func sendInvalidSNACErr(frameIn wire.SNACFrame, rw ResponseWriter) error {
 	frameOut := wire.SNACFrame{
 		FoodGroup: frameIn.FoodGroup,
@@ -31,30 +26,6 @@ func sendInvalidSNACErr(frameIn wire.SNACFrame, rw ResponseWriter) error {
 	return rw.SendSNAC(frameOut, bodyOut)
 }
 
-func consumeFLAPFrames(r io.Reader, msgCh chan incomingMessage, errCh chan error) {
-	defer close(msgCh)
-	defer close(errCh)
-
-	for {
-		in := incomingMessage{}
-		if err := wire.Unmarshal(&in.flap, r); err != nil {
-			errCh <- err
-			return
-		}
-
-		if in.flap.PayloadLength > 0 {
-			buf := make([]byte, in.flap.PayloadLength)
-			if _, err := io.ReadFull(r, buf); err != nil {
-				errCh <- err
-				return
-			}
-			in.payload = bytes.NewBuffer(buf)
-		}
-
-		msgCh <- in
-	}
-}
-
 // 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.
@@ -64,30 +35,46 @@ func consumeFLAPFrames(r io.Reader, msgCh chan incomingMessage, errCh chan error
 //
 // todo: this method has too many params and should be folded into a new type
 func dispatchIncomingMessages(ctx context.Context, sess *state.Session, flapc *wire.FlapClient, r io.Reader, logger *slog.Logger, router Handler, config config.Config) error {
-	// buffered so that the go routine has room to exit
-	msgCh := make(chan incomingMessage, 1)
-	readErrCh := make(chan error, 1)
-	go consumeFLAPFrames(r, msgCh, readErrCh)
-
 	defer func() {
 		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.Unmarshal(&frame, r); err != nil {
+				errCh <- err
+				return
+			}
+			msgCh <- frame
+		}
+	}()
+
 	for {
 		select {
-		case m, ok := <-msgCh:
+		case flap, ok := <-msgCh:
 			if !ok {
 				return nil
 			}
-			switch m.flap.FrameType {
+			switch flap.FrameType {
 			case wire.FLAPFrameData:
+				flapBuf := bytes.NewBuffer(flap.Payload)
+
 				inFrame := wire.SNACFrame{}
-				if err := wire.Unmarshal(&inFrame, m.payload); err != nil {
+				if err := wire.Unmarshal(&inFrame, flapBuf); err != nil {
 					return err
 				}
 				// route a client request to the appropriate service handler. the
 				// handler may write a response to the client connection.
-				if err := router.Handle(ctx, sess, inFrame, m.payload, flapc); err != nil {
+				if err := router.Handle(ctx, sess, inFrame, flapBuf, flapc); err != nil {
 					middleware.LogRequestError(ctx, logger, inFrame, err)
 					if errors.Is(err, ErrRouteNotFound) {
 						if err1 := sendInvalidSNACErr(inFrame, flapc); err1 != nil {
@@ -101,16 +88,16 @@ func dispatchIncomingMessages(ctx context.Context, sess *state.Session, flapc *w
 					return err
 				}
 			case wire.FLAPFrameSignon:
-				return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", m.flap)
+				return fmt.Errorf("shouldn't get FLAPFrameSignon. flap: %v", flap)
 			case wire.FLAPFrameError:
-				return fmt.Errorf("got FLAPFrameError. flap: %v", m.flap)
+				return fmt.Errorf("got FLAPFrameError. flap: %v", flap)
 			case wire.FLAPFrameSignoff:
-				logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", m.flap)
+				logger.InfoContext(ctx, "got FLAPFrameSignoff", "flap", flap)
 				return nil
 			case wire.FLAPFrameKeepAlive:
 				logger.DebugContext(ctx, "keepalive heartbeat")
 			default:
-				return fmt.Errorf("got unknown FLAP frame type. flap: %v", m.flap)
+				return fmt.Errorf("got unknown FLAP frame type. flap: %v", flap)
 			}
 		case m := <-sess.ReceiveMessage():
 			// forward a notification sent from another client to this client
@@ -126,7 +113,7 @@ func dispatchIncomingMessages(ctx context.Context, sess *state.Session, flapc *w
 				return fmt.Errorf("unable to gracefully disconnect user. %w", err)
 			}
 			return nil
-		case err := <-readErrCh:
+		case err := <-errCh:
 			if !errors.Is(io.EOF, err) {
 				logger.ErrorContext(ctx, "client disconnected with error", "err", err)
 			}

+ 4 - 4
server/oscar/connection_test.go

@@ -1,6 +1,7 @@
 package oscar
 
 import (
+	"bytes"
 	"context"
 	"io"
 	"log/slog"
@@ -67,13 +68,12 @@ func TestHandleChatConnection_MessageRelay(t *testing.T) {
 	for i := 0; i < len(inboundMsgs); i++ {
 		flap := wire.FLAPFrame{}
 		assert.NoError(t, wire.Unmarshal(&flap, clientReader))
-		snac, err := flap.ReadBody(clientReader)
-		assert.NoError(t, err)
 		frame := wire.SNACFrame{}
-		assert.NoError(t, wire.Unmarshal(&frame, snac))
+		buf := bytes.NewBuffer(flap.Payload)
+		assert.NoError(t, wire.Unmarshal(&frame, buf))
 		assert.Equal(t, inboundMsgs[i].Frame, frame)
 		body := wire.SNAC_0x0E_0x03_ChatUsersJoined{}
-		assert.NoError(t, wire.Unmarshal(&body, snac))
+		assert.NoError(t, wire.Unmarshal(&body, buf))
 		assert.Equal(t, inboundMsgs[i].Body, body)
 	}
 

+ 25 - 69
wire/frames.go

@@ -20,20 +20,10 @@ const (
 )
 
 type FLAPFrame struct {
-	StartMarker   uint8
-	FrameType     uint8
-	Sequence      uint16
-	PayloadLength uint16
-}
-
-func (f FLAPFrame) ReadBody(r io.Reader) (*bytes.Buffer, error) {
-	b := make([]byte, f.PayloadLength)
-	if f.PayloadLength > 0 {
-		if _, err := io.ReadFull(r, b); err != nil {
-			return nil, err
-		}
-	}
-	return bytes.NewBuffer(b), nil
+	StartMarker uint8
+	FrameType   uint8
+	Sequence    uint16
+	Payload     []byte `len_prefix:"uint16"`
 }
 
 type SNACFrame struct {
@@ -89,19 +79,15 @@ func (f *FlapClient) SendSignonFrame(tlvs []TLV) error {
 	}
 
 	flap := FLAPFrame{
-		StartMarker:   42,
-		FrameType:     FLAPFrameSignon,
-		Sequence:      uint16(f.sequence),
-		PayloadLength: uint16(buf.Len()),
+		StartMarker: 42,
+		FrameType:   FLAPFrameSignon,
+		Sequence:    uint16(f.sequence),
+		Payload:     buf.Bytes(),
 	}
 	if err := Marshal(flap, f.w); err != nil {
 		return err
 	}
 
-	if _, err := f.w.Write(buf.Bytes()); err != nil {
-		return err
-	}
-
 	f.sequence++
 
 	return nil
@@ -114,13 +100,8 @@ func (f *FlapClient) ReceiveSignonFrame() (FLAPSignonFrame, error) {
 		return FLAPSignonFrame{}, err
 	}
 
-	buf, err := flap.ReadBody(f.r)
-	if err != nil {
-		return FLAPSignonFrame{}, err
-	}
-
 	signonFrame := FLAPSignonFrame{}
-	if err := Unmarshal(&signonFrame, buf); err != nil {
+	if err := Unmarshal(&signonFrame, bytes.NewBuffer(flap.Payload)); err != nil {
 		return FLAPSignonFrame{}, err
 	}
 
@@ -129,21 +110,13 @@ func (f *FlapClient) ReceiveSignonFrame() (FLAPSignonFrame, error) {
 
 // ReceiveFLAP receives a FLAP frame and body. It only returns a body if the
 // FLAP frame is a data frame.
-func (f *FlapClient) ReceiveFLAP() (FLAPFrame, *bytes.Buffer, error) {
+func (f *FlapClient) ReceiveFLAP() (FLAPFrame, error) {
 	flap := FLAPFrame{}
-	if err := Unmarshal(&flap, f.r); err != nil {
-		return flap, nil, fmt.Errorf("unable to unmarshal FLAP frame: %w", err)
-	}
-
-	if flap.FrameType != FLAPFrameData {
-		return flap, nil, nil
-	}
-
-	buf, err := flap.ReadBody(f.r)
+	err := Unmarshal(&flap, f.r)
 	if err != nil {
-		err = fmt.Errorf("unable to read FLAP body: %w", err)
+		err = fmt.Errorf("unable to unmarshal FLAP frame: %w", err)
 	}
-	return flap, buf, err
+	return flap, err
 }
 
 // SendSignoffFrame sends a sign-off FLAP frame with attached TLVs as the last
@@ -157,25 +130,16 @@ func (f *FlapClient) SendSignoffFrame(tlvs TLVRestBlock) error {
 	}
 
 	flap := FLAPFrame{
-		StartMarker:   42,
-		FrameType:     FLAPFrameSignoff,
-		Sequence:      uint16(f.sequence),
-		PayloadLength: uint16(tlvBuf.Len()),
+		StartMarker: 42,
+		FrameType:   FLAPFrameSignoff,
+		Sequence:    uint16(f.sequence),
+		Payload:     tlvBuf.Bytes(),
 	}
 
 	if err := Marshal(flap, f.w); err != nil {
 		return err
 	}
 
-	expectLen := tlvBuf.Len()
-	c, err := f.w.Write(tlvBuf.Bytes())
-	if err != nil {
-		return err
-	}
-	if c != expectLen {
-		panic("did not write the expected # of bytes")
-	}
-
 	f.sequence++
 	return nil
 }
@@ -191,19 +155,15 @@ func (f *FlapClient) SendSNAC(frame SNACFrame, body any) error {
 	}
 
 	flap := FLAPFrame{
-		StartMarker:   42,
-		FrameType:     FLAPFrameData,
-		Sequence:      uint16(f.sequence),
-		PayloadLength: uint16(snacBuf.Len()),
+		StartMarker: 42,
+		FrameType:   FLAPFrameData,
+		Sequence:    uint16(f.sequence),
+		Payload:     snacBuf.Bytes(),
 	}
 	if err := Marshal(flap, f.w); err != nil {
 		return err
 	}
 
-	if _, err := f.w.Write(snacBuf.Bytes()); err != nil {
-		return err
-	}
-
 	f.sequence++
 	return nil
 }
@@ -214,10 +174,7 @@ func (f *FlapClient) ReceiveSNAC(frame *SNACFrame, body any) error {
 	if err := Unmarshal(&flap, f.r); err != nil {
 		return err
 	}
-	buf, err := flap.ReadBody(f.r)
-	if err != nil {
-		return err
-	}
+	buf := bytes.NewBuffer(flap.Payload)
 	if err := Unmarshal(frame, buf); err != nil {
 		return err
 	}
@@ -229,10 +186,9 @@ func (f *FlapClient) Disconnect() error {
 	// gracefully disconnect so that the client does not try to
 	// reconnect when the connection closes.
 	flap := FLAPFrame{
-		StartMarker:   42,
-		FrameType:     FLAPFrameSignoff,
-		Sequence:      uint16(f.sequence),
-		PayloadLength: uint16(0),
+		StartMarker: 42,
+		FrameType:   FLAPFrameSignoff,
+		Sequence:    uint16(f.sequence),
 	}
 	return Marshal(flap, f.w)
 }

+ 0 - 28
wire/frames_test.go

@@ -1,28 +0,0 @@
-package wire
-
-import (
-	"bytes"
-	"io"
-	"testing"
-
-	"github.com/stretchr/testify/assert"
-)
-
-func TestFLAPFrame_ReadBody(t *testing.T) {
-	flap := FLAPFrame{
-		PayloadLength: 4,
-	}
-	bufIn := bytes.NewBuffer([]byte{0, 1, 2, 3, 4, 5})
-	buf, err := flap.ReadBody(bufIn)
-	assert.NoError(t, err)
-	assert.Equal(t, []byte{0, 1, 2, 3}, buf.Bytes())
-}
-
-func TestFLAPFrame_ReadBodyError(t *testing.T) {
-	flap := FLAPFrame{
-		PayloadLength: 4,
-	}
-	bufIn := &bytes.Buffer{}
-	_, err := flap.ReadBody(bufIn)
-	assert.ErrorIs(t, err, io.EOF)
-}