summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--internal/client/client.go32
-rw-r--r--internal/messages/eko.go19
-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
-rw-r--r--internal/server/handler.go78
-rw-r--r--internal/server/server.go20
9 files changed, 204 insertions, 62 deletions
diff --git a/internal/client/client.go b/internal/client/client.go
index 9aee12e..80b29a0 100644
--- a/internal/client/client.go
+++ b/internal/client/client.go
@@ -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
}
diff --git a/internal/messages/eko.go b/internal/messages/eko.go
deleted file mode 100644
index d8eaa86..0000000
--- a/internal/messages/eko.go
+++ /dev/null
@@ -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
-}
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
+ }
+}
diff --git a/internal/server/handler.go b/internal/server/handler.go
index 83d3e6e..61ecbae 100644
--- a/internal/server/handler.go
+++ b/internal/server/handler.go
@@ -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
+ 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())
+
+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
+ }
}
- request := string(buffer[:bytesRead])
- log.Printf("Read %v bytes: %v\n", bytesRead, request)
+}
+
+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
- response := []byte(fmt.Sprintf("Eko \"%v\"", request))
- bytesWritten, err := conn.Write(response)
- if err != nil {
- log.Println("failed writing response:", err)
- return
+ 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)
}
diff --git a/internal/server/server.go b/internal/server/server.go
index 48b31f4..89434d3 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -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)
+}