diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/client/api/api.go | 35 | ||||
| -rw-r--r-- | internal/client/client.go | 66 | ||||
| -rw-r--r-- | internal/client/gateway/gateway.go | 5 | ||||
| -rw-r--r-- | internal/client/ui/messagebox/messagebox.go | 132 | ||||
| -rw-r--r-- | internal/data/users.sql.go | 23 | ||||
| -rw-r--r-- | internal/packet/messages.go | 26 | ||||
| -rw-r--r-- | internal/packet/packet.go | 20 | ||||
| -rw-r--r-- | internal/server/api/api.go | 34 | ||||
| -rw-r--r-- | internal/server/server.go | 19 | ||||
| -rw-r--r-- | internal/server/session/session.go | 7 |
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 |
