feat: communication using custom packet protocol

This commit is contained in:
Kyren223
2024-10-15 21:59:08 +03:00
parent c5dba4360a
commit 260ad0e1bd
9 changed files with 205 additions and 63 deletions

View File

@@ -2,6 +2,7 @@ package client
import (
"bufio"
"context"
"crypto/tls"
"crypto/x509"
_ "embed"
@@ -10,6 +11,8 @@ import (
"os"
"strings"
"time"
"github.com/kyren223/eko/internal/packet"
)
//go:embed server.crt
@@ -22,7 +25,7 @@ func Run() {
}
tlsConfig := &tls.Config{
RootCAs: certPool,
RootCAs: certPool,
ServerName: "localhost",
}
@@ -48,20 +51,35 @@ func processRequest(request string, tlsConfig *tls.Config) error {
}
defer conn.Close()
conn.SetDeadline(time.Now().Add(time.Second))
log.Println("established connection with server:", conn.RemoteAddr().String())
_, err = conn.Write([]byte(request))
requestMsg := packet.EkoMessage{Message: request}
encoder, err := packet.NewMsgPackEncoder(&requestMsg)
if err != nil {
return fmt.Errorf("error encoding request: %v", err)
}
requestPacket := packet.NewPacket(encoder)
err = requestPacket.Into(conn)
if err != nil {
return fmt.Errorf("error sending request: %v", err)
}
log.Println("sent request to server")
buffer := make([]byte, 1024)
n, err := conn.Read(buffer)
if err != nil {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
out, outErr := packet.RunFramer(ctx, conn)
var response packet.EkoMessage
select {
case responsePacket := <-out:
if err := responsePacket.DecodePayload(&response); err != nil {
return fmt.Errorf("error decoding response: %v", err)
}
case err := <-outErr:
return fmt.Errorf("error receiving response: %v", err)
}
log.Println("server response:", string(buffer[:n]))
log.Println("server response:", response.Message)
return nil
}

View File

@@ -1,19 +0,0 @@
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
}

View File

@@ -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
}

View File

@@ -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}
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}
}

View File

@@ -1,10 +1,16 @@
package server
import (
"context"
"errors"
"fmt"
"log"
"net"
"sync"
"time"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/pkg/assert"
)
func handleConnection(conn net.Conn, wg *sync.WaitGroup) {
@@ -13,20 +19,66 @@ func handleConnection(conn net.Conn, wg *sync.WaitGroup) {
defer conn.Close()
defer wg.Done()
buffer := make([]byte, 1024)
bytesRead, err := conn.Read(buffer)
if err != nil {
log.Println("failed reading request:", err)
return
}
request := string(buffer[:bytesRead])
log.Printf("Read %v bytes: %v\n", bytesRead, request)
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
out, outErr := packet.RunFramer(ctx, conn)
log.Printf("client %v: running framer\n", conn.RemoteAddr().String())
response := []byte(fmt.Sprintf("Eko \"%v\"", request))
bytesWritten, err := conn.Write(response)
if err != nil {
log.Println("failed writing response:", err)
return
outer:
for {
select {
case packet := <-out:
log.Printf("client %v: request packet: %v\n", conn.RemoteAddr().String(), packet)
responsePacket, err := handlePacket(packet)
log.Printf("client %v: response packet: %v\n", conn.RemoteAddr().String(), responsePacket)
if err != nil {
log.Printf("client %v: error processing request: %v\n", conn.RemoteAddr().String(), err)
break outer
}
err = responsePacket.Into(conn)
if err != nil {
log.Printf("client %v: error writing packet: %v\n", conn.RemoteAddr().String(), err)
break outer
}
case err := <-outErr:
if err == packet.PacketUnsupportedEncoding {
err := unsupportedEncodingErrorPacket.Into(conn)
log.Printf("client %v: error writing unsupported encoding packet: %v\n", conn.RemoteAddr().String(), err)
} else if err == packet.PacketUnsupportedType {
err := unsupportedTypeErrorPacket.Into(conn)
log.Printf("client %v: error writing unsupported type packet: %v\n", conn.RemoteAddr().String(), err)
} else {
log.Printf("client %v: internal error: %v\n", conn.RemoteAddr().String(), err)
}
break outer
case <-ctx.Done():
log.Printf("client %v: %v\n", conn.RemoteAddr().String(), ctx.Err())
break outer
}
}
}
func handlePacket(pkt packet.Packet) (packet.Packet, error) {
switch pkt.Type() {
case packet.TypeEko:
var request packet.EkoMessage
if err := pkt.DecodePayload(&request); err != nil {
return packet.Packet{}, fmt.Errorf("decode error: %v", err)
}
response := packet.EkoMessage{Message: "Eko \"" + request.Message + "\""}
encoder, err := packet.NewMsgPackEncoder(&response)
if err != nil {
return packet.Packet{}, fmt.Errorf("encode error: %v", err)
}
return packet.NewPacket(encoder), nil
case packet.TypeError:
return packet.Packet{}, errors.New("TODO: not implemented yet")
default:
assert.Unreachable("type should be checked for validity before handler, packet = %v", pkt.String())
return packet.Packet{}, nil
}
log.Printf("written %v bytes\n", bytesWritten)
}

View File

@@ -11,6 +11,9 @@ import (
"strconv"
"sync"
"syscall"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/pkg/assert"
)
const port = 7223
@@ -31,6 +34,8 @@ func Start() {
Certificates: []tls.Certificate{cert},
}
prepareConstPackets()
listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(port), tlsConfig)
if err != nil {
log.Fatalf("error starting server: %s", err)
@@ -67,3 +72,18 @@ func listen(listener net.Listener, wg *sync.WaitGroup) {
}
log.Printf("stopped listening on port %v...\n", port)
}
var unsupportedEncodingErrorPacket packet.Packet
var unsupportedTypeErrorPacket packet.Packet
func prepareConstPackets() {
message := packet.ErrorMessage{Error: packet.PacketUnsupportedEncoding.Error()}
encoder, err := packet.NewMsgPackEncoder(&message)
assert.NoError(err, "constant packets should not error")
unsupportedEncodingErrorPacket = packet.NewPacket(encoder)
message = packet.ErrorMessage{Error: packet.PacketUnsupportedType.Error()}
encoder, err = packet.NewMsgPackEncoder(&message)
assert.NoError(err, "constant packets should not error")
unsupportedTypeErrorPacket = packet.NewPacket(encoder)
}