diff options
Diffstat (limited to 'internal/packet')
| -rw-r--r-- | internal/packet/encoders.go | 52 | ||||
| -rw-r--r-- | internal/packet/framer.go | 109 | ||||
| -rw-r--r-- | internal/packet/packet.go | 163 |
3 files changed, 324 insertions, 0 deletions
diff --git a/internal/packet/encoders.go b/internal/packet/encoders.go new file mode 100644 index 0000000..0d190e9 --- /dev/null +++ b/internal/packet/encoders.go @@ -0,0 +1,52 @@ +package packet + +import ( + "bytes" + "encoding/json" + "io" + + "github.com/vmihailenco/msgpack/v5" +) + +type TypedMessage interface{ + Type() PacketType +} + +type defaultPacketEncoder struct { + io.Reader + encoding Encoding + packetType PacketType +} + +func (e defaultPacketEncoder) Encoding() Encoding { + return e.encoding +} + +func (e defaultPacketEncoder) Type() PacketType { + return e.packetType +} + +func NewJsonEncoder(message TypedMessage) (PacketEncoder, error) { + data, err := json.Marshal(message) + if err != nil { + return nil, err + } + + return defaultPacketEncoder{ + Reader: bytes.NewReader(data), + encoding: EncodingJson, + packetType: message.Type(), + }, nil +} + +func NewMsgPackEncoder(message TypedMessage) (PacketEncoder, error) { + data, err := msgpack.Marshal(message) + if err != nil { + return nil, err + } + + return defaultPacketEncoder{ + Reader: bytes.NewReader(data), + encoding: EncodingJson, + packetType: message.Type(), + }, nil diff --git a/internal/packet/framer.go b/internal/packet/framer.go new file mode 100644 index 0000000..58fa9c2 --- /dev/null +++ b/internal/packet/framer.go @@ -0,0 +1,109 @@ +package packet + +import ( + "context" + "encoding/binary" + "errors" + "io" + + "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) { + // Reads a bunch from ioReader + // If bytes r more than header size + // Try parsing 1 or more packets + // Send those packets to a channel + // Have a goroutine read from the channel + + ch := make(chan Packet, framerPacketCapacity) + errCh := make(chan error) + + framer := packetFramer{ + buffer: make([]byte, PACKET_MAX_SIZE, 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:]) + 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 PacketUnsupportedVersion + } + + 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, fullLength) + copy(packetBuffer, f.buffer[HEADER_SIZE:fullLength]) + + f.len = uint16(copy(f.buffer[:fullLength], f.buffer[fullLength:])) + + f.in <- Packet{packetBuffer} + } + return nil +} diff --git a/internal/packet/packet.go b/internal/packet/packet.go new file mode 100644 index 0000000..21ba0cc --- /dev/null +++ b/internal/packet/packet.go @@ -0,0 +1,163 @@ +package packet + +import ( + "encoding/binary" + "encoding/json" + "fmt" + "io" + + "github.com/vmihailenco/msgpack/v5" + + "github.com/kyren223/eko/pkg/assert" +) + +type Encoding uint8 + +func (e Encoding) String() string { + switch e { + case EncodingJson: + return "EncodingJson" + case EncodingMsgPack: + return "EncodingMsgPack" + case EncodingUnused1: + return "EncodingUnused1" + case EncodingUnused2: + return "EncodingUnused2" + default: + return fmt.Sprintf("EncodingInvalid(%v)", e) + } +} + +func (e Encoding) IsSupported() bool { + switch e { + case EncodingJson, EncodingMsgPack: + return true + default: + return false + } +} + +const ( + EncodingJson Encoding = iota + EncodingMsgPack + EncodingUnused1 + EncodingUnused2 +) + +type PacketType uint8 + +func (t PacketType) String() string { + switch t { + case TypeEko: + return "PacketTypeEko" + case TypeError: + return "PacketTypeError" + default: + return fmt.Sprintf("PacketTypeInvalid(%v)", t) + } +} + +func (e PacketType) IsSupported() bool { + switch e { + case TypeEko, TypeError: + return true + default: + return false + } +} + +const ( + TypeEko PacketType = iota + TypeError +) + +const ( + VERSION = 1 + PACKET_MAX_SIZE = ^uint16(0) + PAYLOAD_MAX_SIZE = PACKET_MAX_SIZE - HEADER_SIZE + HEADER_SIZE = 4 + VERSION_OFFSET = 0 + TYPE_OFFSET = 1 + ENCODING_OFFSET = 1 + LENGTH_OFFSET = 2 +) + +type PacketEncoder interface { + io.Reader + Encoding() Encoding + Type() PacketType +} + +// The following diagram shows the packet structure: +// +// 0 1 2 3 +// 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +// +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +// | Version |En.| Type | Payload Length | +// +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +// | Payload... Payload Length bytes ... | +// +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +type Packet struct { + data []byte +} + +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") + + binary.BigEndian.PutUint16(data[LENGTH_OFFSET:], uint16(n)) + + data[VERSION_OFFSET] = VERSION + + packetType, encoding := byte(encoder.Type()), byte(encoder.Encoding()) + assert.Assert(packetType <= 63, "packet type exceeded allowed size type=%v", packetType) + assert.Assert(encoding <= 3, "encoding exceeded allowed permutations encoding=%v", encoding) + data[TYPE_OFFSET] = packetType | encoding<<6 + + return Packet{data[:HEADER_SIZE+n]} +} + +func (p Packet) Version() uint8 { + return p.data[VERSION_OFFSET] +} + +func (p Packet) Type() PacketType { + return PacketType(p.data[TYPE_OFFSET] & 63) // 2^6-1 +} + +func (p Packet) Encoding() Encoding { + return Encoding(p.data[ENCODING_OFFSET] >> 6) +} + +func (p Packet) PayloadLength() uint16 { + return binary.BigEndian.Uint16(p.data[LENGTH_OFFSET:]) +} + +// 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) DecodePayload(v *TypedMessage) error { + if p.Type() != (*v).Type() { + return fmt.Errorf("", p.Type(), (*v).Type()) + } + switch p.Encoding() { + case EncodingJson: + return json.Unmarshal(p.Payload(), v) + case EncodingMsgPack: + return msgpack.Unmarshal(p.Payload(), v) + case EncodingUnused1: + fallthrough + case EncodingUnused2: + return fmt.Errorf("unsupported encoding: ", p.Encoding().String()) + default: + assert.Unreachable("encoding from packet should always be valid encoding=%v", p.Encoding()) + } +} + +func (p Packet) Into(writer io.Writer) (int, error) { + return writer.Write(p.data[:len(p.data)]) +} |
