From 6f83f199f8311bf7782da59bc2e07df38073e65a Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Sun, 20 Oct 2024 20:12:12 +0300 Subject: refactor: finished refactoring server-side code --- internal/packet/encoders.go | 29 +++-- internal/packet/messages.go | 33 ++---- internal/packet/packet.go | 14 +-- internal/packet/packet_test.go | 108 ------------------ internal/server/server.go | 246 +++++++++++++++++++++++------------------ 5 files changed, 170 insertions(+), 260 deletions(-) delete mode 100644 internal/packet/packet_test.go (limited to 'internal') diff --git a/internal/packet/encoders.go b/internal/packet/encoders.go index e70d0e8..f5841ce 100644 --- a/internal/packet/encoders.go +++ b/internal/packet/encoders.go @@ -1,10 +1,9 @@ package packet import ( - "bytes" "encoding/json" - "io" + "github.com/kyren223/eko/pkg/assert" "github.com/vmihailenco/msgpack/v5" ) @@ -13,7 +12,7 @@ type TypedMessage interface { } type defaultPacketEncoder struct { - io.Reader + data []byte encoding Encoding packetType PacketType } @@ -26,28 +25,28 @@ func (e defaultPacketEncoder) Type() PacketType { return e.packetType } -func NewJsonEncoder(message TypedMessage) (PacketEncoder, error) { +func (e defaultPacketEncoder) Payload() []byte { + return e.data +} + +func NewJsonEncoder(message TypedMessage) PacketEncoder { data, err := json.Marshal(message) - if err != nil { - return nil, err - } + assert.NoError(err, "encoding a message with JSON should never fail") return defaultPacketEncoder{ - Reader: bytes.NewReader(data), + data: data, encoding: EncodingJson, packetType: message.Type(), - }, nil + } } -func NewMsgPackEncoder(message TypedMessage) (PacketEncoder, error) { +func NewMsgPackEncoder(message TypedMessage) PacketEncoder { data, err := msgpack.Marshal(message) - if err != nil { - return nil, err - } + assert.NoError(err, "encoding a message with msg pack should never fail") return defaultPacketEncoder{ - Reader: bytes.NewReader(data), + data: data, encoding: EncodingMsgPack, packetType: message.Type(), - }, nil + } } diff --git a/internal/packet/messages.go b/internal/packet/messages.go index b1a6598..c561079 100644 --- a/internal/packet/messages.go +++ b/internal/packet/messages.go @@ -4,43 +4,34 @@ import ( "github.com/kyren223/eko/internal/data" ) -type EkoMessage struct { - Message string `msgpack:"message"` -} - -func (m *EkoMessage) Type() PacketType { - return TypeEko -} - type ErrorMessage struct { Error string `msgpack:"error"` } -func (m *ErrorMessage) Type() PacketType { - return TypeError +func NewOkMessage() *ErrorMessage { + return &ErrorMessage{} } -type GetMessagesMessage struct { - Since *int64 - UpTo *int64 +func (m *ErrorMessage) Type() PacketType { + return PacketError } -func (m *GetMessagesMessage) Type() PacketType { - return TypeGetMessages +func (m *ErrorMessage) IsOk() bool { + return m.Error == "" } -type SendMessageMessage struct { +type SendMessage struct { Content string } -func (m *SendMessageMessage) Type() PacketType { - return TypeSendMessage +func (m *SendMessage) Type() PacketType { + return PacketSendMessage } -type MessagesMessage struct { +type Messages struct { Messages []data.Message } -func (m *MessagesMessage) Type() PacketType { - return TypeMessages +func (m *Messages) Type() PacketType { + return PacketMessages } diff --git a/internal/packet/packet.go b/internal/packet/packet.go index df3ee73..e9d5d94 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -60,11 +60,11 @@ func (t PacketType) String() string { case PacketError: return "PacketError" case PacketSendMessage: - return "PacketTypeSendMessage" + return "PacketSendMessage" case PacketMessages: - return "PacketTypeMessages" + return "PacketMessages" default: - return fmt.Sprintf("PacketTypeInvalid(%v)", byte(t)) + return fmt.Sprintf("PacketInvalidType(%v)", byte(t)) } } @@ -110,15 +110,15 @@ type Packet struct { func NewPacket(encoder PacketEncoder) Packet { payload := encoder.Payload() n := len(payload) - assert.Assert(0 <= n && n <= PAYLOAD_MAX_SIZE, "size of payload must be valid") + assert.Assert(0 <= n && n <= PAYLOAD_MAX_SIZE, "size of payload must be valid", "size", n) data := make([]byte, HEADER_SIZE+n) data[VERSION_OFFSET] = VERSION packetType, encoding := byte(encoder.Type()), byte(encoder.Encoding()) - assert.Assert(packetType <= 63, "packet type exceeded allowed size type=%v", packetType) - assert.Assert(encoding <= 3, "encoding exceeded allowed permutations encoding=%v", encoding) + assert.Assert(packetType <= 63, "packet type exceeded allowed size", "type", packetType) + assert.Assert(encoding <= 3, "encoding exceeded allowed size", "encoding", encoding) data[TYPE_OFFSET] = packetType | encoding<<6 binary.BigEndian.PutUint16(data[LENGTH_OFFSET:], uint16(n)) @@ -170,7 +170,7 @@ func (p Packet) DecodePayload(v TypedMessage) error { case EncodingUnused2: return fmt.Errorf("unsupported encoding: %v", p.Encoding().String()) default: - assert.Never("encoding from packet should always be valid encoding=%v", p.Encoding()) + assert.Never("encoding from packet should always be valid", "encoding", p.Encoding()) return nil } } diff --git a/internal/packet/packet_test.go b/internal/packet/packet_test.go deleted file mode 100644 index a523d59..0000000 --- a/internal/packet/packet_test.go +++ /dev/null @@ -1,108 +0,0 @@ -package packet - -import ( - "context" - "io" - "testing" - "time" - - "github.com/vmihailenco/msgpack/v5" -) - -func TestMsgPackEncoding(t *testing.T) { - request := EkoMessage{Message: "test"} - data, err := msgpack.Marshal(&request) - if err != nil { - t.Errorf("encoding error: %v", err) - return - } - var response EkoMessage - err = msgpack.Unmarshal(data, &response) - if err != nil { - t.Errorf("decoding error: %v", err) - return - } - if request.Message != response.Message { - t.Errorf("%v != %v", request.Message, response.Message) - return - } -} - -func TestPacketMsgPackEncoding(t *testing.T) { - request := EkoMessage{Message: "test"} - encoder1, err := NewMsgPackEncoder(&request) - encoder2, _ := NewMsgPackEncoder(&request) - if err != nil { - t.Errorf("encoding error: %v", err) - return - } - encodedBytes := make([]byte, PACKET_MAX_SIZE) - n, _ := encoder2.Read(encodedBytes[HEADER_SIZE:]) - - packet := NewPacket(encoder1) - var response EkoMessage - err = packet.DecodePayload(&response) - if err != nil { - t.Errorf("decoding error: %#v: packet: %v encoder: %v", err, packet.Payload(), encodedBytes[HEADER_SIZE:HEADER_SIZE+n]) - return - } - if request.Message != response.Message { - t.Errorf("%v != %v", request.Message, response.Message) - return - } -} - -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/server.go b/internal/server/server.go index 933ad29..4ac8969 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -2,6 +2,7 @@ package server import ( "context" + "crypto/ed25519" "crypto/rand" "crypto/tls" _ "embed" @@ -10,6 +11,7 @@ import ( "io" "log" "net" + "os" "strconv" "strings" "sync" @@ -48,44 +50,64 @@ func init() { } message := packet.ErrorMessage{Error: packet.PacketUnsupportedEncoding.Error()} - encoder, err := packet.NewMsgPackEncoder(&message) - assert.NoError(err, "constant packets should not error") - unsupportedEncodingErrorPacket = packet.NewPacket(encoder) + unsupportedEncodingErrorPacket = packet.NewPacket(packet.NewMsgPackEncoder(&message)) message = packet.ErrorMessage{Error: packet.PacketUnsupportedType.Error()} - encoder, err = packet.NewMsgPackEncoder(&message) - assert.NoError(err, "constant packets should not error") - unsupportedTypeErrorPacket = packet.NewPacket(encoder) + unsupportedTypeErrorPacket = packet.NewPacket(packet.NewMsgPackEncoder(&message)) } type server struct { - node *snowflake.Node - port uint16 + Node *snowflake.Node + Port uint16 + sessions map[snowflake.ID]*Session + sessMu sync.RWMutex } func NewServer(port uint16) server { - assert.Assert(nodeId <= snowflake.NodeMax, "maximum amount of servers reached %v", snowflake.NodeMax) + assert.Assert(nodeId <= snowflake.NodeMax, "maximum amount of servers reached") node := snowflake.NewNode(nodeId) nodeId++ return server{ - node: node, - port: port, + Node: node, + Port: port, + sessions: map[snowflake.ID]*Session{}, } } +func (s *server) AddSession(session *Session) { + s.sessMu.Lock() + defer s.sessMu.Unlock() + s.sessions[session.ID] = session +} + +func (s *server) RemoveSession(id snowflake.ID) { + s.sessMu.Lock() + defer s.sessMu.Unlock() + delete(s.sessions, id) +} + +func (s *server) Session(id snowflake.ID) (*Session, bool) { + s.sessMu.RLock() + defer s.sessMu.RUnlock() + session, ok := s.sessions[id] + return session, ok +} + func (s *server) ListenAndServe(ctx context.Context) { - listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.port)), tlsConfig) + listener, err := tls.Listen("tcp4", ":"+strconv.Itoa(int(s.Port)), tlsConfig) if err != nil { log.Fatalf("error starting server: %s", err) } + + assert.AddFlush(listener) defer listener.Close() go func() { <-ctx.Done() listener.Close() }() - log.Printf("started listening on port %v...\n", s.port) + log.Printf("started listening on port %v...\n", s.Port) var wg sync.WaitGroup for { conn, err := listener.Accept() @@ -97,46 +119,61 @@ func (s *server) ListenAndServe(ctx context.Context) { } wg.Add(1) go func() { - handleConnection(ctx, conn) + handleConnection(ctx, conn, s) wg.Done() }() } - log.Printf("stopped listening on port %v\n", s.port) + 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) { +func handleConnection(ctx context.Context, conn net.Conn, server *server) { 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() + log.Println(addr, "accepted") + nonce := [32]byte{} + _, err := rand.Read(nonce[:]) + assert.NoError(err, "random should always produce a value") + pubKey, err := handleAuth(conn, nonce[:]) + if err != nil { + log.Println(addr, err) + conn.Close() + log.Println(addr, "disconnected") + return + } + + // TODO: replace this with DB query for id + id := server.Node.Generate() + session := newSession(server, addr, id, pubKey) + server.AddSession(session) ctx = newContext(ctx, session) framer := packet.NewFramer(ctx) - log.Println(addr, "accepted") - defer func() { conn.Close() + close(framer.Out) + server.RemoveSession(session.ID) + close(session.WriteQueue) log.Println(addr, "disconnected") }() go func() { - var mu sync.Mutex for { - packet, ok := <-writeQueue + packet, ok := <-session.WriteQueue if !ok { break } - mu.Lock() - packet.Into(conn) - mu.Unlock() + if _, err := packet.Into(conn); err != nil { + log.Println(addr, err) + break + } } + session.WriteQueue = nil }() go func() { @@ -146,91 +183,97 @@ func handleConnection(ctx context.Context, conn net.Conn) { break } response := processPacket(ctx, request) - writeQueue <- response + if session.WriteQueue == nil { + break + } + session.WriteQueue <- response } }() buffer := make([]byte, 512) for { + conn.SetReadDeadline(time.Now().Add(time.Second)) n, err := conn.Read(buffer) - if err != nil { + deadlineExceeded := errors.Is(err, os.ErrDeadlineExceeded) + if err != nil && !deadlineExceeded { if !errors.Is(err, io.EOF) { log.Println(addr, "read error:", err) } break } + if ctx.Err() != nil { + log.Println(addr, ctx.Err()) + break + } + err = framer.Push(ctx, buffer[:n]) if err != nil { - // Wrap err and send to client then break + if ctx.Err() != nil { + log.Println(addr, ctx.Err()) + } else { + // TODO: Wrap err and send to client then break + } + break } } } -func _handleConnection(ctx context.Context, conn net.Conn) { - addr, ok := conn.RemoteAddr().(*net.TCPAddr) - assert.Assert(ok, "getting tcp address should be valid") +func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { + conn.SetDeadline(time.Now().Add(time.Second * 5)) - log.Println(addr, "accepted") - defer log.Println(addr, "disconnected") - defer conn.Close() + challengePacket := make([]byte, len(nonce)+1) + challengePacket[0] = packet.VERSION + copy(challengePacket[1:], nonce) - // ctx = newContext(ctx, addr) - // // TODO: consider adding timeout/deadline to ctx? + _, err := conn.Write(challengePacket) + if err != nil { + return nil, fmt.Errorf("error writing challenge: %w", err) + } - 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 - } + challengeResponsePacket := make([]byte, ed25519.PublicKeySize+ed25519.SignatureSize+1) + bytesRead := 0 + for bytesRead < len(challengeResponsePacket) { + n, err := conn.Read(challengeResponsePacket[bytesRead:]) + if err != nil { + return nil, fmt.Errorf("error reading challenge response: %w", err) + } + bytesRead += n + } - 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 + if challengeResponsePacket[0] != packet.VERSION { + return nil, fmt.Errorf("incompatible version: %v", challengeResponsePacket[0]) + } - case <-ctx.Done(): - log.Printf("client %v: %v\n", conn.RemoteAddr().String(), ctx.Err()) - break outer - } + pubKey := ed25519.PublicKey(challengeResponsePacket[1 : 1+ed25519.PublicKeySize]) + signature := ed25519.PrivateKey(challengeResponsePacket[1+ed25519.PublicKeySize:]) + + if ok := ed25519.Verify(pubKey, nonce, signature); !ok { + return nil, errors.New("signature verification failed") } + + conn.SetDeadline(time.Time{}) + return pubKey, nil } type Session struct { + Server *server Addr *net.TCPAddr - WriteQueue <-chan packet.Packet + WriteQueue chan packet.Packet + ID snowflake.ID + PubKey ed25519.PublicKey mu sync.Mutex - challenge []byte + challenge []byte issuedTime time.Time } -func newSession(addr *net.TCPAddr, writeQueue <-chan packet.Packet) *Session { +func newSession(server *server, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session { session := &Session{ + Server: server, Addr: addr, - WriteQueue: writeQueue, + PubKey: pubKey, + WriteQueue: make(chan packet.Packet, 10), challenge: make([]byte, 32), // Recommended nonce size } session.Challenge() // Make sure an initial nonce is generated @@ -263,58 +306,43 @@ func FromContext(ctx context.Context) (*Session, bool) { // TODO: Move everything below this to somewhere else -func handlePacket(pkt packet.Packet) (packet.Packet, error) { +func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet { + session, ok := FromContext(ctx) + assert.Assert(ok, "context in process packet should always have a session") + 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 + var request packet.SendMessage if err := pkt.DecodePayload(&request); err != nil { - return packet.Packet{}, fmt.Errorf("decode error: %v", err) + log.Println("decode error:", err) + response = &packet.ErrorMessage{Error: "malformed payload"} + break } content := strings.TrimSpace(request.Content) if content == "" { - response = &packet.ErrorMessage{Error: "content must not be blank"} + response = &packet.ErrorMessage{Error: "message content must not be blank"} break } + node := session.Server.Node message := data.Message{ Id: node.Generate(), - SenderId: node.Generate(), - FrequencyId: node.Generate(), - NetworkId: node.Generate(), + SenderId: session.ID, + FrequencyId: node.Generate(), // TODO: replace with actual ID + NetworkId: node.Generate(), // TODO: replace with actual ID 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} + response = packet.NewOkMessage() default: - return packet.Packet{}, errors.New("TODO: not implemented yet") + response = &packet.ErrorMessage{Error: "use of unsupported packet type"} } - 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 + assert.NotNil(response, "response must always be assigned to") + return packet.NewPacket(packet.NewMsgPackEncoder(response)) } -var ( - node = snowflake.NewNode(1) - messages []data.Message -) +var messages []data.Message -- cgit v1.3.1