diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/client/ui/core/core.go | 3 | ||||
| -rw-r--r-- | internal/client/ui/core/state/state.go | 15 | ||||
| -rw-r--r-- | internal/data/users.sql.go | 46 | ||||
| -rw-r--r-- | internal/packet/models.go | 1 | ||||
| -rw-r--r-- | internal/packet/packet.go | 8 | ||||
| -rw-r--r-- | internal/packet/types.go | 16 | ||||
| -rw-r--r-- | internal/server/api/api.go | 24 | ||||
| -rw-r--r-- | internal/server/server.go | 3 |
8 files changed, 116 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"} } |
