summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2025-02-13 16:49:07 +0200
committerKyren223 <Kyren223@proton.me>2025-02-13 16:49:07 +0200
commitcfc5fa3d5e9036400b33a2835cc284f26a872ad9 (patch)
tree9d2ba0ba3d87ad98443383358075e6e299e8bb77
parentab4b010e95360405d344e2859df423dc9d3b02fe (diff)
Users that are no longer members (left/kicked/banned) now get fetched
and show their names properly in messages
-rw-r--r--internal/client/ui/core/core.go3
-rw-r--r--internal/client/ui/core/state/state.go15
-rw-r--r--internal/data/users.sql.go46
-rw-r--r--internal/packet/models.go1
-rw-r--r--internal/packet/packet.go8
-rw-r--r--internal/packet/types.go16
-rw-r--r--internal/server/api/api.go24
-rw-r--r--internal/server/server.go3
-rw-r--r--query/users.sql4
9 files changed, 120 insertions, 0 deletions
diff --git a/internal/client/ui/core/core.go b/internal/client/ui/core/core.go
index b46a22b..fce8414 100644
--- a/internal/client/ui/core/core.go
+++ b/internal/client/ui/core/core.go
@@ -332,6 +332,9 @@ func (m *Model) updateConnected(message tea.Msg) tea.Cmd {
case *packet.BlockInfo:
state.UpdateBlockedUsers(msg)
+ case *packet.UsersInfo:
+ state.UpdateUsersInfo(msg)
+
case *packet.NotificationsInfo:
signals := state.UpdateNotifications(msg)
if m.networkList.Index() == networklist.SignalsIndex {
diff --git a/internal/client/ui/core/state/state.go b/internal/client/ui/core/state/state.go
index da894ed..75fd8a7 100644
--- a/internal/client/ui/core/state/state.go
+++ b/internal/client/ui/core/state/state.go
@@ -144,6 +144,7 @@ func UpdateMessages(info *packet.MessagesInfo) {
}
}
+ unknownUsers := []snowflake.ID{}
for _, message := range info.Messages {
msgSource := message.FrequencyID
if msgSource == nil {
@@ -160,8 +161,16 @@ func UpdateMessages(info *packet.MessagesInfo) {
State.Messages[*msgSource] = bt
}
bt.ReplaceOrInsert(message)
+
+ if _, ok := State.Users[message.SenderID]; !ok {
+ unknownUsers = append(unknownUsers, message.SenderID)
+ }
}
+ gateway.SendAsync(&packet.GetUsers{
+ Users: unknownUsers,
+ })
+
// Note: this is a naive approach
// Ideally we check each message that was added/removed
// For the frequency/receiver/sender id and only remove that
@@ -360,3 +369,9 @@ func UpdateBlockedUsers(info *packet.BlockInfo) {
State.BlockingUsers[blocking] = struct{}{}
}
}
+
+func UpdateUsersInfo(info *packet.UsersInfo) {
+ for _, user := range info.Users {
+ State.Users[user.ID] = user
+ }
+}
diff --git a/internal/data/users.sql.go b/internal/data/users.sql.go
index b68a207..dad7aea 100644
--- a/internal/data/users.sql.go
+++ b/internal/data/users.sql.go
@@ -7,6 +7,7 @@ package data
import (
"context"
+ "strings"
"crypto/ed25519"
"github.com/kyren223/eko/pkg/snowflake"
@@ -102,6 +103,51 @@ func (q *Queries) GetUserData(ctx context.Context, userID snowflake.ID) (string,
return data, err
}
+const getUsersByIds = `-- name: GetUsersByIds :many
+SELECT id, name, public_key, description, is_public_dm, is_deleted FROM users
+WHERE id IN (/*SLICE:ids*/?)
+`
+
+func (q *Queries) GetUsersByIds(ctx context.Context, ids []snowflake.ID) ([]User, error) {
+ query := getUsersByIds
+ var queryParams []interface{}
+ if len(ids) > 0 {
+ for _, v := range ids {
+ queryParams = append(queryParams, v)
+ }
+ query = strings.Replace(query, "/*SLICE:ids*/?", strings.Repeat(",?", len(ids))[1:], 1)
+ } else {
+ query = strings.Replace(query, "/*SLICE:ids*/?", "NULL", 1)
+ }
+ rows, err := q.db.QueryContext(ctx, query, queryParams...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ var items []User
+ for rows.Next() {
+ var i User
+ if err := rows.Scan(
+ &i.ID,
+ &i.Name,
+ &i.PublicKey,
+ &i.Description,
+ &i.IsPublicDM,
+ &i.IsDeleted,
+ ); err != nil {
+ return nil, err
+ }
+ items = append(items, i)
+ }
+ if err := rows.Close(); err != nil {
+ return nil, err
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
const setUserData = `-- name: SetUserData :one
INSERT INTO user_data (
user_id, data
diff --git a/internal/packet/models.go b/internal/packet/models.go
index c26d93a..9d7e885 100644
--- a/internal/packet/models.go
+++ b/internal/packet/models.go
@@ -13,6 +13,7 @@ const (
MaxUsernameBytes = 32
MaxUserDescriptionBytes = 200
MaxBanReasonBytes = 64
+ MaxUsersInGetUsers = 64
)
const (
diff --git a/internal/packet/packet.go b/internal/packet/packet.go
index e94086b..fb3e7e8 100644
--- a/internal/packet/packet.go
+++ b/internal/packet/packet.go
@@ -86,6 +86,9 @@ const (
PacketBlockUser
PacketBlockInfo
+ PacketGetUsers
+ PacketUsersInfo
+
PacketMax
)
@@ -261,6 +264,11 @@ func (p Packet) DecodedPayload() (Payload, error) {
case PacketBlockInfo:
payload = &BlockInfo{}
+ case PacketGetUsers:
+ payload = &GetUsers{}
+ case PacketUsersInfo:
+ payload = &UsersInfo{}
+
default:
assert.Never("unexpected packet.PacketType", "type", p.Type())
}
diff --git a/internal/packet/types.go b/internal/packet/types.go
index 572d1ed..f251450 100644
--- a/internal/packet/types.go
+++ b/internal/packet/types.go
@@ -272,3 +272,19 @@ type BlockInfo struct {
func (m *BlockInfo) Type() PacketType {
return PacketBlockInfo
}
+
+type GetUsers struct {
+ Users []snowflake.ID
+}
+
+func (m *GetUsers) Type() PacketType {
+ return PacketGetUsers
+}
+
+type UsersInfo struct {
+ Users []data.User
+}
+
+func (m *UsersInfo) Type() PacketType {
+ return PacketUsersInfo
+}
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 6b33ab8..fe5542a 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -288,6 +288,11 @@ func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.P
}
func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.CreateNetwork) packet.Payload {
+ // TODO: implement private servers
+ if !request.IsPublic {
+ return &ErrNotImplemented
+ }
+
name := strings.TrimSpace(request.Name)
if name == "" {
return &packet.Error{Error: "server name must not be blank"}
@@ -1536,3 +1541,22 @@ func GetBlockedUsers(ctx context.Context, sess *session.Session) packet.Payload
RemovedBlockingUsers: nil,
}
}
+
+func GetUsers(ctx context.Context, sess *session.Session, request *packet.GetUsers) packet.Payload {
+ if len(request.Users) > packet.MaxUsersInGetUsers {
+ return &packet.Error{Error: fmt.Sprintf(
+ "Max users per request may not exceed %v", packet.MaxUsersInGetUsers,
+ )}
+ }
+
+ queries := data.New(db)
+ users, err := queries.GetUsersByIds(ctx, request.Users)
+ if err != nil {
+ log.Println("database error 1:", err)
+ return &ErrInternalError
+ }
+
+ return &packet.UsersInfo{
+ Users: users,
+ }
+}
diff --git a/internal/server/server.go b/internal/server/server.go
index f4224c1..aa34a78 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -374,6 +374,9 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P
case *packet.BlockUser:
response = timeout(10*time.Millisecond, api.BlockUser, ctx, sess, request)
+ case *packet.GetUsers:
+ response = timeout(10*time.Millisecond, api.GetUsers, ctx, sess, request)
+
default:
response = &packet.Error{Error: "use of disallowed packet type for request"}
}
diff --git a/query/users.sql b/query/users.sql
index bd91cda..d2a4a92 100644
--- a/query/users.sql
+++ b/query/users.sql
@@ -40,3 +40,7 @@ RETURNING *;
-- name: GetUserData :one
SELECT data FROM user_data
WHERE user_id = ?;
+
+-- name: GetUsersByIds :many
+SELECT * FROM users
+WHERE id IN (sqlc.slice('ids'));