mirror of
https://github.com/Kyren223/eko.git
synced 2026-08-28 10:21:31 +00:00
feat: communication using custom packet protocol
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
17
internal/packet/messages.go
Normal file
17
internal/packet/messages.go
Normal 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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
50
internal/packet/packet_test.go
Normal file
50
internal/packet/packet_test.go
Normal 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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user