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 - } - 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) } 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) +}