summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/client/api/api.go35
-rw-r--r--internal/client/client.go66
-rw-r--r--internal/client/gateway/gateway.go5
-rw-r--r--internal/client/ui/messagebox/messagebox.go132
-rw-r--r--internal/data/users.sql.go23
-rw-r--r--internal/packet/messages.go26
-rw-r--r--internal/packet/packet.go20
-rw-r--r--internal/server/api/api.go34
-rw-r--r--internal/server/server.go19
-rw-r--r--internal/server/session/session.go7
10 files changed, 292 insertions, 75 deletions
diff --git a/internal/client/api/api.go b/internal/client/api/api.go
index d52b727..7fcde8b 100644
--- a/internal/client/api/api.go
+++ b/internal/client/api/api.go
@@ -14,7 +14,10 @@ import (
"github.com/kyren223/eko/pkg/snowflake"
)
-type AppendMessage data.Message
+type (
+ AppendMessage data.Message
+ UserProfileUpdate data.User
+)
func SendMessage(message string) tea.Cmd {
return func() tea.Msg {
@@ -26,8 +29,7 @@ func SendMessage(message string) tea.Cmd {
}
response, ok := <-gateway.Send(&request)
if !ok {
- log.Println()
- return errors.New("request timeout")
+ return errors.New("request SendMessage timeout")
}
log.Println("request SendMessage received response")
@@ -53,8 +55,7 @@ func GetMessages() tea.Msg {
}
response, ok := <-gateway.Send(&request)
if !ok {
- log.Println()
- return errors.New("request timeout")
+ return errors.New("request GetMessages timeout")
}
log.Println("request GetMessages received response")
@@ -66,3 +67,27 @@ func GetMessages() tea.Msg {
}
return fmt.Errorf("received invalid response from server: %v", response.Type())
}
+
+func GetUserById(id snowflake.ID) tea.Cmd {
+ return func() tea.Msg {
+ log.Println("request GetUserById for ID", id, "sent")
+ request := packet.GetUserByID{UserID: id}
+ response, ok := <-gateway.Send(&request)
+ if !ok {
+ return errors.New("request GetUserById timeout")
+ }
+ log.Println("request GetUserById received response")
+
+ switch response := response.(type) {
+ case *packet.ErrorMessage:
+ return errors.New(response.Error)
+ case *packet.Users:
+ if len(response.Users) == 0 {
+ return fmt.Errorf("requested user id %v not found", id)
+ }
+ assert.Assert(len(response.Users) == 1, "server must return only one user with the matching id")
+ return UserProfileUpdate(response.Users[0])
+ }
+ return fmt.Errorf("received invalid response from server: %v", response.Type())
+ }
+}
diff --git a/internal/client/client.go b/internal/client/client.go
index 2dc432c..10ca6f2 100644
--- a/internal/client/client.go
+++ b/internal/client/client.go
@@ -5,17 +5,16 @@ import (
"crypto/ed25519"
"fmt"
"log"
- "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/client/ui/messagebox"
"github.com/kyren223/eko/internal/data"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/pkg/assert"
@@ -45,11 +44,9 @@ func Run() {
}
type model struct {
- viewport viewport.Model
- messages []string
- textarea textarea.Model
- senderStyle lipgloss.Style
- err error
+ messagebox messagebox.Model
+ textarea textarea.Model
+ err error
}
func initialModel() model {
@@ -68,17 +65,12 @@ func initialModel() model {
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,
+ textarea: ta,
+ messagebox: messagebox.New(30, 20),
+ err: nil,
}
}
@@ -89,7 +81,7 @@ func (m model) Init() tea.Cmd {
func (m model) View() string {
return fmt.Sprintf(
"%s\n%s",
- m.viewport.View(),
+ m.messagebox.View(),
m.textarea.View(),
) + ""
}
@@ -97,17 +89,17 @@ func (m model) View() string {
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.messagebox.Viewport.Width = msg.Width
+ m.messagebox.Viewport.Height = msg.Height - m.textarea.Height()
+ m.messagebox.Viewport.GotoBottom()
m.textarea.SetWidth(msg.Width)
log.Println("resized to:", msg.Width, "x", msg.Height)
- var vpCmd, taCmd tea.Cmd
- m.viewport, vpCmd = m.viewport.Update(msg)
+ var mbCmd, taCmd tea.Cmd
+ m.messagebox, mbCmd = m.messagebox.Update(msg)
m.textarea, taCmd = m.textarea.Update(msg)
- return m, tea.Batch(vpCmd, taCmd)
+ return m, tea.Batch(mbCmd, taCmd)
case tea.KeyMsg:
switch msg.Type {
@@ -121,8 +113,8 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
if content == "" {
return m, nil
}
-
m.textarea.Reset()
+
return m, api.SendMessage(content)
default:
@@ -134,24 +126,10 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
case packet.Payload:
switch msg := msg.(type) {
case *packet.Messages:
- slices.SortFunc(msg.Messages, func(a, b data.Message) int {
- if a.ID-b.ID < 0 {
- return -1
- } else {
- return 1
- }
- })
-
- m.messages = []string{}
- for _, message := range msg.Messages {
- m.messages = append(m.messages, message.Content)
- }
-
- m.viewport.SetContent(strings.Join(m.messages, "\n"))
- m.viewport.GotoBottom()
+ m.messagebox.SetMessages(msg.Messages)
var cmd tea.Cmd
- m.viewport, cmd = m.viewport.Update(msg)
+ m.messagebox, cmd = m.messagebox.Update(msg)
return m, cmd
default:
@@ -159,13 +137,15 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
}
case api.AppendMessage:
- m.messages = append(m.messages, msg.Content)
+ m.messagebox.AppendMessage(data.Message(msg))
- m.viewport.SetContent(strings.Join(m.messages, "\n"))
- m.viewport.GotoBottom()
+ var cmd tea.Cmd
+ m.messagebox, cmd = m.messagebox.Update(msg)
+ return m, cmd
+ case api.UserProfileUpdate:
var cmd tea.Cmd
- m.viewport, cmd = m.viewport.Update(msg)
+ m.messagebox, cmd = m.messagebox.Update(msg)
return m, cmd
case cursor.BlinkMsg:
diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go
index c196b1b..56f1ade 100644
--- a/internal/client/gateway/gateway.go
+++ b/internal/client/gateway/gateway.go
@@ -196,6 +196,11 @@ func Send(request packet.Payload) <-chan packet.Payload {
asyncResponses = append(asyncResponses, responseChan)
responsesMu.Unlock()
+ // TODO: this is a really bad implementation
+ // Should refactor to make it better
+ // Maybe have a timeout function similar to the one on the server?
+ // In any case this is like super bad to just arbitrary hang for 5 seconds
+ // And then try to close the thing so it times out
time.Sleep(5 * time.Second)
responsesMu.Lock()
index := -1
diff --git a/internal/client/ui/messagebox/messagebox.go b/internal/client/ui/messagebox/messagebox.go
new file mode 100644
index 0000000..4b8d3d1
--- /dev/null
+++ b/internal/client/ui/messagebox/messagebox.go
@@ -0,0 +1,132 @@
+package messagebox
+
+import (
+ "slices"
+ "strings"
+ "time"
+
+ "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/data"
+ "github.com/kyren223/eko/pkg/snowflake"
+)
+
+type user struct {
+ name string
+ isFetching bool
+}
+
+type Model struct {
+ Viewport viewport.Model
+ messages []data.Message
+
+ senderStyle lipgloss.Style
+ timestampStyle lipgloss.Style
+
+ users map[snowflake.ID]user
+}
+
+func New(width, height int) Model {
+ return Model{
+ Viewport: viewport.New(width, height),
+ messages: nil,
+ senderStyle: lipgloss.NewStyle().Foreground(lipgloss.Color("5")).Bold(true),
+ timestampStyle: lipgloss.NewStyle().Foreground(lipgloss.Color("#585c62")).Italic(true),
+ users: make(map[snowflake.ID]user),
+ }
+}
+
+func (m Model) Init() tea.Cmd {
+ return api.GetMessages
+}
+
+func (m Model) View() string {
+ return m.Viewport.View()
+}
+
+func (m Model) Update(msg tea.Msg) (Model, tea.Cmd) {
+ var cmds []tea.Cmd
+
+ for id, user := range m.users {
+ if !user.isFetching && user.name == "" {
+ cmds = append(cmds, api.GetUserById(id))
+ user.isFetching = true
+ m.users[id] = user
+ }
+ }
+
+ var cmd tea.Cmd
+ m.Viewport, cmd = m.Viewport.Update(msg)
+ cmds = append(cmds, cmd)
+
+ switch msg := msg.(type) {
+ case api.UserProfileUpdate:
+ m.users[msg.ID] = user{msg.Name, false}
+ m.UpdateMessages()
+ }
+
+ return m, tea.Batch(cmds...)
+}
+
+func (m *Model) AppendMessage(message data.Message) {
+ m.messages = append(m.messages, message)
+ m.UpdateMessages()
+}
+
+func (m *Model) SetMessages(messages []data.Message) {
+ m.messages = messages
+ m.UpdateMessages()
+}
+
+func (m *Model) UpdateMessages() {
+ for _, message := range m.messages {
+ if _, exists := m.users[message.SenderID]; !exists {
+ m.users[message.SenderID] = user{name: "", isFetching: false}
+ }
+ }
+
+ m.sortMessages()
+ m.updateContent()
+ m.Viewport.GotoBottom()
+}
+
+func (m *Model) sortMessages() {
+ slices.SortFunc(m.messages, func(a, b data.Message) int {
+ if a.ID-b.ID < 0 {
+ return -1
+ } else {
+ return 1
+ }
+ })
+}
+
+func (m *Model) updateContent() {
+ if len(m.messages) == 0 {
+ return
+ }
+
+ var builder strings.Builder
+
+ for _, message := range m.messages {
+ sender := m.users[message.SenderID].name
+ builder.WriteString(m.senderStyle.Render(sender))
+ builder.WriteByte(' ')
+
+ timestamp := timestampFromID(message.ID)
+ builder.WriteString(m.timestampStyle.Render(timestamp))
+ builder.WriteByte('\n')
+
+ builder.WriteString(message.Content)
+ builder.WriteByte('\n')
+ }
+
+ m.Viewport.SetContent(builder.String()[:builder.Len()-1])
+}
+
+func timestampFromID(id snowflake.ID) string {
+ localTime := time.UnixMilli(id.Time())
+ return localTime.Format("02/01/2006 3:04 PM")
+}
diff --git a/internal/data/users.sql.go b/internal/data/users.sql.go
index 1ca03ce..c92021d 100644
--- a/internal/data/users.sql.go
+++ b/internal/data/users.sql.go
@@ -16,30 +16,43 @@ const createUser = `-- name: CreateUser :one
INSERT INTO users (
id, name, public_key
) VALUES (
- ?, 'User' || abs(random()) % 1000000, ?
+ ?, ?, ?
)
RETURNING id, name, public_key
`
type CreateUserParams struct {
ID snowflake.ID
+ Name string
PublicKey ed25519.PublicKey
}
func (q *Queries) CreateUser(ctx context.Context, arg CreateUserParams) (User, error) {
- row := q.db.QueryRowContext(ctx, createUser, arg.ID, arg.PublicKey)
+ row := q.db.QueryRowContext(ctx, createUser, arg.ID, arg.Name, arg.PublicKey)
var i User
err := row.Scan(&i.ID, &i.Name, &i.PublicKey)
return i, err
}
-const getUser = `-- name: GetUser :one
+const getUserById = `-- name: GetUserById :one
SELECT id, name, public_key FROM users
WHERE id = ?
`
-func (q *Queries) GetUser(ctx context.Context, id snowflake.ID) (User, error) {
- row := q.db.QueryRowContext(ctx, getUser, id)
+func (q *Queries) GetUserById(ctx context.Context, id snowflake.ID) (User, error) {
+ row := q.db.QueryRowContext(ctx, getUserById, id)
+ var i User
+ err := row.Scan(&i.ID, &i.Name, &i.PublicKey)
+ return i, err
+}
+
+const getUserByPublicKey = `-- name: GetUserByPublicKey :one
+SELECT id, name, public_key FROM users
+WHERE public_key = ?
+`
+
+func (q *Queries) GetUserByPublicKey(ctx context.Context, publicKey ed25519.PublicKey) (User, error) {
+ row := q.db.QueryRowContext(ctx, getUserByPublicKey, publicKey)
var i User
err := row.Scan(&i.ID, &i.Name, &i.PublicKey)
return i, err
diff --git a/internal/packet/messages.go b/internal/packet/messages.go
index 04929cd..58d12e9 100644
--- a/internal/packet/messages.go
+++ b/internal/packet/messages.go
@@ -14,9 +14,9 @@ func (m *ErrorMessage) Type() PacketType {
}
type SendMessage struct {
- ReceiverID *snowflake.ID
+ ReceiverID *snowflake.ID
FrequencyID *snowflake.ID
- Content string
+ Content string
}
func (m *SendMessage) Type() PacketType {
@@ -41,11 +41,27 @@ func (m *Messages) Type() PacketType {
type GetMessagesRange struct {
FrequencyID *snowflake.ID
- ReceiverID *snowflake.ID
- From *int64
- To *int64
+ ReceiverID *snowflake.ID
+ From *int64
+ To *int64
}
func (m *GetMessagesRange) Type() PacketType {
return PacketGetMessageRange
}
+
+type GetUserByID struct {
+ UserID snowflake.ID
+}
+
+func (m *GetUserByID) Type() PacketType {
+ return PacketGetUserById
+}
+
+type Users struct {
+ Users []data.User
+}
+
+func (m *Users) Type() PacketType {
+ return PacketUsers
+}
diff --git a/internal/packet/packet.go b/internal/packet/packet.go
index c7986cb..58e83b4 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -55,6 +55,8 @@ const (
PacketPushedMessages
PacketGetMessageRange
PacketMessages
+ PacketGetUserById
+ PacketUsers
)
func (t PacketType) String() string {
@@ -69,29 +71,25 @@ func (t PacketType) String() string {
return "PacketGetMessageRange"
case PacketMessages:
return "PacketMessages"
+ case PacketGetUserById:
+ return "PacketGetUserById"
+ case PacketUsers:
+ return "PacketUsers"
default:
return fmt.Sprintf("PacketInvalidType(%v)", byte(t))
}
}
func (e PacketType) IsSupported() bool {
- switch e {
- case PacketError, PacketSendMessage, PacketPushedMessages, PacketGetMessageRange, PacketMessages:
- return true
- default:
- return false
- }
+ return e <= PacketUsers
}
// True for all packets that a server may push passively to the client.
func (e PacketType) IsPush() bool {
switch e {
- case PacketError, PacketSendMessage, PacketGetMessageRange, PacketMessages:
- return false
case PacketPushedMessages:
return true
default:
- assert.Never("should never happen")
return false
}
}
@@ -207,6 +205,10 @@ func (p Packet) DecodedPayload() (Payload, error) {
payload = &GetMessagesRange{}
case PacketMessages:
payload = &Messages{}
+ case PacketGetUserById:
+ payload = &GetUserByID{}
+ case PacketUsers:
+ payload = &Users{}
default:
assert.Never("packet type of a packet struct must always be valid")
}
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 3d4d274..f932fbb 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -2,13 +2,17 @@ package api
import (
"context"
+ "crypto/ed25519"
+ "database/sql"
"log"
+ "strconv"
"strings"
"github.com/kyren223/eko/internal/data"
"github.com/kyren223/eko/internal/packet"
"github.com/kyren223/eko/internal/server/session"
"github.com/kyren223/eko/pkg/assert"
+ "github.com/kyren223/eko/pkg/snowflake"
)
func SendMessage(ctx context.Context, request *packet.SendMessage) packet.Payload {
@@ -51,3 +55,33 @@ func GetMessages(ctx context.Context, request *packet.GetMessagesRange) packet.P
}
return &packet.Messages{Messages: messages}
}
+
+func GetUserById(ctx context.Context, request *packet.GetUserByID) packet.Payload {
+ queries := data.New(db)
+ user, err := queries.GetUserById(ctx, request.UserID)
+ if err == sql.ErrNoRows {
+ return &packet.Users{Users: []data.User{}}
+ }
+ if err != nil {
+ log.Println("database error when retrieving user by id:", err)
+ return &packet.ErrorMessage{Error: "internal server error"}
+ }
+ return &packet.Users{Users: []data.User{user}}
+}
+
+func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.PublicKey) (data.User, error) {
+ queries := data.New(db)
+ user, err := queries.GetUserByPublicKey(ctx, pubKey)
+ if err == sql.ErrNoRows {
+ id := node.Generate()
+ user, err = queries.CreateUser(ctx, data.CreateUserParams{
+ ID: id,
+ Name: "User" + strconv.FormatInt(id.Time() % 1000, 10),
+ PublicKey: pubKey,
+ })
+ }
+ if err != nil {
+ return data.User{}, err
+ }
+ return user, nil
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index 09f46e4..07cd537 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -147,9 +147,14 @@ func handleConnection(ctx context.Context, conn net.Conn, server *server) {
return
}
- // TODO: replace this with DB query for id
- id := server.Node().Generate()
- sess := session.NewSession(server, addr, id, pubKey)
+ user, err := api.CreateOrGetUser(ctx, server.Node(), pubKey)
+ if err != nil {
+ log.Println(addr, "user creation/fetching error:", err)
+ conn.Close()
+ log.Println(addr, "disconnected")
+ return
+ }
+ sess := session.NewSession(server, addr, user.ID, pubKey)
server.AddSession(sess)
ctx = session.NewContext(ctx, sess)
framer := packet.NewFramer()
@@ -290,11 +295,15 @@ func processRequest(ctx context.Context, request packet.Payload) packet.Payload
assert.Assert(ok, "context in process packet should always have a session")
log.Println(session.Addr(), "processing", request.Type(), "request:", request)
+ // TODO: add a way to measure the time each request/response took and log it
+ // Potentially even separate time for code vs DB operations
switch request := request.(type) {
case *packet.SendMessage:
- return timeout(50 * time.Millisecond, api.SendMessage, ctx, request)
+ return timeout(20*time.Millisecond, api.SendMessage, ctx, request)
case *packet.GetMessagesRange:
- return timeout(100 * time.Millisecond, api.GetMessages, ctx, request)
+ return timeout(50*time.Millisecond, api.GetMessages, ctx, request)
+ case *packet.GetUserByID:
+ return timeout(50*time.Millisecond, api.GetUserById, ctx, request)
default:
return &packet.ErrorMessage{Error: "use of disallowed packet type for request"}
}
diff --git a/internal/server/session/session.go b/internal/server/session/session.go
index ff56780..863b3c3 100644
--- a/internal/server/session/session.go
+++ b/internal/server/session/session.go
@@ -38,11 +38,12 @@ type Session struct {
func NewSession(manager SessionManager, addr *net.TCPAddr, id snowflake.ID, pubKey ed25519.PublicKey) *Session {
session := &Session{
+ WriteQueue: make(chan packet.Packet, 10),
+ PubKey: pubKey,
manager: manager,
addr: addr,
- PubKey: pubKey,
- WriteQueue: make(chan packet.Packet, 10),
- challenge: make([]byte, 32), // Recommended nonce size
+ id: id,
+ challenge: make([]byte, 32),
}
session.Challenge() // Make sure an initial nonce is generated
return session