summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/client/client.go98
-rw-r--r--internal/client/ui.go124
-rw-r--r--internal/packet/packet.go2
-rw-r--r--internal/server/handler.go54
4 files changed, 212 insertions, 66 deletions
diff --git a/internal/client/client.go b/internal/client/client.go
index 83b3d90..99c53b3 100644
--- a/internal/client/client.go
+++ b/internal/client/client.go
@@ -1,79 +1,94 @@
package client
import (
- "bufio"
"context"
"crypto/tls"
"crypto/x509"
_ "embed"
"fmt"
"log"
- "net"
"os"
- "strings"
"time"
"github.com/kyren223/eko/internal/packet"
+ "github.com/kyren223/eko/pkg/assert"
)
//go:embed server.crt
var certPEM []byte
+var tlsConfig *tls.Config
+
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)
+
certPool := x509.NewCertPool()
if !certPool.AppendCertsFromPEM(certPEM) {
log.Fatalln("failed to append server certificate")
}
- tlsConfig := &tls.Config{
+ tlsConfig = &tls.Config{
RootCAs: certPool,
ServerName: "localhost",
}
log.Println("client started, waiting for user input...")
- for {
- fmt.Print("> ")
- input, _ := bufio.NewReader(os.Stdin).ReadString('\n')
- input = strings.TrimSpace(input)
- if input == ":q" || input == "exit" || input == "quit" {
- break
- }
- err := processRequest(input, tlsConfig)
- if err != nil {
- log.Println(err)
- }
+ 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)
+ // }
+ // }
+}
+
+func sendMessage(message string) error {
+ request := packet.SendMessageMessage{Content: message}
+ var response packet.EkoMessage
+ if err := SendAndReceive(&request, &response); err != nil {
+ return err
+ }
+ 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 processRequest(input string, tlsConfig *tls.Config) error {
+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())
- if input == "SHOW" {
- request := packet.GetMessagesMessage{}
- var response packet.MessagesMessage
- if err := SendAndReceive(conn, &request, &response); err != nil {
- return err
- }
- log.Println("server response:", response.Messages)
- } else {
- request := packet.SendMessageMessage{Content: input}
- var response packet.EkoMessage
- if err := SendAndReceive(conn, &request, &response); err != nil {
- return err
- }
- log.Println("server response:", response.Message)
- }
-
- return nil
-}
-
-func SendAndReceive(conn net.Conn, request packet.TypedMessage, response packet.TypedMessage) error {
encoder, err := packet.NewMsgPackEncoder(request)
if err != nil {
return fmt.Errorf("error encoding request: %v", err)
@@ -92,7 +107,14 @@ func SendAndReceive(conn net.Conn, request packet.TypedMessage, response packet.
select {
case responsePacket := <-out:
if err := responsePacket.DecodePayload(response); err != nil {
- return fmt.Errorf("error decoding response: %v", err)
+ if responsePacket.Type() != packet.TypeError {
+ 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)
+ }
+ return fmt.Errorf("server error: %v", errorResponse.Error)
}
case err := <-outErr:
diff --git a/internal/client/ui.go b/internal/client/ui.go
new file mode 100644
index 0000000..839c0cd
--- /dev/null
+++ b/internal/client/ui.go
@@ -0,0 +1,124 @@
+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)
+ }
+}
+
+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/packet/packet.go b/internal/packet/packet.go
index 3b49249..634643e 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -145,7 +145,7 @@ func (p Packet) PayloadLength() uint16 {
}
func (p Packet) String() string {
- return fmt.Sprintf("{v%v %v %v %v: %v}", p.data[0], p.Encoding().String(), p.Type().String(), p.PayloadLength(), p.Payload())
+ 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
diff --git a/internal/server/handler.go b/internal/server/handler.go
index cf1eb82..0ab5c33 100644
--- a/internal/server/handler.go
+++ b/internal/server/handler.go
@@ -6,11 +6,13 @@ import (
"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"
)
@@ -23,12 +25,14 @@ func handleConnection(conn net.Conn, wg *sync.WaitGroup) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
out, outErr := packet.RunFramer(ctx, conn)
- log.Printf("client %v: running framer\n", conn.RemoteAddr().String())
outer:
for {
select {
- case packet := <-out:
+ 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)
@@ -43,16 +47,13 @@ outer:
}
case err := <-outErr:
- if err == nil {
- continue
- }
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 {
+ } else if err != nil {
log.Printf("client %v: internal error: %v\n", conn.RemoteAddr().String(), err)
}
break outer
@@ -65,6 +66,7 @@ outer:
}
func handlePacket(pkt packet.Packet) (packet.Packet, error) {
+ var response packet.TypedMessage
switch pkt.Type() {
case packet.TypeEko:
var request packet.EkoMessage
@@ -72,51 +74,49 @@ func handlePacket(pkt packet.Packet) (packet.Packet, error) {
return packet.Packet{}, fmt.Errorf("decode error: %v", err)
}
- response := packet.EkoMessage{Message: "Eko \"" + request.Message + "\""}
- encoder, err := packet.NewMsgPackEncoder(&response)
- if err != nil {
- return packet.Packet{}, fmt.Errorf("encode error: %v", err)
- }
- return packet.NewPacket(encoder), nil
+ 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: request.Content,
+ Contents: content,
}
messages = append(messages, message)
- response := packet.EkoMessage{Message: "Eko OK"}
- encoder, err := packet.NewMsgPackEncoder(&response)
- if err != nil {
- return packet.Packet{}, fmt.Errorf("encode error: %v", err)
- }
- return packet.NewPacket(encoder), nil
+ 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}
- encoder, err := packet.NewMsgPackEncoder(&response)
- if err != nil {
- return packet.Packet{}, fmt.Errorf("encode error: %v", err)
- }
- return packet.NewPacket(encoder), nil
+ 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 = make([]data.Message, 10)
+ node = snowflake.NewNode(1)
+ messages []data.Message
)