summaryrefslogtreecommitdiff
path: root/internal/packet
diff options
context:
space:
mode:
Diffstat (limited to 'internal/packet')
-rw-r--r--internal/packet/framer.go106
-rw-r--r--internal/packet/framer_test.go63
-rw-r--r--internal/packet/packet.go145
-rw-r--r--internal/packet/packet_test.go58
4 files changed, 163 insertions, 209 deletions
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,7 +70,7 @@ 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
@@ -73,16 +78,8 @@ func (e PacketType) IsSupported() bool {
}
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
+ }
+}