refactor: mid refactor

This commit is contained in:
Kyren223
2024-10-20 11:36:43 +03:00
parent 6e6ca1d7a1
commit 1997bc8a15
20 changed files with 556 additions and 526 deletions

View File

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

28
cmd/server/main.go Normal file
View File

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

View File

@@ -1,7 +0,0 @@
package main
import "github.com/kyren223/eko/internal/server"
func main() {
server.Start()
}

18
go.mod
View File

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

8
go.sum
View File

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

View File

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

View File

@@ -19,6 +19,7 @@ func startUI() {
if _, err := p.Run(); err != nil {
log.Println("charm ui error:", err)
}
p.
}
type model struct {

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -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".

View File

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

View File

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

View File

@@ -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<<nodeBits - 1
nodeMask = nodeMax << stepBits
NodeMax = 1<<nodeBits - 1
nodeMask = NodeMax << stepBits
stepMask = 1<<stepBits - 1
timeShift = nodeBits + stepBits
nodeShift = stepBits
@@ -35,7 +34,7 @@ type ID int64
func NewNode(node int64) *Node {
assert.Assert(nodeBits+stepBits <= 22, "node and step bits must add up to 22 or less")
assert.Assert(0 <= node && node <= nodeMax, "node and step bits must add up to 22 or less")
assert.Assert(0 <= node && node <= NodeMax, "node and step bits must add up to 22 or less")
// Credit to https://github.com/bwmarrin/snowflake
currentTime := time.Now()

View File

@@ -1,57 +0,0 @@
package util
import (
"context"
"io"
)
const bufferSize = 512
type ChannelReader struct {
Out <-chan []byte
Err <-chan error
}
func NewChannelReader(ctx context.Context, reader io.Reader) ChannelReader {
outCh := make(chan []byte)
errCh := make(chan error)
go func(in chan<- []byte, inErr chan<- error) {
defer close(in)
defer close(inErr)
buffer := make([]byte, 2 * bufferSize)
isD1 := true
d1 := buffer[:bufferSize]
d2 := buffer[bufferSize:]
outer:
for {
select {
case <-ctx.Done():
inErr <- ctx.Err()
break outer
default:
var data []byte
if isD1 {
data = d1
} else {
data = d2
}
isD1 = !isD1
n, err := reader.Read(data)
if err != nil {
if err != io.EOF {
inErr <- err
}
break outer
}
in <- data[:n]
}
}
}(outCh, errCh)
return ChannelReader{
Out: outCh,
Err: errCh,
}
}

View File

@@ -1,59 +0,0 @@
package util
import (
"bytes"
"context"
"testing"
"time"
)
func TestChannelReader(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
b := []byte{1, 2, 3, 4, 5, 6, 7, 8}
reader := NewChannelReader(ctx, bytes.NewReader(b))
select {
case data := <-reader.Out:
if !bytes.Equal(b, data) {
t.Errorf("%v != %v", b, data)
}
case err := <-reader.Err:
t.Errorf("reading err: %v", err)
}
}
func TestChannelReaderMultiPartRead(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
b := []byte("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")
b1 := b[:512]
b2 := b[512:]
reader := NewChannelReader(ctx, bytes.NewReader(b))
counter := 0
outer:
for {
select {
case data := <-reader.Out:
if counter == 0 {
if !bytes.Equal(b1, data) {
t.Errorf("%v != %v", b, data)
}
} else {
if !bytes.Equal(b2[:len(data)], data) {
t.Errorf("%v != %v", b, data)
}
}
counter++
if counter == 2 {
break outer
}
case err := <-reader.Err:
t.Errorf("reading err: %v", err)
break outer
}
}
}

1
xclip Normal file
View File

@@ -0,0 +1 @@
test