From b09a8fba647ff030b24d9ad12215ba330c13080c Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Thu, 9 Jan 2025 19:14:05 +0200 Subject: Fixed everything on the server-side --- internal/data/members.sql.go | 12 +- internal/data/models.go | 5 + internal/data/users.sql.go | 37 +++ internal/data/users_networks.sql.go | 329 --------------------- internal/packet/models.go | 1 + internal/packet/packet.go | 24 +- internal/packet/types.go | 32 +- internal/server/api/api.go | 141 +++++---- internal/server/api/database.go | 1 - internal/server/api/helpers.go | 15 +- .../server/api/migrations/20250109134844_users.sql | 6 + .../migrations/20250109143707_on_user_delete.sql | 2 +- internal/server/server.go | 13 +- 13 files changed, 195 insertions(+), 423 deletions(-) delete mode 100644 internal/data/users_networks.sql.go (limited to 'internal') diff --git a/internal/data/members.sql.go b/internal/data/members.sql.go index c9932cc..fd66dba 100644 --- a/internal/data/members.sql.go +++ b/internal/data/members.sql.go @@ -82,7 +82,7 @@ func (q *Queries) GetMemberById(ctx context.Context, arg GetMemberByIdParams) (M return i, err } -const getMembers = `-- name: GetMembers :many +const getNetworkMembers = `-- name: GetNetworkMembers :many SELECT users.id, users.name, users.public_key, users.description, users.is_public_dm, users.is_deleted, members.user_id, members.network_id, members.joined_at, members.is_member, members.is_admin, members.is_muted, members.is_banned, members.ban_reason @@ -91,20 +91,20 @@ JOIN users ON users.id = members.user_id WHERE network_id = ? AND is_member = true ` -type GetMembersRow struct { +type GetNetworkMembersRow struct { User User Member Member } -func (q *Queries) GetMembers(ctx context.Context, networkID snowflake.ID) ([]GetMembersRow, error) { - rows, err := q.db.QueryContext(ctx, getMembers, networkID) +func (q *Queries) GetNetworkMembers(ctx context.Context, networkID snowflake.ID) ([]GetNetworkMembersRow, error) { + rows, err := q.db.QueryContext(ctx, getNetworkMembers, networkID) if err != nil { return nil, err } defer rows.Close() - var items []GetMembersRow + var items []GetNetworkMembersRow for rows.Next() { - var i GetMembersRow + var i GetNetworkMembersRow if err := rows.Scan( &i.User.ID, &i.User.Name, diff --git a/internal/data/models.go b/internal/data/models.go index bb0b271..8815e0b 100644 --- a/internal/data/models.go +++ b/internal/data/models.go @@ -67,3 +67,8 @@ type User struct { IsPublicDM bool IsDeleted bool } + +type UserData struct { + UserID snowflake.ID + Data string +} diff --git a/internal/data/users.sql.go b/internal/data/users.sql.go index f7a7f9a..db28ca7 100644 --- a/internal/data/users.sql.go +++ b/internal/data/users.sql.go @@ -89,3 +89,40 @@ func (q *Queries) GetUserByPublicKey(ctx context.Context, publicKey ed25519.Publ ) return i, err } + +const getUserData = `-- name: GetUserData :one +SELECT data FROM user_data +WHERE user_id = ? +` + +func (q *Queries) GetUserData(ctx context.Context, userID snowflake.ID) (string, error) { + row := q.db.QueryRowContext(ctx, getUserData, userID) + var data string + err := row.Scan(&data) + return data, err +} + +const setUserData = `-- name: SetUserData :one +INSERT INTO user_data ( + user_id, data +) VALUES ( + ?1, ?2 +) +ON CONFLICT DO +UPDATE SET + user_id = EXCLUDED.user_id, data = EXCLUDED.data +WHERE user_id = EXCLUDED.user_id +RETURNING user_id, data +` + +type SetUserDataParams struct { + UserID snowflake.ID + Data string +} + +func (q *Queries) SetUserData(ctx context.Context, arg SetUserDataParams) (UserData, error) { + row := q.db.QueryRowContext(ctx, setUserData, arg.UserID, arg.Data) + var i UserData + err := row.Scan(&i.UserID, &i.Data) + return i, err +} diff --git a/internal/data/users_networks.sql.go b/internal/data/users_networks.sql.go deleted file mode 100644 index e180b51..0000000 --- a/internal/data/users_networks.sql.go +++ /dev/null @@ -1,329 +0,0 @@ -// Code generated by sqlc. DO NOT EDIT. -// versions: -// sqlc v1.27.0 -// source: users_networks.sql - -package data - -import ( - "context" - "strings" - - "github.com/kyren223/eko/pkg/snowflake" -) - -const filterUsersInNetwork = `-- name: FilterUsersInNetwork :many -SELECT user_id FROM users_networks -WHERE network_id = ? AND user_id IN (/*SLICE:users*/?) -` - -type FilterUsersInNetworkParams struct { - NetworkID snowflake.ID - Users []snowflake.ID -} - -func (q *Queries) FilterUsersInNetwork(ctx context.Context, arg FilterUsersInNetworkParams) ([]snowflake.ID, error) { - query := filterUsersInNetwork - var queryParams []interface{} - queryParams = append(queryParams, arg.NetworkID) - if len(arg.Users) > 0 { - for _, v := range arg.Users { - queryParams = append(queryParams, v) - } - query = strings.Replace(query, "/*SLICE:users*/?", strings.Repeat(",?", len(arg.Users))[1:], 1) - } else { - query = strings.Replace(query, "/*SLICE:users*/?", "NULL", 1) - } - rows, err := q.db.QueryContext(ctx, query, queryParams...) - if err != nil { - return nil, err - } - defer rows.Close() - var items []snowflake.ID - for rows.Next() { - var user_id snowflake.ID - if err := rows.Scan(&user_id); err != nil { - return nil, err - } - items = append(items, user_id) - } - if err := rows.Close(); err != nil { - return nil, err - } - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -const getNetworkBannedUsers = `-- name: GetNetworkBannedUsers :many -SELECT - users.id, users.name, users.public_key, users.description, users.is_public_dm, users.is_deleted, - users_networks.ban_reason -FROM users_networks -JOIN users ON users.id = users_networks.user_id -WHERE users_networks.network_id = ? -` - -type GetNetworkBannedUsersRow struct { - User User - BanReason *string -} - -func (q *Queries) GetNetworkBannedUsers(ctx context.Context, networkID snowflake.ID) ([]GetNetworkBannedUsersRow, error) { - rows, err := q.db.QueryContext(ctx, getNetworkBannedUsers, networkID) - if err != nil { - return nil, err - } - defer rows.Close() - var items []GetNetworkBannedUsersRow - for rows.Next() { - var i GetNetworkBannedUsersRow - if err := rows.Scan( - &i.User.ID, - &i.User.Name, - &i.User.PublicKey, - &i.User.Description, - &i.User.IsPublicDM, - &i.User.IsDeleted, - &i.BanReason, - ); 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 getNetworkMemberById = `-- name: GetNetworkMemberById :one -SELECT user_id, network_id, joined_at, is_member, is_admin, is_muted, is_banned, ban_reason, position -FROM users_networks -WHERE users_networks.network_id = ? AND users_networks.user_id = ? -` - -type GetNetworkMemberByIdParams struct { - NetworkID snowflake.ID - UserID snowflake.ID -} - -func (q *Queries) GetNetworkMemberById(ctx context.Context, arg GetNetworkMemberByIdParams) (UserNetwork, error) { - row := q.db.QueryRowContext(ctx, getNetworkMemberById, arg.NetworkID, arg.UserID) - var i UserNetwork - err := row.Scan( - &i.UserID, - &i.NetworkID, - &i.JoinedAt, - &i.IsMember, - &i.IsAdmin, - &i.IsMuted, - &i.IsBanned, - &i.BanReason, - &i.Position, - ) - return i, err -} - -const getNetworkMembers = `-- name: GetNetworkMembers :many -SELECT - users.id, users.name, users.public_key, users.description, users.is_public_dm, users.is_deleted, - users_networks.joined_at, - users_networks.is_admin, - users_networks.is_muted -FROM users_networks -JOIN users ON users.id = users_networks.user_id -WHERE users_networks.network_id = ? AND is_member = true -` - -type GetNetworkMembersRow struct { - User User - JoinedAt string - IsAdmin bool - IsMuted bool -} - -func (q *Queries) GetNetworkMembers(ctx context.Context, networkID snowflake.ID) ([]GetNetworkMembersRow, error) { - rows, err := q.db.QueryContext(ctx, getNetworkMembers, networkID) - if err != nil { - return nil, err - } - defer rows.Close() - var items []GetNetworkMembersRow - for rows.Next() { - var i GetNetworkMembersRow - if err := rows.Scan( - &i.User.ID, - &i.User.Name, - &i.User.PublicKey, - &i.User.Description, - &i.User.IsPublicDM, - &i.User.IsDeleted, - &i.JoinedAt, - &i.IsAdmin, - &i.IsMuted, - ); 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 getUserNetwork = `-- name: GetUserNetwork :one -SELECT user_id, network_id, joined_at, is_member, is_admin, is_muted, is_banned, ban_reason, position FROM users_networks -WHERE user_id = ? AND network_id = ? -` - -type GetUserNetworkParams struct { - UserID snowflake.ID - NetworkID snowflake.ID -} - -func (q *Queries) GetUserNetwork(ctx context.Context, arg GetUserNetworkParams) (UserNetwork, error) { - row := q.db.QueryRowContext(ctx, getUserNetwork, arg.UserID, arg.NetworkID) - var i UserNetwork - err := row.Scan( - &i.UserID, - &i.NetworkID, - &i.JoinedAt, - &i.IsMember, - &i.IsAdmin, - &i.IsMuted, - &i.IsBanned, - &i.BanReason, - &i.Position, - ) - return i, err -} - -const getUserNetworks = `-- name: GetUserNetworks :many -SELECT networks.id, networks.owner_id, networks.name, networks.icon, networks.bg_hex_color, networks.fg_hex_color, networks.is_public, users_networks.position FROM networks -JOIN users_networks ON networks.id = users_networks.network_id -WHERE users_networks.user_id = ? -ORDER BY users_networks.position -` - -type GetUserNetworksRow struct { - Network Network - Position *int64 -} - -func (q *Queries) GetUserNetworks(ctx context.Context, userID snowflake.ID) ([]GetUserNetworksRow, error) { - rows, err := q.db.QueryContext(ctx, getUserNetworks, userID) - if err != nil { - return nil, err - } - defer rows.Close() - var items []GetUserNetworksRow - for rows.Next() { - var i GetUserNetworksRow - if err := rows.Scan( - &i.Network.ID, - &i.Network.OwnerID, - &i.Network.Name, - &i.Network.Icon, - &i.Network.BgHexColor, - &i.Network.FgHexColor, - &i.Network.IsPublic, - &i.Position, - ); 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 setMember = `-- name: SetMember :one -INSERT INTO users_networks ( - user_id, network_id, - is_member, is_admin, is_muted, - is_banned, ban_reason, position -) VALUES ( - ?1, ?2, - ?3, ?4, ?5, - ?6, ?7, - CASE - WHEN ?3 = false THEN NULL - ELSE (SELECT COUNT(*) FROM users_networks WHERE user_id = ?1) - END -) -ON CONFLICT DO -UPDATE SET - is_member = EXCLUDED.is_member, is_admin = EXCLUDED.is_admin, is_muted = EXCLUDED.is_muted, - is_banned = EXCLUDED.is_banned, ban_reason = EXCLUDED.ban_reason, position = EXCLUDED.position -WHERE user_id = EXCLUDED.user_id AND network_id = EXCLUDED.network_id -RETURNING user_id, network_id, joined_at, is_member, is_admin, is_muted, is_banned, ban_reason, position -` - -type SetMemberParams struct { - UserID snowflake.ID - NetworkID snowflake.ID - IsMember bool - IsAdmin bool - IsMuted bool - IsBanned bool - BanReason *string -} - -func (q *Queries) SetMember(ctx context.Context, arg SetMemberParams) (UserNetwork, error) { - row := q.db.QueryRowContext(ctx, setMember, - arg.UserID, - arg.NetworkID, - arg.IsMember, - arg.IsAdmin, - arg.IsMuted, - arg.IsBanned, - arg.BanReason, - ) - var i UserNetwork - err := row.Scan( - &i.UserID, - &i.NetworkID, - &i.JoinedAt, - &i.IsMember, - &i.IsAdmin, - &i.IsMuted, - &i.IsBanned, - &i.BanReason, - &i.Position, - ) - return i, err -} - -const swapUserNetworks = `-- name: SwapUserNetworks :exec -UPDATE users_networks SET - position = CASE - WHEN position = ?1 THEN ?2 - WHEN position = ?2 THEN ?1 - END -WHERE user_id = ?3 AND position IN (?1, ?2) -` - -type SwapUserNetworksParams struct { - Pos1 *int64 - Pos2 *int64 - UserID snowflake.ID -} - -func (q *Queries) SwapUserNetworks(ctx context.Context, arg SwapUserNetworksParams) error { - _, err := q.db.ExecContext(ctx, swapUserNetworks, arg.Pos1, arg.Pos2, arg.UserID) - return err -} diff --git a/internal/packet/models.go b/internal/packet/models.go index 562e782..f7162d7 100644 --- a/internal/packet/models.go +++ b/internal/packet/models.go @@ -5,6 +5,7 @@ const ( DefaultFrequencyName = "main" DefaultFrequencyColor = "#FFFFFF" MaxFrequencyName = 32 + MaxUserDataBytes = 8192 ) const ( diff --git a/internal/packet/packet.go b/internal/packet/packet.go index 0067c9a..57260da 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -52,16 +52,15 @@ type PacketType uint8 const ( PacketError PacketType = iota + PacketSetUserData + PacketGetUserData + PacketCreateNetwork PacketUpdateNetwork PacketTransferNetwork PacketDeleteNetwork - PacketSwapUserNetworks PacketNetworksInfo - PacketSetMember - PacketMembersInfo - PacketCreateFrequency PacketUpdateFrequency PacketDeleteFrequency @@ -74,6 +73,9 @@ const ( PacketRequestMessages PacketMessagesInfo + PacketSetMember + PacketMembersInfo + PacketMax ) @@ -189,6 +191,11 @@ func (p Packet) DecodedPayload() (Payload, error) { case PacketError: payload = &Error{} + case PacketSetUserData: + payload = &SetUserData{} + case PacketGetUserData: + payload = &GetUserData{} + case PacketCreateNetwork: payload = &CreateNetwork{} case PacketUpdateNetwork: @@ -197,10 +204,6 @@ func (p Packet) DecodedPayload() (Payload, error) { payload = &TransferNetwork{} case PacketDeleteNetwork: payload = &DeleteNetwork{} - case PacketSwapUserNetworks: - payload = &SwapUserNetworks{} - case PacketSetMember: - payload = &SetMember{} case PacketNetworksInfo: payload = &NetworksInfo{} @@ -226,6 +229,11 @@ func (p Packet) DecodedPayload() (Payload, error) { case PacketMessagesInfo: payload = &MessagesInfo{} + case PacketSetMember: + payload = &SetMember{} + case PacketMembersInfo: + payload = &MembersInfo{} + default: assert.Never("unexpected packet.PacketType", "type", p.Type()) } diff --git a/internal/packet/types.go b/internal/packet/types.go index b107243..370de21 100644 --- a/internal/packet/types.go +++ b/internal/packet/types.go @@ -52,15 +52,6 @@ func (m *DeleteNetwork) Type() PacketType { return PacketDeleteNetwork } -type SwapUserNetworks struct { - Pos1 int - Pos2 int -} - -func (m *SwapUserNetworks) Type() PacketType { - return PacketSwapUserNetworks -} - type SetMember struct { Member *bool Admin *bool @@ -78,8 +69,8 @@ func (m *SetMember) Type() PacketType { type FullNetwork struct { data.Network Frequencies []data.Frequency - Members []data.GetNetworkMembersRow - Position int + Members []data.Member + Users []data.User } type NetworksInfo struct { @@ -189,10 +180,27 @@ func (m *MessagesInfo) Type() PacketType { type MembersInfo struct { RemovedMembers []snowflake.ID - Members []data.GetNetworkMembersRow + Members []data.Member + Users []data.User Network snowflake.ID } func (m *MembersInfo) Type() PacketType { return PacketMessagesInfo } + +type SetUserData struct{ + Data string +} + +func (m *SetUserData) Type() PacketType { + return PacketSetUserData +} + +type GetUserData struct{ + Data string +} + +func (m *GetUserData) Type() PacketType { + return PacketGetUserData +} diff --git a/internal/server/api/api.go b/internal/server/api/api.go index cfba30a..28d7072 100644 --- a/internal/server/api/api.go +++ b/internal/server/api/api.go @@ -43,7 +43,10 @@ func SendMessage(ctx context.Context, sess *session.Session, request *packet.Sen return &ErrInternalError } - return &packet.MessagesInfo{Messages: []data.Message{message}} + return &packet.MessagesInfo{ + Messages: []data.Message{message}, + RemovedMessages: nil, + } } func RequestMessages(ctx context.Context, sess *session.Session, request *packet.RequestMessages) packet.Payload { @@ -59,14 +62,27 @@ func RequestMessages(ctx context.Context, sess *session.Session, request *packet User2: request.ReceiverID, }) } else { - return &packet.Error{Error: "either receiver id or frequency id must exist"} + return &packet.Error{Error: "either receiver id or frequency id must be specified"} } if err != nil { log.Println("database error when retrieving messages:", err) return &packet.Error{Error: "internal server error"} } - return &packet.MessagesInfo{Messages: messages} + + // TODO: + // FIXME: If an attacker knows the frequency ID of a private frequency + // Such in a case where they used to have access to it but no longer do + // Then it's possible to view all messages including current ones + // Add a check to make sure the user is inside the network of this + // Frequency, and even if it is, if the frequency is "no access" then + // an admin-check needs to happen, and if he's not an admin it needs to be + // denied + + return &packet.MessagesInfo{ + Messages: messages, + RemovedMessages: nil, + } } func CreateOrGetUser(ctx context.Context, node *snowflake.Node, pubKey ed25519.PublicKey) (data.User, error) { @@ -107,7 +123,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C log.Println("database error:", err) return &ErrInternalError } - defer tx.Rollback() //nolint + defer tx.Rollback() queries := data.New(db) qtx := queries.WithTx(tx) @@ -138,7 +154,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C return &ErrInternalError } - networkUser, err := qtx.SetMember(ctx, data.SetMemberParams{ + member, err := qtx.SetMember(ctx, data.SetMemberParams{ UserID: network.OwnerID, NetworkID: network.ID, IsMember: true, @@ -167,13 +183,8 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C fullNetwork := packet.FullNetwork{ Network: network, Frequencies: []data.Frequency{frequency}, - Members: []data.GetNetworkMembersRow{{ - JoinedAt: networkUser.JoinedAt, - User: user, - IsAdmin: networkUser.IsAdmin, - IsMuted: networkUser.IsMuted, - }}, - Position: int(*networkUser.Position), + Members: []data.Member{member}, + Users: []data.User{user}, } return &packet.NetworksInfo{ Networks: []packet.FullNetwork{fullNetwork}, @@ -189,40 +200,33 @@ func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload if err != nil { return nil, err } - defer tx.Rollback() //nolint + defer tx.Rollback() queries := data.New(db) qtx := queries.WithTx(tx) - networks, err := qtx.GetUserNetworks(ctx, sess.ID()) + networks, err := qtx.GetNetworksOfUser(ctx, sess.ID()) if err != nil { return nil, err } - for _, userNetwork := range networks { - network := userNetwork.Network - - // User was in this network but is no longer there (left/kicked/banned) - if userNetwork.Position == nil { - continue - } - - position := int(*userNetwork.Position) + for _, network := range networks { frequencies, err := qtx.GetNetworkFrequencies(ctx, network.ID) if err != nil { return nil, err } - members, err := qtx.GetNetworkMembers(ctx, network.ID) + membersAndUsers, err := qtx.GetNetworkMembers(ctx, network.ID) if err != nil { return nil, err } + members, users := SplitMembersAndUsers(membersAndUsers) fullNetworks = append(fullNetworks, packet.FullNetwork{ Network: network, Frequencies: frequencies, Members: members, - Position: position, + Users: users, }) } @@ -238,22 +242,6 @@ func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload }, nil } -func SwapUserNetworks(ctx context.Context, sess *session.Session, request *packet.SwapUserNetworks) packet.Payload { - queries := data.New(db) - pos1, pos2 := int64(request.Pos1), int64(request.Pos2) - err := queries.SwapUserNetworks(ctx, data.SwapUserNetworksParams{ - Pos1: &pos1, - Pos2: &pos2, - UserID: sess.ID(), - }) - if err != nil { - log.Println("database error:", err) - return &ErrInternalError - } - - return request -} - func CreateFrequency(ctx context.Context, sess *session.Session, request *packet.CreateFrequency) packet.Payload { queries := data.New(db) @@ -425,7 +413,7 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe return &ErrInternalError } - member, err := queries.GetNetworkMemberById(ctx, data.GetNetworkMemberByIdParams{ + member, err := queries.GetMemberById(ctx, data.GetMemberByIdParams{ NetworkID: request.Network, UserID: request.User, }) @@ -454,13 +442,9 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe NetworkPropagate(ctx, sess, request.Network, &packet.MembersInfo{ RemovedMembers: nil, - Members: []data.GetNetworkMembersRow{{ - User: user, - JoinedAt: newMember.JoinedAt, - IsAdmin: newMember.IsAdmin, - IsMuted: newMember.IsMuted, - }}, - Network: request.Network, + Members: []data.Member{newMember}, + Users: []data.User{user}, + Network: request.Network, }) frequencies, err := queries.GetNetworkFrequencies(ctx, network.ID) @@ -469,18 +453,19 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe return &ErrInternalError } - members, err := queries.GetNetworkMembers(ctx, network.ID) + membersAndUsers, err := queries.GetNetworkMembers(ctx, network.ID) if err != nil { log.Println("database error 5:", err) return &ErrInternalError } + members, users := SplitMembersAndUsers(membersAndUsers) return &packet.NetworksInfo{ Networks: []packet.FullNetwork{{ Network: network, Frequencies: frequencies, Members: members, - Position: int(*newMember.Position), + Users: users, }}, RemovedNetworks: nil, Set: false, @@ -553,6 +538,7 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe return NetworkPropagate(ctx, sess, request.Network, &packet.MembersInfo{ RemovedMembers: []snowflake.ID{newMember.UserID}, Members: nil, + Users: nil, Network: request.Network, }) } @@ -565,13 +551,9 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe payload := NetworkPropagate(ctx, sess, request.Network, &packet.MembersInfo{ RemovedMembers: nil, - Members: []data.GetNetworkMembersRow{{ - User: user, - JoinedAt: newMember.JoinedAt, - IsAdmin: newMember.IsAdmin, - IsMuted: newMember.IsMuted, - }}, - Network: request.Network, + Members: []data.Member{newMember}, + Users: []data.User{user}, + Network: request.Network, }) // Joined @@ -582,18 +564,19 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe return &ErrInternalError } - members, err := queries.GetNetworkMembers(ctx, network.ID) + membersAndUsers, err := queries.GetNetworkMembers(ctx, network.ID) if err != nil { log.Println("database error 11:", err) return &ErrInternalError } + members, users := SplitMembersAndUsers(membersAndUsers) return &packet.NetworksInfo{ Networks: []packet.FullNetwork{{ Network: network, Frequencies: frequencies, Members: members, - Position: int(*newMember.Position), + Users: users, }}, RemovedNetworks: nil, Set: false, @@ -603,3 +586,41 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe // Normal case, was already in the server return payload } + +func SetUserData(ctx context.Context, sess *session.Session, request *packet.SetUserData) packet.Payload { + queries := data.New(db) + + if len(request.Data) > packet.MaxUserDataBytes { + return &packet.Error{ + Error: "data bytes may not exceed " + + strconv.FormatInt(packet.MaxUserDataBytes, 10) + " bytes", + } + } + + _, err := queries.SetUserData(ctx, data.SetUserDataParams{ + UserID: sess.ID(), + Data: request.Data, + }) + if err != nil { + log.Println("database error:", err) + return &ErrInternalError + } + + return &packet.GetUserData{ + Data: request.Data, + } +} + +func GetUserData(ctx context.Context, sess *session.Session, request *packet.GetUserData) packet.Payload { + queries := data.New(db) + + data, err := queries.GetUserData(ctx, sess.ID()) + if err != nil { + log.Println("database error:", err) + return &ErrInternalError + } + + return &packet.GetUserData{ + Data: data, + } +} diff --git a/internal/server/api/database.go b/internal/server/api/database.go index 00e8d4d..b9d7413 100644 --- a/internal/server/api/database.go +++ b/internal/server/api/database.go @@ -46,7 +46,6 @@ func ConnectToDatabase() { db.Close() log.Fatalln("error running up migrations:", err) } - log.Println("database up migrations applied successfully") log.Println("database connection ready to be used") } diff --git a/internal/server/api/helpers.go b/internal/server/api/helpers.go index 0ab6ae5..f98e2ad 100644 --- a/internal/server/api/helpers.go +++ b/internal/server/api/helpers.go @@ -33,9 +33,9 @@ func isValidHexColor(color string) (bool, string) { } func IsNetworkAdmin(ctx context.Context, queries *data.Queries, userId, networkId snowflake.ID) (bool, error) { - userNetwork, err := queries.GetUserNetwork(ctx, data.GetUserNetworkParams{ - UserID: userId, + userNetwork, err := queries.GetMemberById(ctx, data.GetMemberByIdParams{ NetworkID: networkId, + UserID: userId, }) if err != nil { return false, err @@ -87,3 +87,14 @@ func NetworkPropagate( return payload } + +func SplitMembersAndUsers(membersAndUsers []data.GetNetworkMembersRow) ([]data.Member, []data.User) { + members := make([]data.Member, 0, len(membersAndUsers)) + users := make([]data.User, 0, len(membersAndUsers)) + for _, memberAndUser := range membersAndUsers { + members = append(members, memberAndUser.Member) + users = append(users, memberAndUser.User) + } + + return members, users +} diff --git a/internal/server/api/migrations/20250109134844_users.sql b/internal/server/api/migrations/20250109134844_users.sql index 8ebb9eb..c22cc19 100644 --- a/internal/server/api/migrations/20250109134844_users.sql +++ b/internal/server/api/migrations/20250109134844_users.sql @@ -8,5 +8,11 @@ CREATE TABLE IF NOT EXISTS users ( is_deleted BOOLEAN NOT NULL CHECK (is_deleted IN (false, true)) DEFAULT false ); +CREATE TABLE IF NOT EXISTS user_data ( + user_id INTEGER PRIMARY KEY REFERENCES users(id), + data TEXT NOT NULL +); + -- +goose Down DROP TABLE IF EXISTS users; +DROP TABLE IF EXISTS user_data; diff --git a/internal/server/api/migrations/20250109143707_on_user_delete.sql b/internal/server/api/migrations/20250109143707_on_user_delete.sql index 940bec4..9e9bdde 100644 --- a/internal/server/api/migrations/20250109143707_on_user_delete.sql +++ b/internal/server/api/migrations/20250109143707_on_user_delete.sql @@ -1,6 +1,6 @@ -- +goose Up -- +goose StatementBegin -CREATE TRIGGER on_user_delete +CREATE TRIGGER IF NOT EXISTS on_user_delete AFTER UPDATE OF is_deleted ON users WHEN NEW.is_deleted = true BEGIN diff --git a/internal/server/server.go b/internal/server/server.go index 56c3a50..9abf51f 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -332,14 +332,16 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P // Potentially even separate time for code vs DB operations var response packet.Payload switch request := request.(type) { + + case *packet.SetUserData: + response = timeout(5*time.Millisecond, api.SetUserData, ctx, sess, request) + case *packet.GetUserData: + response = timeout(5*time.Millisecond, api.GetUserData, ctx, sess, request) + case *packet.CreateNetwork: response = timeout(10*time.Millisecond, api.CreateNetwork, ctx, sess, request) case *packet.DeleteNetwork: response = timeout(500*time.Millisecond, api.DeleteNetwork, ctx, sess, request) - case *packet.SwapUserNetworks: - response = timeout(5*time.Millisecond, api.SwapUserNetworks, ctx, sess, request) - case *packet.SetMember: - response = timeout(50*time.Millisecond, api.SetMember, ctx, sess, request) case *packet.CreateFrequency: response = timeout(5*time.Millisecond, api.CreateFrequency, ctx, sess, request) @@ -353,6 +355,9 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P case *packet.RequestMessages: response = timeout(50*time.Millisecond, api.RequestMessages, ctx, sess, request) + case *packet.SetMember: + response = timeout(50*time.Millisecond, api.SetMember, ctx, sess, request) + default: response = &packet.Error{Error: "use of disallowed packet type for request"} } -- cgit v1.3.1