summaryrefslogtreecommitdiff
path: root/internal/packet
diff options
context:
space:
mode:
Diffstat (limited to 'internal/packet')
-rw-r--r--internal/packet/encoders.go7
-rw-r--r--internal/packet/framer.go19
-rw-r--r--internal/packet/messages.go17
-rw-r--r--internal/packet/packet.go24
-rw-r--r--internal/packet/packet_test.go50
5 files changed, 94 insertions, 23 deletions
diff --git a/internal/packet/encoders.go b/internal/packet/encoders.go
index 0d190e9..e70d0e8 100644
--- a/internal/packet/encoders.go
+++ b/internal/packet/encoders.go
@@ -8,13 +8,13 @@ import (
"github.com/vmihailenco/msgpack/v5"
)
-type TypedMessage interface{
+type TypedMessage interface {
Type() PacketType
}
type defaultPacketEncoder struct {
io.Reader
- encoding Encoding
+ encoding Encoding
packetType PacketType
}
@@ -47,6 +47,7 @@ func NewMsgPackEncoder(message TypedMessage) (PacketEncoder, error) {
return defaultPacketEncoder{
Reader: bytes.NewReader(data),
- encoding: EncodingJson,
+ encoding: EncodingMsgPack,
packetType: message.Type(),
}, nil
+}
diff --git a/internal/packet/framer.go b/internal/packet/framer.go
index 58fa9c2..1a76ea6 100644
--- a/internal/packet/framer.go
+++ b/internal/packet/framer.go
@@ -4,8 +4,10 @@ import (
"context"
"encoding/binary"
"errors"
+ "fmt"
"io"
+ "github.com/kyren223/eko/pkg/assert"
"github.com/kyren223/eko/pkg/util"
)
@@ -25,17 +27,11 @@ type packetFramer struct {
}
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),
+ buffer: make([]byte, PACKET_MAX_SIZE),
len: 0,
in: ch,
inErr: errCh,
@@ -57,6 +53,7 @@ outer:
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 {
@@ -79,7 +76,7 @@ outer:
func (f *packetFramer) parse() error {
for f.len > HEADER_SIZE {
if f.buffer[VERSION_OFFSET] != VERSION {
- return PacketUnsupportedVersion
+ return fmt.Errorf("%w version=%v", PacketUnsupportedVersion, f.buffer[VERSION_OFFSET])
}
encoding := Encoding(f.buffer[ENCODING_OFFSET] >> 6)
@@ -98,10 +95,10 @@ func (f *packetFramer) parse() error {
}
fullLength := HEADER_SIZE + length
- packetBuffer := make([]byte, fullLength, fullLength)
- copy(packetBuffer, f.buffer[HEADER_SIZE:fullLength])
+ packetBuffer := make([]byte, fullLength)
+ copy(packetBuffer, f.buffer[:fullLength])
- f.len = uint16(copy(f.buffer[:fullLength], f.buffer[fullLength:]))
+ f.len = uint16(copy(f.buffer, f.buffer[fullLength:f.len]))
f.in <- Packet{packetBuffer}
}
diff --git a/internal/packet/messages.go b/internal/packet/messages.go
new file mode 100644
index 0000000..f8390f8
--- /dev/null
+++ b/internal/packet/messages.go
@@ -0,0 +1,17 @@
+package packet
+
+type EkoMessage struct {
+ Message string `msgpack:"message"`
+}
+
+func (m *EkoMessage) Type() PacketType {
+ return TypeEko
+}
+
+type ErrorMessage struct {
+ Error string `msgpack:"error"`
+}
+
+func (m *ErrorMessage) Type() PacketType {
+ return TypeError
+}
diff --git a/internal/packet/packet.go b/internal/packet/packet.go
index 21ba0cc..068a8f3 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -24,7 +24,7 @@ func (e Encoding) String() string {
case EncodingUnused2:
return "EncodingUnused2"
default:
- return fmt.Sprintf("EncodingInvalid(%v)", e)
+ return fmt.Sprintf("EncodingInvalid(%v)", byte(e))
}
}
@@ -53,7 +53,7 @@ func (t PacketType) String() string {
case TypeError:
return "PacketTypeError"
default:
- return fmt.Sprintf("PacketTypeInvalid(%v)", t)
+ return fmt.Sprintf("PacketTypeInvalid(%v)", byte(t))
}
}
@@ -72,7 +72,7 @@ const (
)
const (
- VERSION = 1
+ VERSION = byte(1)
PACKET_MAX_SIZE = ^uint16(0)
PAYLOAD_MAX_SIZE = PACKET_MAX_SIZE - HEADER_SIZE
HEADER_SIZE = 4
@@ -135,14 +135,18 @@ func (p Packet) PayloadLength() uint16 {
return binary.BigEndian.Uint16(p.data[LENGTH_OFFSET:])
}
+func (p Packet) String() string {
+ return fmt.Sprintf("{v%v %v %v %v: %v}", p.data[0], p.Encoding().String(), p.Type().String(), p.PayloadLength(), p.Payload())
+}
+
// 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())
+func (p Packet) DecodePayload(v TypedMessage) error {
+ if p.Type() != v.Type() {
+ return fmt.Errorf("type mismatch: want %v got %v", p.Type(), v.Type())
}
switch p.Encoding() {
case EncodingJson:
@@ -152,12 +156,14 @@ func (p Packet) DecodePayload(v *TypedMessage) error {
case EncodingUnused1:
fallthrough
case EncodingUnused2:
- return fmt.Errorf("unsupported encoding: ", p.Encoding().String())
+ return fmt.Errorf("unsupported encoding: %v", p.Encoding().String())
default:
assert.Unreachable("encoding from packet should always be valid encoding=%v", p.Encoding())
+ return nil
}
}
-func (p Packet) Into(writer io.Writer) (int, error) {
- return writer.Write(p.data[:len(p.data)])
+func (p Packet) Into(writer io.Writer) error {
+ _, err := writer.Write(p.data[:len(p.data)])
+ return err
}
diff --git a/internal/packet/packet_test.go b/internal/packet/packet_test.go
new file mode 100644
index 0000000..20370af
--- /dev/null
+++ b/internal/packet/packet_test.go
@@ -0,0 +1,50 @@
+package packet
+
+import (
+ "testing"
+
+ "github.com/vmihailenco/msgpack/v5"
+)
+
+func TestMsgPackEncoding(t *testing.T) {
+ request := EkoMessage{Message: "test"}
+ data, err := msgpack.Marshal(&request)
+ if err != nil {
+ t.Errorf("encoding error: %v", err)
+ return
+ }
+ var response EkoMessage
+ err = msgpack.Unmarshal(data, &response)
+ if err != nil {
+ t.Errorf("decoding error: %v", err)
+ return
+ }
+ if request.Message != response.Message {
+ t.Errorf("%v != %v", request.Message, response.Message)
+ return
+ }
+}
+
+func TestPacketMsgPackEncoding(t *testing.T) {
+ request := EkoMessage{Message: "test"}
+ encoder1, err := NewMsgPackEncoder(&request)
+ encoder2, _ := NewMsgPackEncoder(&request)
+ if err != nil {
+ t.Errorf("encoding error: %v", err)
+ return
+ }
+ encodedBytes := make([]byte, PACKET_MAX_SIZE)
+ n, _ := encoder2.Read(encodedBytes[HEADER_SIZE:])
+
+ packet := NewPacket(encoder1)
+ var response EkoMessage
+ err = packet.DecodePayload(&response)
+ if err != nil {
+ t.Errorf("decoding error: %#v: packet: %v encoder: %v", err, packet.Payload(), encodedBytes[HEADER_SIZE:HEADER_SIZE+n])
+ return
+ }
+ if request.Message != response.Message {
+ t.Errorf("%v != %v", request.Message, response.Message)
+ return
+ }
+}