From 688668f6d34d4d962247bd5df60c3fca61185927 Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Thu, 24 Oct 2024 16:39:14 +0300 Subject: feat: worked on client --- internal/client/api/api.go | 29 +++++ internal/client/client.go | 230 ++++++++++++++++++++++--------------- internal/client/gateway/gateway.go | 215 ++++++++++++++++++++++++++++++++++ internal/client/gateway/server.crt | 21 ++++ internal/client/server.crt | 21 ---- internal/client/ui.go | 125 -------------------- internal/data/chat.go | 29 ++++- internal/packet/encoders.go | 17 +-- internal/packet/packet.go | 37 +++++- internal/server/server.go | 74 ++++++++---- 10 files changed, 518 insertions(+), 280 deletions(-) create mode 100644 internal/client/api/api.go create mode 100644 internal/client/gateway/gateway.go create mode 100644 internal/client/gateway/server.crt delete mode 100644 internal/client/server.crt delete mode 100644 internal/client/ui.go (limited to 'internal') diff --git a/internal/client/api/api.go b/internal/client/api/api.go new file mode 100644 index 0000000..56cf054 --- /dev/null +++ b/internal/client/api/api.go @@ -0,0 +1,29 @@ +package api + +import ( + "log" + + tea "github.com/charmbracelet/bubbletea" + + "github.com/kyren223/eko/internal/client/gateway" + "github.com/kyren223/eko/internal/packet" + "github.com/kyren223/eko/pkg/assert" +) + +func SendMessage(message string) tea.Cmd { + return func() tea.Msg { + log.Println("request SendMessage sent") + request := packet.SendMessage{Content: message} + responsePayload, ok := <-gateway.Send(&request) + if !ok { + log.Println("request timeout") + return "" + } + log.Println("request SendMessage received response") + + response := responsePayload.(*packet.ErrorMessage) + assert.Assert(ok, "server should return an ErrorMessage response for a SendMessage request") + assert.Assert(response.IsOk(), "server should always return an OK if we didn't mess up") + return "" + } +} diff --git a/internal/client/client.go b/internal/client/client.go index 5f42c8f..1276f01 100644 --- a/internal/client/client.go +++ b/internal/client/client.go @@ -2,125 +2,167 @@ package client import ( "context" - "crypto/tls" - "crypto/x509" - _ "embed" + "crypto/ed25519" "fmt" "log" - "os" - "time" - + "slices" + "strings" + + "github.com/charmbracelet/bubbles/cursor" + "github.com/charmbracelet/bubbles/textarea" + "github.com/charmbracelet/bubbles/viewport" + tea "github.com/charmbracelet/bubbletea" + "github.com/charmbracelet/lipgloss" + + "github.com/kyren223/eko/internal/client/api" + "github.com/kyren223/eko/internal/client/gateway" + "github.com/kyren223/eko/internal/data" "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/pkg/assert" - "github.com/vmihailenco/msgpack/v5" ) -//go:embed server.crt -var certPEM []byte +type BubbleTeaCloser struct { + program *tea.Program +} -var tlsConfig *tls.Config +func (c BubbleTeaCloser) Close() error { + c.program.Quit() + return nil +} func Run() { - logFile, err := os.OpenFile("client.log", os.O_APPEND|os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0666) - if err != nil { - log.Fatal(err) - } - defer logFile.Close() - log.SetOutput(logFile) + log.Println("client started") + program := tea.NewProgram(initialModel(), tea.WithAltScreen()) + assert.AddFlush(BubbleTeaCloser{program}) - certPool := x509.NewCertPool() - if !certPool.AppendCertsFromPEM(certPEM) { - log.Fatalln("failed to append server certificate") - } + _, privKey, err := ed25519.GenerateKey(nil) + assert.NoError(err, "private key gen should not error") - tlsConfig = &tls.Config{ - RootCAs: certPool, - ServerName: "localhost", + gateway.Connect(context.Background(), program, privKey) + if _, err := program.Run(); err != nil { + log.Println(err) } +} - log.Println("client started, waiting for user input...") - startUI() - - // for { - // fmt.Print("> ") - // input, _ := bufio.NewReader(os.Stdin).ReadString('\n') - // input = strings.TrimSpace(input) - // if input == ":q" || input == "exit" || input == "quit" { - // break - // } - // if input == "" { - // continue - // } - // err := processRequest(input) - // if err != nil { - // log.Println(err) - // fmt.Println(err) - // } - // } +type model struct { + viewport viewport.Model + messages []string + textarea textarea.Model + senderStyle lipgloss.Style + err error } -func sendMessage(message string) error { - request := packet.SendMessageMessage{Content: message} - var response packet.EkoMessage - if err := SendAndReceive(&request, &response); err != nil { - return err +func initialModel() model { + ta := textarea.New() + ta.Placeholder = "Send a message..." + ta.Focus() + + ta.Prompt = "┃ " + ta.CharLimit = 280 + + ta.SetWidth(30) + ta.SetHeight(3) + + // Remove cursor line styling + ta.FocusedStyle.CursorLine = lipgloss.NewStyle() + + ta.ShowLineNumbers = false + + vp := viewport.New(30, 20) + // vp.SetContent("Welcome to Eko!\n Type a message and press Enter to send.") + + ta.KeyMap.InsertNewline.SetEnabled(false) + + return model{ + textarea: ta, + messages: []string{}, + viewport: vp, + senderStyle: lipgloss.NewStyle().Foreground(lipgloss.Color("5")), + err: nil, } - assert.Assert(response.Message == "Eko OK", "server should return an OK status") - return nil } -func getMessages() ([]string, error) { - request := packet.GetMessagesMessage{} - var response packet.MessagesMessage - if err := SendAndReceive(&request, &response); err != nil { - return nil, err - } - var messages []string - for _, message := range response.Messages { - messages = append(messages, message.Contents) - } - return messages, nil +func (m model) Init() tea.Cmd { + return textarea.Blink } -func SendAndReceive(request packet.TypedMessage, response packet.TypedMessage) error { - conn, err := tls.Dial("tcp4", ":7223", tlsConfig) - if err != nil { - return fmt.Errorf("error establishing connection with server: %v", err) - } - defer conn.Close() - log.Println("established connection with server:", conn.RemoteAddr().String()) +func (m model) View() string { + return fmt.Sprintf( + "%s\n%s", + m.viewport.View(), + m.textarea.View(), + ) + "" +} - encoder, err := packet.NewMsgPackEncoder(request) - if err != nil { - return fmt.Errorf("error encoding request: %v", err) - } - requestPacket := packet.NewPacket(encoder) - err = requestPacket.Into(conn) - if err != nil { - return fmt.Errorf("error sending request: %v", err) - } - log.Println("sent request to server") +func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { + switch msg := msg.(type) { + case tea.WindowSizeMsg: + m.viewport.Width = msg.Width + m.viewport.Height = msg.Height - m.textarea.Height() + m.viewport.GotoBottom() + + m.textarea.SetWidth(msg.Width) + log.Println("resized to:", msg.Width, "x", msg.Height) + + var vpCmd tea.Cmd + var taCmd tea.Cmd + m.viewport, vpCmd = m.viewport.Update(msg) + m.textarea, taCmd = m.textarea.Update(msg) + return m, tea.Batch(vpCmd, taCmd) + + case tea.KeyMsg: + switch msg.Type { + case tea.KeyEsc, tea.KeyCtrlC: + return m, tea.Quit + + case tea.KeyEnter: + value := m.textarea.Value() + + content := strings.TrimSpace(value) + if content == "" { + return m, nil + } - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - out, outErr := packet.RunFramer(ctx, conn) + m.textarea.Reset() + return m, api.SendMessage(content) - select { - case responsePacket := <-out: - if err := responsePacket.DecodePayload(response); err != nil { - if responsePacket.Type() != packet.PacketError { - return fmt.Errorf("error decoding response: %v", err) - } - var errorResponse packet.ErrorMessage - if err := responsePacket.DecodePayload(&errorResponse); err != nil { - return fmt.Errorf("error decoding error packet: %w", err) + default: + // Send all other keypresses to the textarea. + var cmd tea.Cmd + m.textarea, cmd = m.textarea.Update(msg) + return m, cmd + } + + case packet.Payload: + switch msg := msg.(type) { + case *packet.Messages: + slices.SortFunc(msg.Messages, func(a, b data.Message) int { + return a.CmpTimestamp(b) + }) + + m.messages = []string{} + for _, message := range msg.Messages { + m.messages = append(m.messages, message.Contents) } - return fmt.Errorf("server error: %v", errorResponse.Error) + + m.viewport.SetContent(strings.Join(m.messages, "\n")) + m.viewport.GotoBottom() + + var cmd tea.Cmd + m.viewport, cmd = m.viewport.Update(msg) + return m, cmd + + default: + return m, nil } - case err := <-outErr: - return fmt.Errorf("error receiving response: %v", err) - } + case cursor.BlinkMsg: + // Textarea should also process cursor blinks. + var cmd tea.Cmd + m.textarea, cmd = m.textarea.Update(msg) + return m, cmd - return nil + default: + return m, nil + } } diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go new file mode 100644 index 0000000..2798950 --- /dev/null +++ b/internal/client/gateway/gateway.go @@ -0,0 +1,215 @@ +package gateway + +import ( + "context" + "crypto/ed25519" + "crypto/tls" + "crypto/x509" + _ "embed" + "errors" + "io" + "log" + "net" + "os" + "sync" + "time" + + tea "github.com/charmbracelet/bubbletea" + + "github.com/kyren223/eko/internal/packet" + "github.com/kyren223/eko/pkg/assert" +) + +//go:embed server.crt +var certPEM []byte + +var ( + tlsConfig *tls.Config + + asyncResponses []chan packet.Payload + responsesMu sync.Mutex + + connection net.Conn + connMu sync.Mutex +) + +func init() { + certPool := x509.NewCertPool() + if !certPool.AppendCertsFromPEM(certPEM) { + log.Fatalln("failed to append server certificate") + } + + tlsConfig = &tls.Config{ + RootCAs: certPool, + ServerName: "localhost", + } +} + +func Connect(ctx context.Context, program *tea.Program, privKey ed25519.PrivateKey) { + conn, err := tls.Dial("tcp4", ":7223", tlsConfig) + if err != nil { + assert.NoError(err, "TODO handle error") + } + log.Println("established connection with server") + + if err := handleAuth(ctx, conn, privKey); err != nil { + assert.NoError(err, "TODO handle error") + } + log.Println("successfully authenticated with server") + + framer := packet.NewFramer(ctx) + + go func() { + connection = conn + handleConnection(ctx, conn, framer) + close(framer.Out) + conn.Close() + connection = nil + }() + + go handlePacketStream(framer, program) +} + +func handleAuth(ctx context.Context, conn net.Conn, privKey ed25519.PrivateKey) error { + const nonceSize = 32 + challengeRequest := make([]byte, 1+nonceSize) + + err := conn.SetDeadline(time.Now().Add(10 * time.Second)) + assert.NoError(err, "setting deadline should not error") + defer func() { + err := conn.SetDeadline(time.Time{}) + assert.NoError(err, "unsetting deadline should not error") + }() + bytesRead := 0 + for bytesRead < 1+nonceSize { + n, err := conn.Read(challengeRequest[bytesRead:]) + if err != nil { + return err + } + bytesRead += n + } + + assert.Assert(challengeRequest[0] == packet.VERSION, "client should always have the same version as the server") + + challengeResponse := make([]byte, 1+ed25519.PublicKeySize+ed25519.SignatureSize) + challengeResponse[0] = packet.VERSION + copy(challengeResponse[1:1+ed25519.PublicKeySize], privKey.Public().(ed25519.PublicKey)) + signedNonce := ed25519.Sign(privKey, challengeRequest[1:]) + n := copy(challengeResponse[1+ed25519.PublicKeySize:], signedNonce) + assert.Assert(n == ed25519.SignatureSize, "copy should've copied the entire signature exactly") + + _, err = conn.Write(challengeResponse) + if err != nil { + return err + } + + return nil +} + +func handleConnection(ctx context.Context, conn net.Conn, framer packet.PacketFramer) { + buffer := make([]byte, 512) + for { + err := conn.SetReadDeadline(time.Now().Add(time.Second)) + assert.NoError(err, "setting a read deadline should not error") + n, err := conn.Read(buffer) + deadlineExceeded := errors.Is(err, os.ErrDeadlineExceeded) + if err != nil && !deadlineExceeded { + if !errors.Is(err, io.EOF) { + log.Println("read error:", err) + } + break + } + + if ctx.Err() != nil { + log.Println("context error:", ctx.Err()) + break + } + + err = framer.Push(ctx, buffer[:n]) + if ctx.Err() != nil { + log.Println("context error:", ctx.Err()) + break + } + assert.NoError(err, "packets from server should always be correct") + } +} + +func handlePacketStream(framer packet.PacketFramer, program *tea.Program) { + for { + pkt, ok := <-framer.Out + if !ok { + break + } + + payload, err := pkt.DecodedPayload() + assert.NoError(err, "server should always provide a decodeable packet") + + if pkt.Type().IsPush() { + log.Println("received streamed packet:", payload) + program.Send(payload) + continue + } + + responsesMu.Lock() + assert.Assert(len(asyncResponses) != 0, "there must always be at least 1 response waiting") + responseChan := asyncResponses[0] + copy(asyncResponses, asyncResponses[1:]) + asyncResponses = asyncResponses[:len(asyncResponses)-1] + responsesMu.Unlock() + + go func() { + responseChan <- payload + }() + } +} + +func conn() net.Conn { + return connection +} + +func Send(request packet.Payload) <-chan packet.Payload { + responseChan := make(chan packet.Payload) + go func() { + pkt := packet.NewPacket(packet.NewMsgPackEncoder(request)) + + conn := conn() + if conn == nil { + log.Println("request send error:", "connection is closed") + close(responseChan) + return + } + + connMu.Lock() + errDeadline := conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + assert.NoError(errDeadline, "setting a write deadline should not error") + _, err := pkt.Into(conn) + errDeadline = conn.SetWriteDeadline(time.Time{}) + assert.NoError(errDeadline, "setting a write deadline should not error") + connMu.Unlock() + if err != nil { + log.Println("request send error:", err) + close(responseChan) + return + } + + responsesMu.Lock() + asyncResponses = append(asyncResponses, responseChan) + responsesMu.Unlock() + + time.Sleep(5 * time.Second) + responsesMu.Lock() + index := -1 + for i, ch := range asyncResponses { + if ch == responseChan { + index = i + } + } + if index != -1 { + copy(asyncResponses[index:], asyncResponses[index+1:]) + asyncResponses = asyncResponses[:len(asyncResponses)-1] + close(responseChan) + } + responsesMu.Unlock() + }() + return responseChan +} diff --git a/internal/client/gateway/server.crt b/internal/client/gateway/server.crt new file mode 100644 index 0000000..caf2384 --- /dev/null +++ b/internal/client/gateway/server.crt @@ -0,0 +1,21 @@ +-----BEGIN CERTIFICATE----- +MIIDbTCCAlWgAwIBAgIUZyvzq7LOxKqTttTRjChoJWOO4pYwDQYJKoZIhvcNAQEL +BQAwOjEhMB8GCSqGSIb3DQEJARYSa3lyZW4yMjNAcHJvdG9uLm1lMRUwEwYDVQQD +DAxreXJlbjIyMy5kZXYwHhcNMjQxMDA5MjE1NzQzWhcNMjUxMDA5MjE1NzQzWjA6 +MSEwHwYJKoZIhvcNAQkBFhJreXJlbjIyM0Bwcm90b24ubWUxFTATBgNVBAMMDGt5 +cmVuMjIzLmRldjCCASIwDQYJKoZIhvcNAQEBBQADggEPADCCAQoCggEBAK7hd+zT +kqrn/8EhLEO0uMKKHgfoyczYWTlA9uPFADOsjdzXRLuR/Y3rK0PBE4u55xcjYZSf +mzJmVHuv1rEFOt634YOoE2UwJd9V2M0p+cD716XIEDNPfVCUe77FoZoYaH1h8QF5 +Mrx2eDH5JZt690F05O39zYzbb7+RlChWlt1kBcmLEZ1GKJeXznbL6lLMh20deYX9 +7oemqYMqP9DFbFeHkubeZ20yQvKW9cOWae9M+IhE9dAa8fm5WdfiDoTdAHfbIawx +r1OB4YqfXlXler9wAHfHWeCS0KgZCTdghF1h6wtYlwyQZcUuv+dHN7SP7zVo8pOD +b7NUqjFAMGlNgf0CAwEAAaNrMGkwCQYDVR0TBAIwADALBgNVHQ8EBAMCBaAwMAYD +VR0RBCkwJ4IJbG9jYWxob3N0ggsqLmxvY2FsaG9zdIINZWtvLmxvY2FsaG9zdDAd +BgNVHQ4EFgQUf32SXO976zgO0K/wlgWdyT3EPzcwDQYJKoZIhvcNAQELBQADggEB +AHVMGCkaZv5eIOQwevfrsEJQo3dNG34om8wBVGS5iQyho0VJZpKZSiQ16yv4x2kc +UICfVEFcfO/7/hRlA5yLWE/wpeqCgTSgtQ74gvc8D6H26wCznSPj9MIRWxYhSmPM +YO+7UKqyvFoaKiW4OkqJvCRzrpwr/lbXcGpD47UqT5gRvjJ91ULCHIUt8qDUS6+8 +mEGJAe/xFkiJ6zT0bThlqMaCA4v5g9tHGXzooIZ+YSgTvlWhAM6mVwt34l2rDSOw +4YNGUJXKCoGpy8U0NteIOOs6HhaslJpKe1mSSxmMQcgBcaf6yBT08mYfQPSsaeOk +OSoncuVBnT64liAtShpsgTc= +-----END CERTIFICATE----- diff --git a/internal/client/server.crt b/internal/client/server.crt deleted file mode 100644 index caf2384..0000000 --- a/internal/client/server.crt +++ /dev/null @@ -1,21 +0,0 @@ ------BEGIN CERTIFICATE----- -MIIDbTCCAlWgAwIBAgIUZyvzq7LOxKqTttTRjChoJWOO4pYwDQYJKoZIhvcNAQEL -BQAwOjEhMB8GCSqGSIb3DQEJARYSa3lyZW4yMjNAcHJvdG9uLm1lMRUwEwYDVQQD -DAxreXJlbjIyMy5kZXYwHhcNMjQxMDA5MjE1NzQzWhcNMjUxMDA5MjE1NzQzWjA6 -MSEwHwYJKoZIhvcNAQkBFhJreXJlbjIyM0Bwcm90b24ubWUxFTATBgNVBAMMDGt5 -cmVuMjIzLmRldjCCASIwDQYJKoZIhvcNAQEBBQADggEPADCCAQoCggEBAK7hd+zT -kqrn/8EhLEO0uMKKHgfoyczYWTlA9uPFADOsjdzXRLuR/Y3rK0PBE4u55xcjYZSf -mzJmVHuv1rEFOt634YOoE2UwJd9V2M0p+cD716XIEDNPfVCUe77FoZoYaH1h8QF5 -Mrx2eDH5JZt690F05O39zYzbb7+RlChWlt1kBcmLEZ1GKJeXznbL6lLMh20deYX9 -7oemqYMqP9DFbFeHkubeZ20yQvKW9cOWae9M+IhE9dAa8fm5WdfiDoTdAHfbIawx -r1OB4YqfXlXler9wAHfHWeCS0KgZCTdghF1h6wtYlwyQZcUuv+dHN7SP7zVo8pOD -b7NUqjFAMGlNgf0CAwEAAaNrMGkwCQYDVR0TBAIwADALBgNVHQ8EBAMCBaAwMAYD -VR0RBCkwJ4IJbG9jYWxob3N0ggsqLmxvY2FsaG9zdIINZWtvLmxvY2FsaG9zdDAd -BgNVHQ4EFgQUf32SXO976zgO0K/wlgWdyT3EPzcwDQYJKoZIhvcNAQELBQADggEB -AHVMGCkaZv5eIOQwevfrsEJQo3dNG34om8wBVGS5iQyho0VJZpKZSiQ16yv4x2kc -UICfVEFcfO/7/hRlA5yLWE/wpeqCgTSgtQ74gvc8D6H26wCznSPj9MIRWxYhSmPM -YO+7UKqyvFoaKiW4OkqJvCRzrpwr/lbXcGpD47UqT5gRvjJ91ULCHIUt8qDUS6+8 -mEGJAe/xFkiJ6zT0bThlqMaCA4v5g9tHGXzooIZ+YSgTvlWhAM6mVwt34l2rDSOw -4YNGUJXKCoGpy8U0NteIOOs6HhaslJpKe1mSSxmMQcgBcaf6yBT08mYfQPSsaeOk -OSoncuVBnT64liAtShpsgTc= ------END CERTIFICATE----- diff --git a/internal/client/ui.go b/internal/client/ui.go deleted file mode 100644 index 9de322c..0000000 --- a/internal/client/ui.go +++ /dev/null @@ -1,125 +0,0 @@ -package client - -import ( - "fmt" - "log" - "strings" - - "github.com/charmbracelet/bubbles/cursor" - "github.com/charmbracelet/bubbles/textarea" - "github.com/charmbracelet/bubbles/viewport" - tea "github.com/charmbracelet/bubbletea" - "github.com/charmbracelet/lipgloss" - - "github.com/kyren223/eko/pkg/assert" -) - -func startUI() { - p := tea.NewProgram(initialModel(), tea.WithAltScreen()) - if _, err := p.Run(); err != nil { - log.Println("charm ui error:", err) - } - p. -} - -type model struct { - viewport viewport.Model - messages []string - textarea textarea.Model - senderStyle lipgloss.Style - err error -} - -func initialModel() model { - ta := textarea.New() - ta.Placeholder = "Send a message..." - ta.Focus() - - ta.Prompt = "┃ " - ta.CharLimit = 280 - - ta.SetWidth(30) - ta.SetHeight(3) - - // Remove cursor line styling - ta.FocusedStyle.CursorLine = lipgloss.NewStyle() - - ta.ShowLineNumbers = false - - vp := viewport.New(30, 20) - vp.SetContent("Welcome to Eko!\n Type a message and press Enter to send.") - - ta.KeyMap.InsertNewline.SetEnabled(false) - - messages, err := getMessages() - assert.NoError(err, "TODO HANDLE ERROR") - vp.SetContent(strings.Join(messages, "\n")) - - return model{ - textarea: ta, - messages: messages, - viewport: vp, - senderStyle: lipgloss.NewStyle().Foreground(lipgloss.Color("5")), - err: nil, - } -} - -func (m model) Init() tea.Cmd { - return textarea.Blink -} - -func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { - switch msg := msg.(type) { - case tea.WindowSizeMsg: - m.viewport.Width = msg.Width - m.textarea.SetWidth(msg.Width) - return m, tea.Printf("%dx%d", msg.Width, msg.Height) - case tea.KeyMsg: - switch msg.Type { - case tea.KeyEsc, tea.KeyCtrlC: - // Quit. - fmt.Println(m.textarea.Value()) - return m, tea.Quit - case tea.KeyEnter: - value := m.textarea.Value() - - content := strings.TrimSpace(value) - if content == "" { - // Don't send empty messages. - return m, nil - } - assert.NoError(sendMessage(content), "TODO HANDLE ERROR") - messages, err := getMessages() - assert.NoError(err, "TODO HANDLE ERROR") - - // m.messages = append(m.messages, m.senderStyle.Render("You: ")+content) - m.messages = messages - m.viewport.SetContent(strings.Join(m.messages, "\n")) - m.textarea.Reset() - m.viewport.GotoBottom() - return m, nil - default: - // Send all other keypresses to the textarea. - var cmd tea.Cmd - m.textarea, cmd = m.textarea.Update(msg) - return m, cmd - } - - case cursor.BlinkMsg: - // Textarea should also process cursor blinks. - var cmd tea.Cmd - m.textarea, cmd = m.textarea.Update(msg) - return m, cmd - - default: - return m, nil - } -} - -func (m model) View() string { - return fmt.Sprintf( - "%s\n\n%s", - m.viewport.View(), - m.textarea.View(), - ) + "\n\n" -} diff --git a/internal/data/chat.go b/internal/data/chat.go index 9d2443d..ae5dc67 100644 --- a/internal/data/chat.go +++ b/internal/data/chat.go @@ -1,6 +1,9 @@ package data -import "github.com/kyren223/eko/pkg/snowflake" +import ( + "github.com/kyren223/eko/pkg/snowflake" + "github.com/kyren223/eko/pkg/utils" +) // Represents an Eko Network, equivalent to a "discord server". type Network struct { @@ -23,11 +26,27 @@ type Signal struct { // Represents a message type Message struct { - Id snowflake.ID - SenderId snowflake.ID + Id snowflake.ID + SenderId snowflake.ID FrequencyId snowflake.ID - NetworkId snowflake.ID - Contents string + NetworkId snowflake.ID + Contents string +} + +func (a Message) CmpTimestamp(b Message) int { + cmpTime := a.Id.Time() - b.Id.Time() + if cmpTime != 0 { + return int(utils.Clamp(cmpTime, -1, 1)) + } + cmpStep := a.Id.Step() - b.Id.Step() + if cmpStep != 0 { + return int(utils.Clamp(cmpStep, -1, 1)) + } + cmpNode := a.Id.Node() - b.Id.Node() + if cmpNode != 0 { + return int(utils.Clamp(cmpNode, -1, 1)) + } + return 0 } // Represents an Eko User diff --git a/internal/packet/encoders.go b/internal/packet/encoders.go index f5841ce..760c11c 100644 --- a/internal/packet/encoders.go +++ b/internal/packet/encoders.go @@ -3,11 +3,12 @@ package packet import ( "encoding/json" - "github.com/kyren223/eko/pkg/assert" "github.com/vmihailenco/msgpack/v5" + + "github.com/kyren223/eko/pkg/assert" ) -type TypedMessage interface { +type Payload interface { Type() PacketType } @@ -29,24 +30,24 @@ func (e defaultPacketEncoder) Payload() []byte { return e.data } -func NewJsonEncoder(message TypedMessage) PacketEncoder { - data, err := json.Marshal(message) +func NewJsonEncoder(payload Payload) PacketEncoder { + data, err := json.Marshal(payload) assert.NoError(err, "encoding a message with JSON should never fail") return defaultPacketEncoder{ data: data, encoding: EncodingJson, - packetType: message.Type(), + packetType: payload.Type(), } } -func NewMsgPackEncoder(message TypedMessage) PacketEncoder { - data, err := msgpack.Marshal(message) +func NewMsgPackEncoder(payload Payload) PacketEncoder { + data, err := msgpack.Marshal(payload) assert.NoError(err, "encoding a message with msg pack should never fail") return defaultPacketEncoder{ data: data, encoding: EncodingMsgPack, - packetType: message.Type(), + packetType: payload.Type(), } } diff --git a/internal/packet/packet.go b/internal/packet/packet.go index e9d5d94..f38ad14 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -77,6 +77,15 @@ func (e PacketType) IsSupported() bool { } } +func (e PacketType) IsPush() bool { + switch e { + case PacketError, PacketSendMessage: + return false + default: + return true + } +} + const ( VERSION = byte(1) PACKET_MAX_SIZE = math.MaxUint16 @@ -156,7 +165,7 @@ func (p Packet) Into(writer io.Writer) (int, error) { return writer.Write(p.data) } -func (p Packet) DecodePayload(v TypedMessage) error { +func (p Packet) DecodePayloadInto(v Payload) error { if p.Type() != v.Type() { return fmt.Errorf("type mismatch: want %v got %v", p.Type(), v.Type()) } @@ -175,24 +184,40 @@ func (p Packet) DecodePayload(v TypedMessage) error { } } +func (p Packet) DecodedPayload() (Payload, error) { + var payload Payload + switch p.Type() { + case PacketError: + payload = &ErrorMessage{} + case PacketMessages: + payload = &Messages{} + case PacketSendMessage: + payload = &SendMessage{} + default: + assert.Never("packet type of a packet struct must always be valid") + } + err := p.DecodePayloadInto(payload) + return payload, 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 { +type PacketFramer struct { buffer []byte Out chan Packet } -func NewFramer(ctx context.Context) packetFramer { - return packetFramer{ +func NewFramer(ctx context.Context) PacketFramer { + return PacketFramer{ Out: make(chan Packet, 10), } } -func (f *packetFramer) Push(ctx context.Context, data []byte) error { +func (f *PacketFramer) Push(ctx context.Context, data []byte) error { f.buffer = append(f.buffer, data...) for { @@ -208,7 +233,7 @@ func (f *packetFramer) Push(ctx context.Context, data []byte) error { } } -func (f *packetFramer) parse() (*Packet, error) { +func (f *PacketFramer) parse() (*Packet, error) { if len(f.buffer) < HEADER_SIZE { return nil, nil } diff --git a/internal/server/server.go b/internal/server/server.go index 4ac8969..ff28ad7 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -168,6 +168,9 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { if !ok { break } + if packet.Type().IsPush() { + log.Println(addr, "streaming packet:", packet) + } if _, err := packet.Into(conn); err != nil { log.Println(addr, err) break @@ -192,7 +195,8 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { buffer := make([]byte, 512) for { - conn.SetReadDeadline(time.Now().Add(time.Second)) + err := conn.SetReadDeadline(time.Now().Add(time.Second)) + assert.NoError(err, "setting read deadline should not error") n, err := conn.Read(buffer) deadlineExceeded := errors.Is(err, os.ErrDeadlineExceeded) if err != nil && !deadlineExceeded { @@ -220,13 +224,19 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) { } func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { - conn.SetDeadline(time.Now().Add(time.Second * 5)) + err := conn.SetDeadline(time.Now().Add(time.Second * 5)) + assert.NoError(err, "setting read deadline should not error") + + defer func() { + err := conn.SetDeadline(time.Time{}) + assert.NoError(err, "unsetting read deadline should not error") + }() challengePacket := make([]byte, len(nonce)+1) challengePacket[0] = packet.VERSION copy(challengePacket[1:], nonce) - _, err := conn.Write(challengePacket) + _, err = conn.Write(challengePacket) if err != nil { return nil, fmt.Errorf("error writing challenge: %w", err) } @@ -252,7 +262,6 @@ func handleAuth(conn net.Conn, nonce []byte) (ed25519.PublicKey, error) { return nil, errors.New("signature verification failed") } - conn.SetDeadline(time.Time{}) return pubKey, nil } @@ -310,20 +319,31 @@ 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.PacketSendMessage: - var request packet.SendMessage - if err := pkt.DecodePayload(&request); err != nil { - log.Println("decode error:", err) - response = &packet.ErrorMessage{Error: "malformed payload"} - break - } + var response packet.Payload + + request, err := pkt.DecodedPayload() + if err != nil { + response = &packet.ErrorMessage{Error: "malformed payload"} + } else { + response = processRequest(ctx, request) + } + + assert.NotNil(response, "response must always be assigned to") + log.Println(session.Addr, "sending ", response.Type(), "response:", response) + return packet.NewPacket(packet.NewMsgPackEncoder(response)) +} + +func processRequest(ctx context.Context, request packet.Payload) packet.Payload { + session, ok := FromContext(ctx) + assert.Assert(ok, "context in process packet should always have a session") + + log.Println(session.Addr, "processing", request.Type(), "request:", request) + switch request := request.(type) { + case *packet.SendMessage: content := strings.TrimSpace(request.Content) if content == "" { - response = &packet.ErrorMessage{Error: "message content must not be blank"} - break + return &packet.ErrorMessage{Error: "message content must not be blank"} } node := session.Server.Node @@ -334,15 +354,27 @@ func processPacket(ctx context.Context, pkt packet.Packet) packet.Packet { NetworkId: node.Generate(), // TODO: replace with actual ID Contents: content, } - messages = append(messages, message) + sendMessage(ctx, message) + + return packet.NewOkMessage() - response = packet.NewOkMessage() default: - response = &packet.ErrorMessage{Error: "use of unsupported packet type"} + return &packet.ErrorMessage{Error: "use of unsupported packet type"} } - - assert.NotNil(response, "response must always be assigned to") - return packet.NewPacket(packet.NewMsgPackEncoder(response)) } var messages []data.Message + +func sendMessage(ctx context.Context, msg data.Message) { + session, ok := FromContext(ctx) + assert.Assert(ok, "context in process packet should always have a session") + + messages = append(messages, msg) + + payload := &packet.Messages{ + Messages: messages, + } + + pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload)) + session.WriteQueue <- pkt +} -- cgit v1.3.1