From 1997bc8a150b92783060cc7129e1ee72a761182b Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Sun, 20 Oct 2024 11:36:43 +0300 Subject: refactor: mid refactor --- internal/packet/framer.go | 106 ------------------------------ internal/packet/framer_test.go | 63 ------------------ internal/packet/packet.go | 145 +++++++++++++++++++++++++++++------------ internal/packet/packet_test.go | 58 +++++++++++++++++ 4 files changed, 163 insertions(+), 209 deletions(-) delete mode 100644 internal/packet/framer.go delete mode 100644 internal/packet/framer_test.go (limited to 'internal/packet') diff --git a/internal/packet/framer.go b/internal/packet/framer.go deleted file mode 100644 index 1a76ea6..0000000 --- a/internal/packet/framer.go +++ /dev/null @@ -1,106 +0,0 @@ -package packet - -import ( - "context" - "encoding/binary" - "errors" - "fmt" - "io" - - "github.com/kyren223/eko/pkg/assert" - "github.com/kyren223/eko/pkg/util" -) - -const framerPacketCapacity = 10 - -var ( - PacketUnsupportedVersion error = errors.New("packet error: unsupported version") - PacketUnsupportedEncoding error = errors.New("packet error: unsupported encoding") - PacketUnsupportedType error = errors.New("packet error: unsupported type") -) - -type packetFramer struct { - buffer []byte - len uint16 - in chan<- Packet - inErr chan<- error -} - -func RunFramer(ctx context.Context, reader io.Reader) (out <-chan Packet, outErr <-chan error) { - ch := make(chan Packet, framerPacketCapacity) - errCh := make(chan error) - - framer := packetFramer{ - buffer: make([]byte, PACKET_MAX_SIZE), - len: 0, - in: ch, - inErr: errCh, - } - - go framer.run(ctx, util.NewChannelReader(ctx, reader)) - - return ch, errCh -} - -func (f *packetFramer) run(ctx context.Context, reader util.ChannelReader) { - defer close(f.in) - defer close(f.inErr) - -outer: - for { - select { - case data := <-reader.Out: - dataRead := 0 - for dataRead < len(data) { - n := copy(f.buffer[f.len:], data[dataRead:]) - assert.Assert(0 <= n && n <= int(PACKET_MAX_SIZE), "n must fit in a u16") - f.len += uint16(n) - dataRead += n - if err := f.parse(); err != nil { - f.inErr <- err - break outer - } - } - - case err := <-reader.Err: - f.inErr <- err - break outer - - case <-ctx.Done(): - f.inErr <- ctx.Err() - break outer - } - } -} - -func (f *packetFramer) parse() error { - for f.len > HEADER_SIZE { - if f.buffer[VERSION_OFFSET] != VERSION { - return fmt.Errorf("%w version=%v", PacketUnsupportedVersion, f.buffer[VERSION_OFFSET]) - } - - encoding := Encoding(f.buffer[ENCODING_OFFSET] >> 6) - packetType := PacketType(f.buffer[TYPE_OFFSET] & 63) - if !encoding.IsSupported() { - return PacketUnsupportedEncoding - } - if !packetType.IsSupported() { - return PacketUnsupportedType - } - - length := binary.BigEndian.Uint16(f.buffer[LENGTH_OFFSET:]) - if f.len-HEADER_SIZE < length { - // Wait for more data to arrive - return nil - } - - fullLength := HEADER_SIZE + length - packetBuffer := make([]byte, fullLength) - copy(packetBuffer, f.buffer[:fullLength]) - - f.len = uint16(copy(f.buffer, f.buffer[fullLength:f.len])) - - f.in <- Packet{packetBuffer} - } - return nil -} diff --git a/internal/packet/framer_test.go b/internal/packet/framer_test.go deleted file mode 100644 index 41a4396..0000000 --- a/internal/packet/framer_test.go +++ /dev/null @@ -1,63 +0,0 @@ -package packet - -import ( - "context" - "io" - "testing" - "time" -) - -type TestIoReader struct { - data []byte - len int -} - -func (r *TestIoReader) Read(data []byte) (int, error) { - if len(r.data) <= r.len { - // log.Println("RETURNING EOF:", len(data), r.len) - return 0, io.EOF - } - n := copy(data, r.data[r.len:]) - r.len += n - // log.Println("RETURNING N:", len(r.data), r.len) - return n, nil -} - -func (r *TestIoReader) start(msg TypedMessage, t *testing.T) { - encoder, err := NewMsgPackEncoder(msg) - if err != nil { - t.Errorf("encoding error: %v", err) - return - } - r.data = NewPacket(encoder).data - // log.Println("LEN:", len(r.data)) -} - -func TestPacketFramer(t *testing.T) { - reader := &TestIoReader{} - - // Long msg to test multiple - msg := EkoMessage{"Testing FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting Framer"} - reader.start(&msg, t) - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - out, outErr := RunFramer(ctx, reader) - - var message EkoMessage - select { - case packet := <-out: - if err := packet.DecodePayload(&message); err != nil { - t.Errorf("error decoding response: %v", err) - return - } - if msg.Message != message.Message { - t.Errorf("%v != %v", msg.Message, message.Message) - return - } - - case err := <-outErr: - t.Errorf("error receiving packet: %v", err) - return - } -} diff --git a/internal/packet/packet.go b/internal/packet/packet.go index 634643e..df3ee73 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -1,10 +1,13 @@ package packet import ( + "context" "encoding/binary" "encoding/json" + "errors" "fmt" "io" + "math" "github.com/vmihailenco/msgpack/v5" @@ -13,6 +16,13 @@ import ( type Encoding uint8 +const ( + EncodingJson Encoding = iota + EncodingMsgPack + EncodingUnused1 + EncodingUnused2 +) + func (e Encoding) String() string { switch e { case EncodingJson: @@ -37,26 +47,21 @@ func (e Encoding) IsSupported() bool { } } +type PacketType uint8 + const ( - EncodingJson Encoding = iota - EncodingMsgPack - EncodingUnused1 - EncodingUnused2 + PacketError PacketType = iota + PacketSendMessage + PacketMessages ) -type PacketType uint8 - func (t PacketType) String() string { switch t { - case TypeEko: - return "PacketTypeEko" - case TypeError: - return "PacketTypeError" - case TypeGetMessages: - return "PacketTypeGetMessages" - case TypeSendMessage: + case PacketError: + return "PacketError" + case PacketSendMessage: return "PacketTypeSendMessage" - case TypeMessages: + case PacketMessages: return "PacketTypeMessages" default: return fmt.Sprintf("PacketTypeInvalid(%v)", byte(t)) @@ -65,24 +70,16 @@ func (t PacketType) String() string { func (e PacketType) IsSupported() bool { switch e { - case TypeEko, TypeError, TypeGetMessages, TypeSendMessage, TypeMessages: + case PacketError, PacketSendMessage, PacketMessages: return true default: return false } } -const ( - TypeEko PacketType = iota - TypeError - TypeGetMessages - TypeSendMessage - TypeMessages -) - const ( VERSION = byte(1) - PACKET_MAX_SIZE = ^uint16(0) + PACKET_MAX_SIZE = math.MaxUint16 PAYLOAD_MAX_SIZE = PACKET_MAX_SIZE - HEADER_SIZE HEADER_SIZE = 4 VERSION_OFFSET = 0 @@ -92,7 +89,7 @@ const ( ) type PacketEncoder interface { - io.Reader + Payload() []byte Encoding() Encoding Type() PacketType } @@ -111,12 +108,11 @@ type Packet struct { } func NewPacket(encoder PacketEncoder) Packet { - data := make([]byte, PACKET_MAX_SIZE) - - n, err := encoder.Read(data[HEADER_SIZE:]) - assert.NoError(err, "packet encoder should never error when reading") + payload := encoder.Payload() + n := len(payload) + assert.Assert(0 <= n && n <= PAYLOAD_MAX_SIZE, "size of payload must be valid") - binary.BigEndian.PutUint16(data[LENGTH_OFFSET:], uint16(n)) + data := make([]byte, HEADER_SIZE+n) data[VERSION_OFFSET] = VERSION @@ -125,7 +121,11 @@ func NewPacket(encoder PacketEncoder) Packet { assert.Assert(encoding <= 3, "encoding exceeded allowed permutations encoding=%v", encoding) data[TYPE_OFFSET] = packetType | encoding<<6 - return Packet{data[:HEADER_SIZE+n]} + binary.BigEndian.PutUint16(data[LENGTH_OFFSET:], uint16(n)) + + copy(data[HEADER_SIZE:], payload) + + return Packet{data} } func (p Packet) Version() uint8 { @@ -133,7 +133,7 @@ func (p Packet) Version() uint8 { } func (p Packet) Type() PacketType { - return PacketType(p.data[TYPE_OFFSET] & 63) // 2^6-1 + return PacketType(p.data[TYPE_OFFSET] & 63) } func (p Packet) Encoding() Encoding { @@ -144,13 +144,16 @@ func (p Packet) PayloadLength() uint16 { return binary.BigEndian.Uint16(p.data[LENGTH_OFFSET:]) } +func (p Packet) Payload() []byte { + return p.data[HEADER_SIZE:] +} + func (p Packet) String() string { - return fmt.Sprintf("{v%v %v %v [%v bytes...]}", p.Version(), p.Encoding().String(), p.Type().String(), p.PayloadLength()) + return fmt.Sprintf("Packet(v%v %v %v [%v bytes...])", p.Version(), p.Encoding().String(), p.Type().String(), p.PayloadLength()) } -// The payload data, caller must not modify the returned slice, even temporarily -func (p Packet) Payload() []byte { - return p.data[HEADER_SIZE:] +func (p Packet) Into(writer io.Writer) (int, error) { + return writer.Write(p.data) } func (p Packet) DecodePayload(v TypedMessage) error { @@ -167,12 +170,74 @@ func (p Packet) DecodePayload(v TypedMessage) error { case EncodingUnused2: return fmt.Errorf("unsupported encoding: %v", p.Encoding().String()) default: - assert.Unreachable("encoding from packet should always be valid encoding=%v", p.Encoding()) + assert.Never("encoding from packet should always be valid encoding=%v", p.Encoding()) return nil } } -func (p Packet) Into(writer io.Writer) error { - _, err := writer.Write(p.data[:len(p.data)]) - return err +var ( + PacketUnsupportedVersion error = errors.New("packet error: unsupported version") + PacketUnsupportedEncoding error = errors.New("packet error: unsupported encoding") + PacketUnsupportedType error = errors.New("packet error: unsupported type") +) + +type packetFramer struct { + buffer []byte + Out chan Packet +} + +func NewFramer(ctx context.Context) packetFramer { + return packetFramer{ + Out: make(chan Packet, 10), + } +} + +func (f *packetFramer) Push(ctx context.Context, data []byte) error { + f.buffer = append(f.buffer, data...) + + for { + packet, err := f.parse() + if packet == nil || err != nil { + return err + } + select { + case f.Out <- *packet: + case <-ctx.Done(): + return ctx.Err() + } + } +} + +func (f *packetFramer) parse() (*Packet, error) { + if len(f.buffer) < HEADER_SIZE { + return nil, nil + } + + if f.buffer[VERSION_OFFSET] != VERSION { + return nil, fmt.Errorf("%w version=%v", PacketUnsupportedVersion, f.buffer[VERSION_OFFSET]) + } + + encoding := Encoding(f.buffer[ENCODING_OFFSET] >> 6) + if !encoding.IsSupported() { + return nil, PacketUnsupportedEncoding + } + + packetType := PacketType(f.buffer[TYPE_OFFSET] & 63) + if !packetType.IsSupported() { + return nil, PacketUnsupportedType + } + + length := binary.BigEndian.Uint16(f.buffer[LENGTH_OFFSET:]) + if len(f.buffer)-HEADER_SIZE < int(length) { + // Wait for more data to arrive + return nil, nil + } + + fullLength := HEADER_SIZE + length + packetBuffer := make([]byte, fullLength) + copy(packetBuffer, f.buffer[:fullLength]) + copy(f.buffer, f.buffer[fullLength:]) + f.buffer = f.buffer[:len(f.buffer)-int(fullLength)] + + return &Packet{packetBuffer}, nil } diff --git a/internal/packet/packet_test.go b/internal/packet/packet_test.go index 20370af..a523d59 100644 --- a/internal/packet/packet_test.go +++ b/internal/packet/packet_test.go @@ -1,7 +1,10 @@ package packet import ( + "context" + "io" "testing" + "time" "github.com/vmihailenco/msgpack/v5" ) @@ -48,3 +51,58 @@ func TestPacketMsgPackEncoding(t *testing.T) { return } } + +type TestIoReader struct { + data []byte + len int +} + +func (r *TestIoReader) Read(data []byte) (int, error) { + if len(r.data) <= r.len { + // log.Println("RETURNING EOF:", len(data), r.len) + return 0, io.EOF + } + n := copy(data, r.data[r.len:]) + r.len += n + // log.Println("RETURNING N:", len(r.data), r.len) + return n, nil +} + +func (r *TestIoReader) start(msg TypedMessage, t *testing.T) { + encoder, err := NewMsgPackEncoder(msg) + if err != nil { + t.Errorf("encoding error: %v", err) + return + } + r.data = NewPacket(encoder).data + // log.Println("LEN:", len(r.data)) +} + +func TestPacketFramer(t *testing.T) { + reader := &TestIoReader{} + + // Long msg to test multiple + msg := EkoMessage{"Testing FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting Framer"} + reader.start(&msg, t) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + out, outErr := RunFramer(ctx, reader) + + var message EkoMessage + select { + case packet := <-out: + if err := packet.DecodePayload(&message); err != nil { + t.Errorf("error decoding response: %v", err) + return + } + if msg.Message != message.Message { + t.Errorf("%v != %v", msg.Message, message.Message) + return + } + + case err := <-outErr: + t.Errorf("error receiving packet: %v", err) + return + } +} -- cgit v1.3.1