From 1997bc8a150b92783060cc7129e1ee72a761182b Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Sun, 20 Oct 2024 11:36:43 +0300 Subject: [PATCH] refactor: mid refactor --- cmd/client/client.log | 3 - cmd/client/{client.go => main.go} | 0 cmd/server/main.go | 28 +++ cmd/server/server.go | 7 - go.mod | 18 +- go.sum | 8 +- internal/client/client.go | 3 +- internal/client/ui.go | 1 + internal/packet/framer.go | 106 ---------- internal/packet/framer_test.go | 63 ------ internal/packet/packet.go | 149 ++++++++++---- internal/packet/packet_test.go | 58 ++++++ internal/server/handler.go | 122 ----------- internal/server/protocol.txt | 59 ++++++ internal/server/server.go | 331 +++++++++++++++++++++++++----- pkg/assert/assert.go | 2 +- pkg/snowflake/snowflake.go | 7 +- pkg/util/io.go | 57 ----- pkg/util/io_test.go | 59 ------ xclip | 1 + 20 files changed, 556 insertions(+), 526 deletions(-) delete mode 100644 cmd/client/client.log rename cmd/client/{client.go => main.go} (100%) create mode 100644 cmd/server/main.go delete mode 100644 cmd/server/server.go delete mode 100644 internal/packet/framer.go delete mode 100644 internal/packet/framer_test.go delete mode 100644 internal/server/handler.go create mode 100644 internal/server/protocol.txt delete mode 100644 pkg/util/io.go delete mode 100644 pkg/util/io_test.go create mode 100644 xclip diff --git a/cmd/client/client.log b/cmd/client/client.log deleted file mode 100644 index 6777b26..0000000 --- a/cmd/client/client.log +++ /dev/null @@ -1,3 +0,0 @@ -2024/10/17 16:45:33 client started, waiting for user input... -2024/10/17 16:45:33 established connection with server: 127.0.0.1:7223 -2024/10/17 16:45:33 sent request to server diff --git a/cmd/client/client.go b/cmd/client/main.go similarity index 100% rename from cmd/client/client.go rename to cmd/client/main.go diff --git a/cmd/server/main.go b/cmd/server/main.go new file mode 100644 index 0000000..e9cb4ad --- /dev/null +++ b/cmd/server/main.go @@ -0,0 +1,28 @@ +package main + +import ( + "context" + "log" + "os" + "os/signal" + "syscall" + + "github.com/kyren223/eko/internal/server" +) + +const port = 7223 + +func main() { + server := server.NewServer(port) + + ctx, cancel := context.WithCancel(context.Background()) + signalChan := make(chan os.Signal, 1) + signal.Notify(signalChan, syscall.SIGINT, syscall.SIGTERM) + go func() { + signal := <-signalChan + log.Println("signal:", signal.String()) + cancel() + }() + + server.ListenAndServe(ctx) +} diff --git a/cmd/server/server.go b/cmd/server/server.go deleted file mode 100644 index 168264a..0000000 --- a/cmd/server/server.go +++ /dev/null @@ -1,7 +0,0 @@ -package main - -import "github.com/kyren223/eko/internal/server" - -func main() { - server.Start() -} diff --git a/go.mod b/go.mod index 121360c..d4e63e5 100644 --- a/go.mod +++ b/go.mod @@ -2,12 +2,17 @@ module github.com/kyren223/eko go 1.23.2 -require github.com/vmihailenco/msgpack/v5 v5.4.1 +require ( + github.com/charmbracelet/bubbles v0.20.0 + github.com/charmbracelet/bubbletea v1.1.1 + github.com/charmbracelet/lipgloss v0.13.0 + github.com/vmihailenco/msgpack/v5 v5.4.1 +) require ( + github.com/ProtonMail/gopenpgp/v3 v3.0.0-beta.2-proton // indirect github.com/atotto/clipboard v0.1.4 // indirect github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect - github.com/charmbracelet/lipgloss v0.13.0 // indirect github.com/charmbracelet/x/ansi v0.2.3 // indirect github.com/charmbracelet/x/term v0.2.0 // indirect github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f // indirect @@ -19,13 +24,8 @@ require ( github.com/muesli/cancelreader v0.2.2 // indirect github.com/muesli/termenv v0.15.2 // indirect github.com/rivo/uniseg v0.4.7 // indirect + github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect golang.org/x/sync v0.8.0 // indirect golang.org/x/sys v0.24.0 // indirect - golang.org/x/text v0.3.8 // indirect -) - -require ( - github.com/charmbracelet/bubbles v0.20.0 - github.com/charmbracelet/bubbletea v1.1.1 - github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect + golang.org/x/text v0.14.0 // indirect ) diff --git a/go.sum b/go.sum index c3fb316..51f9466 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,7 @@ +github.com/MakeNowJust/heredoc v1.0.0 h1:cXCdzVdstXyiTqTvfqk9SDHpKNjxuom+DOlyEeQ4pzQ= +github.com/MakeNowJust/heredoc v1.0.0/go.mod h1:mG5amYoWBHf8vpLOuehzbGGw0EHxpZZ6lCpQ4fNJ8LE= +github.com/ProtonMail/gopenpgp/v3 v3.0.0-beta.2-proton h1:XFu8VgaGnb5MGOnwUr/l25HGLwfI/XFz12yTb3qhUYQ= +github.com/ProtonMail/gopenpgp/v3 v3.0.0-beta.2-proton/go.mod h1:TBpqWZ9IzA7g3TEzNA9Fwv/nA/eYpjcvYQBq+FX+tE4= github.com/atotto/clipboard v0.1.4 h1:EH0zSVneZPSuFR11BlR9YppQTVDbh5+16AmcJi4g1z4= github.com/atotto/clipboard v0.1.4/go.mod h1:ZY9tmq7sm5xIbd9bOK4onWV4S6X0u6GY7Vn0Yu86PYI= github.com/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiErDT4WkJ2k= @@ -22,8 +26,6 @@ github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWE github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mattn/go-localereader v0.0.1 h1:ygSAOl7ZXTx4RdPYinUpg6W99U8jWvWi9Ye2JC/oIi4= github.com/mattn/go-localereader v0.0.1/go.mod h1:8fBrzywKY7BI3czFoHkuzRoWE9C+EiG4R1k4Cjx5p88= -github.com/mattn/go-runewidth v0.0.15 h1:UNAjwbU9l54TA3KzvqLGxwWjHmMgBUVhBiTjelZgg3U= -github.com/mattn/go-runewidth v0.0.15/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w= github.com/mattn/go-runewidth v0.0.16 h1:E5ScNMtiwvlvB5paMFdw9p4kSQzbXFikJ5SQO6TULQc= github.com/mattn/go-runewidth v0.0.16/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w= github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 h1:ZK8zHtRHOkbHy6Mmr5D264iyp3TiX5OmNcI5cIARiQI= @@ -39,6 +41,7 @@ github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= github.com/stretchr/testify v1.6.1 h1:hDPOHmpOpP40lSULcqw7IrRb/u7w6RpDC9399XyoNd0= github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY= 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= @@ -51,5 +54,6 @@ golang.org/x/sys v0.24.0 h1:Twjiwq9dn6R1fQcyiK+wQyHWfaz/BJB+YIpzU/Cv3Xg= golang.org/x/sys v0.24.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/text v0.3.8 h1:nAL+RVCQ9uMn3vJZbV+MRnydTJFPf8qqY42YiA6MrqY= golang.org/x/text v0.3.8/go.mod h1:E6s5w1FMmriuDzIBO73fBruAKo1PCIq6d2Q6DHfQ8WQ= +golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/client/client.go b/internal/client/client.go index 99c53b3..5f42c8f 100644 --- a/internal/client/client.go +++ b/internal/client/client.go @@ -12,6 +12,7 @@ import ( "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/pkg/assert" + "github.com/vmihailenco/msgpack/v5" ) //go:embed server.crt @@ -107,7 +108,7 @@ func SendAndReceive(request packet.TypedMessage, response packet.TypedMessage) e select { case responsePacket := <-out: if err := responsePacket.DecodePayload(response); err != nil { - if responsePacket.Type() != packet.TypeError { + if responsePacket.Type() != packet.PacketError { return fmt.Errorf("error decoding response: %v", err) } var errorResponse packet.ErrorMessage diff --git a/internal/client/ui.go b/internal/client/ui.go index 839c0cd..9de322c 100644 --- a/internal/client/ui.go +++ b/internal/client/ui.go @@ -19,6 +19,7 @@ func startUI() { if _, err := p.Run(); err != nil { log.Println("charm ui error:", err) } + p. } type model struct { diff --git a/internal/packet/framer.go b/internal/packet/framer.go deleted file mode 100644 index 1a76ea6..0000000 --- a/internal/packet/framer.go +++ /dev/null @@ -1,106 +0,0 @@ -package packet - -import ( - "context" - "encoding/binary" - "errors" - "fmt" - "io" - - "github.com/kyren223/eko/pkg/assert" - "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) { - ch := make(chan Packet, framerPacketCapacity) - errCh := make(chan error) - - framer := packetFramer{ - buffer: make([]byte, 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:]) - 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 { - 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 fmt.Errorf("%w version=%v", PacketUnsupportedVersion, f.buffer[VERSION_OFFSET]) - } - - 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) - copy(packetBuffer, f.buffer[:fullLength]) - - f.len = uint16(copy(f.buffer, f.buffer[fullLength:f.len])) - - f.in <- Packet{packetBuffer} - } - return nil -} diff --git a/internal/packet/framer_test.go b/internal/packet/framer_test.go deleted file mode 100644 index 41a4396..0000000 --- a/internal/packet/framer_test.go +++ /dev/null @@ -1,63 +0,0 @@ -package packet - -import ( - "context" - "io" - "testing" - "time" -) - -type TestIoReader struct { - data []byte - len int -} - -func (r *TestIoReader) Read(data []byte) (int, error) { - if len(r.data) <= r.len { - // log.Println("RETURNING EOF:", len(data), r.len) - return 0, io.EOF - } - n := copy(data, r.data[r.len:]) - r.len += n - // log.Println("RETURNING N:", len(r.data), r.len) - return n, nil -} - -func (r *TestIoReader) start(msg TypedMessage, t *testing.T) { - encoder, err := NewMsgPackEncoder(msg) - if err != nil { - t.Errorf("encoding error: %v", err) - return - } - r.data = NewPacket(encoder).data - // log.Println("LEN:", len(r.data)) -} - -func TestPacketFramer(t *testing.T) { - reader := &TestIoReader{} - - // Long msg to test multiple - msg := EkoMessage{"Testing FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting Framer"} - reader.start(&msg, t) - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - out, outErr := RunFramer(ctx, reader) - - var message EkoMessage - select { - case packet := <-out: - if err := packet.DecodePayload(&message); err != nil { - t.Errorf("error decoding response: %v", err) - return - } - if msg.Message != message.Message { - t.Errorf("%v != %v", msg.Message, message.Message) - return - } - - case err := <-outErr: - t.Errorf("error receiving packet: %v", err) - return - } -} diff --git a/internal/packet/packet.go b/internal/packet/packet.go index 634643e..df3ee73 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -1,10 +1,13 @@ package packet import ( + "context" "encoding/binary" "encoding/json" + "errors" "fmt" "io" + "math" "github.com/vmihailenco/msgpack/v5" @@ -13,6 +16,13 @@ import ( type Encoding uint8 +const ( + EncodingJson Encoding = iota + EncodingMsgPack + EncodingUnused1 + EncodingUnused2 +) + func (e Encoding) String() string { switch e { case EncodingJson: @@ -37,26 +47,21 @@ func (e Encoding) IsSupported() bool { } } -const ( - EncodingJson Encoding = iota - EncodingMsgPack - EncodingUnused1 - EncodingUnused2 -) - type PacketType uint8 +const ( + PacketError PacketType = iota + PacketSendMessage + PacketMessages +) + func (t PacketType) String() string { switch t { - case TypeEko: - return "PacketTypeEko" - case TypeError: - return "PacketTypeError" - case TypeGetMessages: - return "PacketTypeGetMessages" - case TypeSendMessage: + case PacketError: + return "PacketError" + case PacketSendMessage: return "PacketTypeSendMessage" - case TypeMessages: + case PacketMessages: return "PacketTypeMessages" default: return fmt.Sprintf("PacketTypeInvalid(%v)", byte(t)) @@ -65,24 +70,16 @@ func (t PacketType) String() string { func (e PacketType) IsSupported() bool { switch e { - case TypeEko, TypeError, TypeGetMessages, TypeSendMessage, TypeMessages: + case PacketError, PacketSendMessage, PacketMessages: return true default: return false } } -const ( - TypeEko PacketType = iota - TypeError - TypeGetMessages - TypeSendMessage - TypeMessages -) - const ( VERSION = byte(1) - PACKET_MAX_SIZE = ^uint16(0) + PACKET_MAX_SIZE = math.MaxUint16 PAYLOAD_MAX_SIZE = PACKET_MAX_SIZE - HEADER_SIZE HEADER_SIZE = 4 VERSION_OFFSET = 0 @@ -92,7 +89,7 @@ const ( ) type PacketEncoder interface { - io.Reader + Payload() []byte Encoding() Encoding Type() PacketType } @@ -111,12 +108,11 @@ type Packet struct { } func NewPacket(encoder PacketEncoder) Packet { - data := make([]byte, PACKET_MAX_SIZE) + payload := encoder.Payload() + n := len(payload) + assert.Assert(0 <= n && n <= PAYLOAD_MAX_SIZE, "size of payload must be valid") - 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 := make([]byte, HEADER_SIZE+n) data[VERSION_OFFSET] = VERSION @@ -125,7 +121,11 @@ func NewPacket(encoder PacketEncoder) Packet { assert.Assert(encoding <= 3, "encoding exceeded allowed permutations encoding=%v", encoding) data[TYPE_OFFSET] = packetType | encoding<<6 - return Packet{data[:HEADER_SIZE+n]} + binary.BigEndian.PutUint16(data[LENGTH_OFFSET:], uint16(n)) + + copy(data[HEADER_SIZE:], payload) + + return Packet{data} } func (p Packet) Version() uint8 { @@ -133,7 +133,7 @@ func (p Packet) Version() uint8 { } func (p Packet) Type() PacketType { - return PacketType(p.data[TYPE_OFFSET] & 63) // 2^6-1 + return PacketType(p.data[TYPE_OFFSET] & 63) } func (p Packet) Encoding() Encoding { @@ -144,15 +144,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 bytes...]}", p.Version(), p.Encoding().String(), p.Type().String(), p.PayloadLength()) -} - -// 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) String() string { + return fmt.Sprintf("Packet(v%v %v %v [%v bytes...])", p.Version(), p.Encoding().String(), p.Type().String(), p.PayloadLength()) +} + +func (p Packet) Into(writer io.Writer) (int, error) { + return writer.Write(p.data) +} + 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()) @@ -167,12 +170,74 @@ func (p Packet) DecodePayload(v TypedMessage) error { case EncodingUnused2: return fmt.Errorf("unsupported encoding: %v", p.Encoding().String()) default: - assert.Unreachable("encoding from packet should always be valid encoding=%v", p.Encoding()) + assert.Never("encoding from packet should always be valid encoding=%v", p.Encoding()) return nil } } -func (p Packet) Into(writer io.Writer) error { - _, err := writer.Write(p.data[:len(p.data)]) - return err +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 + Out chan Packet +} + +func NewFramer(ctx context.Context) packetFramer { + return packetFramer{ + Out: make(chan Packet, 10), + } +} + +func (f *packetFramer) Push(ctx context.Context, data []byte) error { + f.buffer = append(f.buffer, data...) + + for { + packet, err := f.parse() + if packet == nil || err != nil { + return err + } + select { + case f.Out <- *packet: + case <-ctx.Done(): + return ctx.Err() + } + } +} + +func (f *packetFramer) parse() (*Packet, error) { + if len(f.buffer) < HEADER_SIZE { + return nil, nil + } + + if f.buffer[VERSION_OFFSET] != VERSION { + return nil, fmt.Errorf("%w version=%v", PacketUnsupportedVersion, f.buffer[VERSION_OFFSET]) + } + + encoding := Encoding(f.buffer[ENCODING_OFFSET] >> 6) + if !encoding.IsSupported() { + return nil, PacketUnsupportedEncoding + } + + packetType := PacketType(f.buffer[TYPE_OFFSET] & 63) + if !packetType.IsSupported() { + return nil, PacketUnsupportedType + } + + length := binary.BigEndian.Uint16(f.buffer[LENGTH_OFFSET:]) + if len(f.buffer)-HEADER_SIZE < int(length) { + // Wait for more data to arrive + return nil, nil + } + + fullLength := HEADER_SIZE + length + packetBuffer := make([]byte, fullLength) + copy(packetBuffer, f.buffer[:fullLength]) + copy(f.buffer, f.buffer[fullLength:]) + f.buffer = f.buffer[:len(f.buffer)-int(fullLength)] + + return &Packet{packetBuffer}, nil } diff --git a/internal/packet/packet_test.go b/internal/packet/packet_test.go index 20370af..a523d59 100644 --- a/internal/packet/packet_test.go +++ b/internal/packet/packet_test.go @@ -1,7 +1,10 @@ package packet import ( + "context" + "io" "testing" + "time" "github.com/vmihailenco/msgpack/v5" ) @@ -48,3 +51,58 @@ func TestPacketMsgPackEncoding(t *testing.T) { return } } + +type TestIoReader struct { + data []byte + len int +} + +func (r *TestIoReader) Read(data []byte) (int, error) { + if len(r.data) <= r.len { + // log.Println("RETURNING EOF:", len(data), r.len) + return 0, io.EOF + } + n := copy(data, r.data[r.len:]) + r.len += n + // log.Println("RETURNING N:", len(r.data), r.len) + return n, nil +} + +func (r *TestIoReader) start(msg TypedMessage, t *testing.T) { + encoder, err := NewMsgPackEncoder(msg) + if err != nil { + t.Errorf("encoding error: %v", err) + return + } + r.data = NewPacket(encoder).data + // log.Println("LEN:", len(r.data)) +} + +func TestPacketFramer(t *testing.T) { + reader := &TestIoReader{} + + // Long msg to test multiple + msg := EkoMessage{"Testing FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting FramerTesting Framer"} + reader.start(&msg, t) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + out, outErr := RunFramer(ctx, reader) + + var message EkoMessage + select { + case packet := <-out: + if err := packet.DecodePayload(&message); err != nil { + t.Errorf("error decoding response: %v", err) + return + } + if msg.Message != message.Message { + t.Errorf("%v != %v", msg.Message, message.Message) + return + } + + case err := <-outErr: + t.Errorf("error receiving packet: %v", err) + return + } +} diff --git a/internal/server/handler.go b/internal/server/handler.go deleted file mode 100644 index 0ab5c33..0000000 --- a/internal/server/handler.go +++ /dev/null @@ -1,122 +0,0 @@ -package server - -import ( - "context" - "errors" - "fmt" - "log" - "net" - "strings" - "sync" - "time" - - "github.com/kyren223/eko/internal/data" - "github.com/kyren223/eko/internal/packet" - "github.com/kyren223/eko/pkg/assert" - "github.com/kyren223/eko/pkg/snowflake" -) - -func handleConnection(conn net.Conn, wg *sync.WaitGroup) { - log.Println("accepted client:", conn.RemoteAddr().String()) - defer log.Println("disconnected client:", conn.RemoteAddr().String()) - defer conn.Close() - defer wg.Done() - - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() - out, outErr := packet.RunFramer(ctx, conn) - -outer: - for { - select { - case packet, ok := <-out: - if !ok { - break outer - } - 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 if err != nil { - 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) { - var response packet.TypedMessage - 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 + "\""} - case packet.TypeSendMessage: - var request packet.SendMessageMessage - if err := pkt.DecodePayload(&request); err != nil { - return packet.Packet{}, fmt.Errorf("decode error: %v", err) - } - - content := strings.TrimSpace(request.Content) - if content == "" { - response = &packet.ErrorMessage{Error: "content must not be blank"} - break - } - - message := data.Message{ - Id: node.Generate(), - SenderId: node.Generate(), - FrequencyId: node.Generate(), - NetworkId: node.Generate(), - Contents: content, - } - messages = append(messages, message) - - response = &packet.EkoMessage{Message: "Eko OK"} - case packet.TypeGetMessages: - var request packet.GetMessagesMessage - if err := pkt.DecodePayload(&request); err != nil { - return packet.Packet{}, fmt.Errorf("decode error: %v", err) - } - - response = &packet.MessagesMessage{Messages: messages} - default: - return packet.Packet{}, errors.New("TODO: not implemented yet") - } - - assert.NotNil(response, "response must always be set") - encoder, err := packet.NewMsgPackEncoder(response) - if err != nil { - return packet.Packet{}, fmt.Errorf("encode error: %v", err) - } - return packet.NewPacket(encoder), nil -} - -var ( - node = snowflake.NewNode(1) - messages []data.Message -) diff --git a/internal/server/protocol.txt b/internal/server/protocol.txt new file mode 100644 index 0000000..9a0070d --- /dev/null +++ b/internal/server/protocol.txt @@ -0,0 +1,59 @@ +# Eko Protocol V1 + +## 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 ... | ++-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + +Order of bytes is from left to right, top to bottom. +The first byte is always the version, any bytes after it +depend on the specific value of the first byte. + +- Encoding: 0-3, determines the way the payload was encoded + - 0: JSON + - 1: MsgPack + - 2: Reserved for future use + - 3: Reserved for future use +- Type: 0-63, determines the type ("schema"), of the payload +- Payload Length: 0-65531, determines how long the payload is in bytes +- Payload: 0 to 65531 bytes long, depending on the payload size (~64kb) + +## Handshake + +The first time a connection is established, the following packets are exchanged. + +- Server sends a special 1-byte for version then 32-byte the challenge nonce packet +- Client sends a Challenge Response packet with: + * version (1-byte long) + - the client's ed25519 public key (32-bytes long) + - the client's ed25519 signature for the server-given nonce (64-bytes long) + +After the handshake the server may close the connection, +for example due to an invalid signature. + +## Error handling + +The server may abruptly close a connection in these cases: + +- After the initial handshake +- After any response + A server may not close the connection if it received a request, it must first response then close. + +The client may abruptly close a connection at any time + +### Malformed Packets + +- unsupported/invalid version: connection can be closed immediately +- unsupported encoding: server must respond with an error type, may use any encoding, client may close the connection +- unknown type: server must respond with an error, client may close the connection +- malformed paylod: server must respond with an error, client may close the connection + +For application errors such as a client asking to send a message in a non-existent Frequency, +the server must respond with an error packet. +For internal errors such as database failure, the server must respond, it may choose to +disclose as much information as it wants, or just say "internal server error". diff --git a/internal/server/server.go b/internal/server/server.go index 89434d3..933ad29 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1,82 +1,52 @@ package server import ( + "context" + "crypto/rand" "crypto/tls" _ "embed" "errors" + "fmt" + "io" "log" "net" - "os" - "os/signal" "strconv" + "strings" "sync" - "syscall" + "time" + "github.com/kyren223/eko/internal/data" "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/pkg/assert" + "github.com/kyren223/eko/pkg/snowflake" ) -const port = 7223 - //go:embed server.crt var certPEM []byte //go:embed server.key var keyPEM []byte -func Start() { +var ( + nodeId int64 = 0 + tlsConfig *tls.Config + + unsupportedEncodingErrorPacket packet.Packet + unsupportedTypeErrorPacket packet.Packet +) + +var ErrClosedNilListener error = errors.New("server: close on nil listener") + +func init() { cert, err := tls.X509KeyPair(certPEM, keyPEM) if err != nil { log.Fatalln("error loading certificate:", err) } - tlsConfig := &tls.Config{ + tlsConfig = &tls.Config{ Certificates: []tls.Certificate{cert}, } - prepareConstPackets() - - listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(port), tlsConfig) - if err != nil { - log.Fatalf("error starting server: %s", err) - } - defer listener.Close() - - signalChan := make(chan os.Signal, 1) - signal.Notify(signalChan, syscall.SIGINT, syscall.SIGTERM) - go handleInterrupt(listener, signalChan) - - var wg sync.WaitGroup - listen(listener, &wg) - wg.Wait() -} - -func handleInterrupt(listener net.Listener, stopChan <-chan os.Signal) { - signal := <-stopChan - log.Println("signal:", signal.String()) - listener.Close() -} - -func listen(listener net.Listener, wg *sync.WaitGroup) { - log.Printf("started listening on port %v...\n", port) - for { - conn, err := listener.Accept() - if err != nil { - if !errors.Is(err, net.ErrClosed) { - log.Println("error accepting connection:", err) - } - break - } - wg.Add(1) - go handleConnection(conn, wg) - } - 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") @@ -87,3 +57,264 @@ func prepareConstPackets() { assert.NoError(err, "constant packets should not error") unsupportedTypeErrorPacket = packet.NewPacket(encoder) } + +type server struct { + node *snowflake.Node + port uint16 +} + +func NewServer(port uint16) server { + assert.Assert(nodeId <= snowflake.NodeMax, "maximum amount of servers reached %v", snowflake.NodeMax) + node := snowflake.NewNode(nodeId) + nodeId++ + + return server{ + node: node, + port: port, + } +} + +func (s *server) ListenAndServe(ctx context.Context) { + listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.port)), tlsConfig) + if err != nil { + log.Fatalf("error starting server: %s", err) + } + defer listener.Close() + go func() { + <-ctx.Done() + listener.Close() + }() + + log.Printf("started listening on port %v...\n", s.port) + var wg sync.WaitGroup + for { + conn, err := listener.Accept() + if err != nil { + if !errors.Is(err, net.ErrClosed) { + log.Println("error accepting connection:", err) + } + break + } + wg.Add(1) + go func() { + handleConnection(ctx, conn) + wg.Done() + }() + } + log.Printf("stopped listening on port %v\n", s.port) + + log.Println("waiting for all active connections to close...") + wg.Wait() + log.Println("server shutdown complete") +} + +func handleConnection(ctx context.Context, conn net.Conn) { + addr, ok := conn.RemoteAddr().(*net.TCPAddr) + assert.Assert(ok, "getting tcp address should never fail as we are using tcp connections") + + writeQueue := make(chan packet.Packet, 10) + session := newSession(addr, writeQueue) + nonce := session.Challenge() + + ctx = newContext(ctx, session) + framer := packet.NewFramer(ctx) + + log.Println(addr, "accepted") + + defer func() { + conn.Close() + log.Println(addr, "disconnected") + }() + + go func() { + var mu sync.Mutex + for { + packet, ok := <-writeQueue + if !ok { + break + } + mu.Lock() + packet.Into(conn) + mu.Unlock() + } + }() + + go func() { + for { + request, ok := <-framer.Out + if !ok { + break + } + response := processPacket(ctx, request) + writeQueue <- response + } + }() + + buffer := make([]byte, 512) + for { + n, err := conn.Read(buffer) + if err != nil { + if !errors.Is(err, io.EOF) { + log.Println(addr, "read error:", err) + } + break + } + + err = framer.Push(ctx, buffer[:n]) + if err != nil { + // Wrap err and send to client then break + } + } +} + +func _handleConnection(ctx context.Context, conn net.Conn) { + addr, ok := conn.RemoteAddr().(*net.TCPAddr) + assert.Assert(ok, "getting tcp address should be valid") + + log.Println(addr, "accepted") + defer log.Println(addr, "disconnected") + defer conn.Close() + + // ctx = newContext(ctx, addr) + // // TODO: consider adding timeout/deadline to ctx? + + out, outErr := packet.RunFramer(ctx, conn) +outer: + for { + select { + case packet, ok := <-out: + if !ok { + break outer + } + 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 if err != nil { + 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 + } + } +} + +type Session struct { + Addr *net.TCPAddr + WriteQueue <-chan packet.Packet + + mu sync.Mutex + challenge []byte + issuedTime time.Time +} + +func newSession(addr *net.TCPAddr, writeQueue <-chan packet.Packet) *Session { + session := &Session{ + Addr: addr, + WriteQueue: writeQueue, + challenge: make([]byte, 32), // Recommended nonce size + } + session.Challenge() // Make sure an initial nonce is generated + return session +} + +func (s *Session) Challenge() []byte { + s.mu.Lock() + defer s.mu.Unlock() + if time.Since(s.issuedTime) > time.Minute { + s.issuedTime = time.Now() + _, err := rand.Read(s.challenge) + assert.NoError(err, "random should always produce a value") + } + return s.challenge +} + +type key struct{} + +var sessKey key + +func newContext(ctx context.Context, sess *Session) context.Context { + return context.WithValue(ctx, sessKey, sess) +} + +func FromContext(ctx context.Context) (*Session, bool) { + sess, ok := ctx.Value(sessKey).(*Session) + return sess, ok +} + +// TODO: Move everything below this to somewhere else + +func handlePacket(pkt packet.Packet) (packet.Packet, error) { + var response packet.TypedMessage + switch pkt.Type() { + case packet.PacketEko: + 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 + "\""} + case packet.PacketSendMessage: + var request packet.SendMessageMessage + if err := pkt.DecodePayload(&request); err != nil { + return packet.Packet{}, fmt.Errorf("decode error: %v", err) + } + + content := strings.TrimSpace(request.Content) + if content == "" { + response = &packet.ErrorMessage{Error: "content must not be blank"} + break + } + + message := data.Message{ + Id: node.Generate(), + SenderId: node.Generate(), + FrequencyId: node.Generate(), + NetworkId: node.Generate(), + Contents: content, + } + messages = append(messages, message) + + response = &packet.EkoMessage{Message: "Eko OK"} + case packet.PacketGetMessages: + var request packet.GetMessagesMessage + if err := pkt.DecodePayload(&request); err != nil { + return packet.Packet{}, fmt.Errorf("decode error: %v", err) + } + + response = &packet.MessagesMessage{Messages: messages} + default: + return packet.Packet{}, errors.New("TODO: not implemented yet") + } + + assert.NotNil(response, "response must always be set") + encoder, err := packet.NewMsgPackEncoder(response) + if err != nil { + return packet.Packet{}, fmt.Errorf("encode error: %v", err) + } + return packet.NewPacket(encoder), nil +} + +var ( + node = snowflake.NewNode(1) + messages []data.Message +) diff --git a/pkg/assert/assert.go b/pkg/assert/assert.go index 6c4a4f5..7281a95 100644 --- a/pkg/assert/assert.go +++ b/pkg/assert/assert.go @@ -14,7 +14,7 @@ func NoError(err error, message string, a ...any) { } } -func Unreachable(message string, a ...any) { +func Never(message string, a ...any) { log.Fatalf(message+"\n", a...) } diff --git a/pkg/snowflake/snowflake.go b/pkg/snowflake/snowflake.go index 5b58cbb..bb4a07b 100644 --- a/pkg/snowflake/snowflake.go +++ b/pkg/snowflake/snowflake.go @@ -13,11 +13,10 @@ const ( // Epoch is set to the twitter snowflake epoch of Nov 04 2010 01:42:54 UTC in milliseconds // TODO: change this to eko epoch when eko is production ready Epoch int64 = 1288834974657 - nodeBits = 10 stepBits = 12 - nodeMax = 1<