mirror of
https://github.com/Kyren223/eko.git
synced 2026-08-29 10:51:31 +00:00
refactor: mid refactor
This commit is contained in:
@@ -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
28
cmd/server/main.go
Normal 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)
|
||||
}
|
||||
@@ -1,7 +0,0 @@
|
||||
package main
|
||||
|
||||
import "github.com/kyren223/eko/internal/server"
|
||||
|
||||
func main() {
|
||||
server.Start()
|
||||
}
|
||||
18
go.mod
18
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
|
||||
)
|
||||
|
||||
8
go.sum
8
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=
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -19,6 +19,7 @@ func startUI() {
|
||||
if _, err := p.Run(); err != nil {
|
||||
log.Println("charm ui error:", err)
|
||||
}
|
||||
p.
|
||||
}
|
||||
|
||||
type model struct {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
)
|
||||
59
internal/server/protocol.txt
Normal file
59
internal/server/protocol.txt
Normal 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".
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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...)
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user