summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--go.mod5
-rw-r--r--go.sum4
-rw-r--r--internal/client/client.go2
-rw-r--r--internal/messages/eko.go19
-rw-r--r--internal/packet/encoders.go52
-rw-r--r--internal/packet/framer.go109
-rw-r--r--internal/packet/packet.go163
-rw-r--r--internal/server/server.go6
8 files changed, 356 insertions, 4 deletions
diff --git a/go.mod b/go.mod
index 4e100c4..c0d621e 100644
--- a/go.mod
+++ b/go.mod
@@ -1,3 +1,8 @@
module github.com/kyren223/eko
go 1.23.2
+
+require (
+ github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
+ github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
+)
diff --git a/go.sum b/go.sum
new file mode 100644
index 0000000..84eba6c
--- /dev/null
+++ b/go.sum
@@ -0,0 +1,4 @@
+github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IUPn0Bjt8=
+github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
+github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g=
+github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds=
diff --git a/internal/client/client.go b/internal/client/client.go
index 3345a9f..9aee12e 100644
--- a/internal/client/client.go
+++ b/internal/client/client.go
@@ -18,7 +18,7 @@ var certPEM []byte
func Run() {
certPool := x509.NewCertPool()
if !certPool.AppendCertsFromPEM(certPEM) {
- log.Fatalf("failed to append server certificate")
+ log.Fatalln("failed to append server certificate")
}
tlsConfig := &tls.Config{
diff --git a/internal/messages/eko.go b/internal/messages/eko.go
new file mode 100644
index 0000000..d8eaa86
--- /dev/null
+++ b/internal/messages/eko.go
@@ -0,0 +1,19 @@
+package messages
+
+import "github.com/kyren223/eko/internal/packet"
+
+type EkoMessage struct {
+ Message string `msgpack:"message"`
+}
+
+func (m EkoMessage) Type() packet.PacketType {
+ return packet.TypeEko
+}
+
+type ErrorMessage struct {
+ Error string `msgpack:"error"`
+}
+
+func (m ErrorMessage) Type() packet.PacketType {
+ return packet.TypeError
+}
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)])
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index 55e0763..48b31f4 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -24,7 +24,7 @@ var keyPEM []byte
func Start() {
cert, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
- log.Fatalf("error loading certificate: %s", err)
+ log.Fatalln("error loading certificate:", err)
}
tlsConfig := &tls.Config{
@@ -33,7 +33,7 @@ func Start() {
listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(port), tlsConfig)
if err != nil {
- log.Fatal("error starting server: %s", err)
+ log.Fatalf("error starting server: %s", err)
}
defer listener.Close()
@@ -48,7 +48,7 @@ func Start() {
func handleInterrupt(listener net.Listener, stopChan <-chan os.Signal) {
signal := <-stopChan
- log.Println(signal.String())
+ log.Println("signal:", signal.String())
listener.Close()
}