summaryrefslogtreecommitdiff
path: root/internal/client
diff options
context:
space:
mode:
Diffstat (limited to 'internal/client')
-rw-r--r--internal/client/api/api.go29
-rw-r--r--internal/client/client.go226
-rw-r--r--internal/client/gateway/gateway.go215
-rw-r--r--internal/client/gateway/server.crt (renamed from internal/client/server.crt)0
-rw-r--r--internal/client/ui.go125
5 files changed, 378 insertions, 217 deletions
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)
- ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
- defer cancel()
- out, outErr := packet.RunFramer(ctx, conn)
+ 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)
- select {
- case responsePacket := <-out:
- if err := responsePacket.DecodePayload(response); err != nil {
- if responsePacket.Type() != packet.PacketError {
- return fmt.Errorf("error decoding response: %v", err)
+ 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
}
- var errorResponse packet.ErrorMessage
- if err := responsePacket.DecodePayload(&errorResponse); err != nil {
- return fmt.Errorf("error decoding error packet: %w", err)
+
+ m.textarea.Reset()
+ return m, api.SendMessage(content)
+
+ 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/server.crt b/internal/client/gateway/server.crt
index caf2384..caf2384 100644
--- a/internal/client/server.crt
+++ b/internal/client/gateway/server.crt
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"
-}