diff options
| -rw-r--r-- | internal/client/ui/core/core.go | 25 | ||||
| -rw-r--r-- | internal/data/models.go | 14 | ||||
| -rw-r--r-- | internal/data/users_networks.sql.go | 31 | ||||
| -rw-r--r-- | internal/packet/models.go | 1 | ||||
| -rw-r--r-- | internal/packet/packet.go | 51 | ||||
| -rw-r--r-- | internal/packet/types.go | 11 | ||||
| -rw-r--r-- | internal/server/api/api.go | 97 | ||||
| -rw-r--r-- | internal/server/server.go | 3 | ||||
| -rw-r--r-- | query/users_networks.sql | 4 | ||||
| -rw-r--r-- | sqlc.yml | 1 |
10 files changed, 190 insertions, 48 deletions
diff --git a/internal/client/ui/core/core.go b/internal/client/ui/core/core.go index 27c6b96..5031dc0 100644 --- a/internal/client/ui/core/core.go +++ b/internal/client/ui/core/core.go @@ -19,6 +19,7 @@ import ( "github.com/kyren223/eko/internal/client/ui/core/networklist" "github.com/kyren223/eko/internal/client/ui/core/state" "github.com/kyren223/eko/internal/client/ui/loadscreen" + "github.com/kyren223/eko/internal/data" "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/pkg/assert" "github.com/kyren223/eko/pkg/snowflake" @@ -165,7 +166,6 @@ func (m *Model) updateConnected(msg tea.Msg) tea.Cmd { if msg.Set { state.State.Networks = msg.Networks } else { - // state.State.Networks = append(state.State.Networks, msg.Networks...) networks := state.State.Networks networks = append(networks, msg.Networks...) networks = slices.DeleteFunc(networks, func(network packet.FullNetwork) bool { @@ -178,6 +178,29 @@ func (m *Model) updateConnected(msg tea.Msg) tea.Cmd { state.State.Networks = networks } + case *packet.FrequenciesInfo: + var network *packet.FullNetwork + for i, fullNetwork := range state.State.Networks { + if fullNetwork.ID == msg.Network { + network = &state.State.Networks[i] + } + } + + if msg.Set { + network.Frequencies = msg.Frequencies + } else { + frequencies := network.Frequencies + frequencies = append(frequencies, msg.Frequencies...) + frequencies = slices.DeleteFunc(frequencies, func(frequency data.Frequency) bool { + return slices.Contains(msg.RemoveFrequencies, frequency.ID) + }) + slices.SortFunc(frequencies, func(a, b data.Frequency) int { + return int(a.Position - b.Position) + }) + log.Println(frequencies) + network.Frequencies = frequencies + } + case ui.QuitMsg: gateway.Disconnect() diff --git a/internal/data/models.go b/internal/data/models.go index 566d2ef..f53d726 100644 --- a/internal/data/models.go +++ b/internal/data/models.go @@ -51,13 +51,7 @@ type UserBlockedUser struct { BlockedUserID snowflake.ID } -type UserTrustedUser struct { - TrusterUserID snowflake.ID - TrustedUserID snowflake.ID - TrustedPublicKey []byte -} - -type UsersNetwork struct { +type UserNetwork struct { UserID snowflake.ID NetworkID snowflake.ID JoinedAt string @@ -68,3 +62,9 @@ type UsersNetwork struct { BanReason *string Position *int64 } + +type UserTrustedUser struct { + TrusterUserID snowflake.ID + TrustedUserID snowflake.ID + TrustedPublicKey []byte +} diff --git a/internal/data/users_networks.sql.go b/internal/data/users_networks.sql.go index 4c1d84f..42733da 100644 --- a/internal/data/users_networks.sql.go +++ b/internal/data/users_networks.sql.go @@ -107,6 +107,33 @@ func (q *Queries) GetNetworkMembers(ctx context.Context, networkID snowflake.ID) 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 @@ -183,7 +210,7 @@ type SetNetworkUserParams struct { BanReason *string } -func (q *Queries) SetNetworkUser(ctx context.Context, arg SetNetworkUserParams) (UsersNetwork, error) { +func (q *Queries) SetNetworkUser(ctx context.Context, arg SetNetworkUserParams) (UserNetwork, error) { row := q.db.QueryRowContext(ctx, setNetworkUser, arg.UserID, arg.NetworkID, @@ -193,7 +220,7 @@ func (q *Queries) SetNetworkUser(ctx context.Context, arg SetNetworkUserParams) arg.IsBanned, arg.BanReason, ) - var i UsersNetwork + var i UserNetwork err := row.Scan( &i.UserID, &i.NetworkID, diff --git a/internal/packet/models.go b/internal/packet/models.go index f1326e2..5e4a70b 100644 --- a/internal/packet/models.go +++ b/internal/packet/models.go @@ -3,6 +3,7 @@ package packet const ( MaxIconBytes = 16 DefaultFrequencyName = "main" + MaxFrequencyName = 32 ) const ( diff --git a/internal/packet/packet.go b/internal/packet/packet.go index df31dd0..8d06100 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -64,6 +64,7 @@ const ( PacketUpdateFrequency PacketDeleteFrequency PacketSwapFrequencies + PacketFrequenciesInfo PacketSendMessage PacketEditMessage @@ -185,38 +186,44 @@ func (p Packet) DecodedPayload() (Payload, error) { switch p.Type() { case PacketError: payload = &Error{} - case PacketCreateFrequency: - payload = &CreateFrequency{} + case PacketCreateNetwork: payload = &CreateNetwork{} - case PacketDeleteFrequency: - payload = &DeleteFrequency{} - case PacketDeleteMessage: - payload = &DeleteMessage{} + case PacketUpdateNetwork: + payload = &UpdateNetwork{} + case PacketTransferNetwork: + payload = &TransferNetwork{} case PacketDeleteNetwork: payload = &DeleteNetwork{} case PacketSwapUserNetworks: payload = &SwapUserNetworks{} - case PacketEditMessage: - payload = &EditMessage{} - case PacketMessagesInfo: - payload = &MessagesInfo{} - case PacketNetworksInfo: - payload = &NetworksInfo{} - case PacketRequestMessages: - payload = &RequestMessages{} - case PacketSendMessage: - payload = &SendMessage{} case PacketSetNetworkUser: payload = &SetNetworkUser{} - case PacketSwapFrequencies: - payload = &SwapFrequencies{} - case PacketTransferNetwork: - payload = &TransferNetwork{} + case PacketNetworksInfo: + payload = &NetworksInfo{} + + case PacketCreateFrequency: + payload = &CreateFrequency{} case PacketUpdateFrequency: payload = &UpdateFrequency{} - case PacketUpdateNetwork: - payload = &UpdateNetwork{} + case PacketDeleteFrequency: + payload = &DeleteFrequency{} + case PacketSwapFrequencies: + payload = &SwapFrequencies{} + case PacketFrequenciesInfo: + payload = &FrequenciesInfo{} + + case PacketSendMessage: + payload = &SendMessage{} + case PacketEditMessage: + payload = &EditMessage{} + case PacketDeleteMessage: + payload = &DeleteMessage{} + case PacketRequestMessages: + payload = &RequestMessages{} + case PacketMessagesInfo: + payload = &MessagesInfo{} + default: assert.Never("unexpected packet.PacketType", "type", p.Type()) } diff --git a/internal/packet/types.go b/internal/packet/types.go index 82b8397..6df3d65 100644 --- a/internal/packet/types.go +++ b/internal/packet/types.go @@ -132,6 +132,17 @@ func (m *SwapFrequencies) Type() PacketType { return PacketSwapFrequencies } +type FrequenciesInfo struct { + RemoveFrequencies []snowflake.ID + Frequencies []data.Frequency + Network snowflake.ID + Set bool +} + +func (m *FrequenciesInfo) Type() PacketType { + return PacketFrequenciesInfo +} + type SendMessage struct { ReceiverID *snowflake.ID FrequencyID *snowflake.ID diff --git a/internal/server/api/api.go b/internal/server/api/api.go index 9c6474e..4370d90 100644 --- a/internal/server/api/api.go +++ b/internal/server/api/api.go @@ -15,7 +15,10 @@ import ( "github.com/kyren223/eko/pkg/snowflake" ) -var internalError = packet.Error{Error: "internal server error"} +var ( + ErrInternalError = packet.Error{Error: "internal server error"} + ErrPermissionDenied = packet.Error{Error: "permission denied"} +) func SendMessage(ctx context.Context, sess *session.Session, request *packet.SendMessage) packet.Payload { if (request.ReceiverID != nil) == (request.FrequencyID != nil) { @@ -37,7 +40,7 @@ func SendMessage(ctx context.Context, sess *session.Session, request *packet.Sen }) if err != nil { log.Println(sess.Addr(), "database error:", err, "in SendMessage") - return &internalError + return &ErrInternalError } return &packet.MessagesInfo{Messages: []data.Message{message}} @@ -88,8 +91,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C if len(request.Icon) > packet.MaxIconBytes { return &packet.Error{Error: fmt.Sprintf( - "icon is too large, must be smaller than %v bytes", - packet.MaxIconBytes, + "exceeded allowed icon size in bytes: %v", packet.MaxIconBytes, )} } @@ -103,7 +105,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C tx, err := db.BeginTx(ctx, nil) if err != nil { log.Println("database error:", err) - return &internalError + return &ErrInternalError } defer tx.Rollback() //nolint @@ -121,7 +123,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C }) if err != nil { log.Println("database error 1:", err) - return &internalError + return &ErrInternalError } frequency, err := qtx.CreateFrequency(ctx, data.CreateFrequencyParams{ @@ -133,7 +135,7 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C }) if err != nil { log.Println("database error 2:", err) - return &internalError + return &ErrInternalError } networkUser, err := qtx.SetNetworkUser(ctx, data.SetNetworkUserParams{ @@ -147,19 +149,19 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C }) if err != nil { log.Println("database error 3:", err) - return &internalError + return &ErrInternalError } user, err := qtx.GetUserById(ctx, network.OwnerID) if err != nil { log.Println("database error 4:", err) - return &internalError + return &ErrInternalError } err = tx.Commit() if err != nil { log.Println("database error 5:", err) - return &internalError + return &ErrInternalError } fullNetwork := packet.FullNetwork{ @@ -174,12 +176,14 @@ func CreateNetwork(ctx context.Context, sess *session.Session, request *packet.C Position: int(*networkUser.Position), } return &packet.NetworksInfo{ - Networks: []packet.FullNetwork{fullNetwork}, + Networks: []packet.FullNetwork{fullNetwork}, + Set: false, + RemoveNetworks: nil, } } func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload, error) { - networksInfo := &packet.NetworksInfo{} + var fullNetworks []packet.FullNetwork tx, err := db.BeginTx(ctx, nil) if err != nil { @@ -208,7 +212,7 @@ func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload return nil, err } - networksInfo.Networks = append(networksInfo.Networks, packet.FullNetwork{ + fullNetworks = append(fullNetworks, packet.FullNetwork{ Network: network, Frequencies: frequencies, Members: members, @@ -221,10 +225,15 @@ func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload return nil, err } - return networksInfo, nil + return &packet.NetworksInfo{ + Networks: fullNetworks, + RemoveNetworks: nil, + Set: true, + }, nil } -// FIXME: Deletion of a network is not handled! it needs to be shifted like how it works with frequencies! +// FIXME: Deletion of a network is not handled! +// it needs to be shifted like how it works with frequencies! 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) @@ -235,8 +244,64 @@ func SwapUserNetworks(ctx context.Context, sess *session.Session, request *packe }) if err != nil { log.Println("database error:", err) - return &internalError + return &ErrInternalError } return request } + +func CreateFrequency(ctx context.Context, sess *session.Session, request *packet.CreateFrequency) packet.Payload { + queries := data.New(db) + + // Check if authorized + userNetwork, err := queries.GetUserNetwork(ctx, data.GetUserNetworkParams{ + UserID: sess.ID(), + NetworkID: request.Network, + }) + if err == sql.ErrNoRows { + return &packet.Error{Error: "user entry in network doesn't exist"} + } + if err != nil { + log.Println("database error 1:", err) + return &ErrInternalError + } + if !userNetwork.IsAdmin || !userNetwork.IsMember || userNetwork.IsBanned { + return &ErrPermissionDenied + } + + if len(request.Name) > packet.MaxFrequencyName { + return &packet.Error{Error: fmt.Sprintf( + "exceeded allowed frequency name length in bytes: %v", + packet.MaxFrequencyName, + )} + } + + if ok, err := isValidHexColor(request.HexColor); !ok { + return &packet.Error{Error: err} + } + + if request.Perms < 0 || request.Perms > packet.PermMax { + return &packet.Error{Error: fmt.Sprintf( + "exceeded allowed perms value: 0 <= perms < %v", packet.PermMax, + )} + } + + frequency, err := queries.CreateFrequency(ctx, data.CreateFrequencyParams{ + ID: sess.Manager().Node().Generate(), + NetworkID: request.Network, + Name: request.Name, + HexColor: &request.HexColor, + Perms: int64(request.Perms), + }) + if err != nil { + log.Println("database error 2:", err) + return &ErrInternalError + } + + return &packet.FrequenciesInfo{ + RemoveFrequencies: nil, + Frequencies: []data.Frequency{frequency}, + Network: request.Network, + Set: false, + } +} diff --git a/internal/server/server.go b/internal/server/server.go index 2ca0c2a..b6c4aad 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -318,6 +318,9 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P case *packet.SwapUserNetworks: response = timeout(5*time.Millisecond, api.SwapUserNetworks, ctx, sess, request) + case *packet.CreateFrequency: + response = timeout(5*time.Millisecond, api.CreateFrequency, ctx, sess, request) + case *packet.SendMessage: response = timeout(20*time.Millisecond, api.SendMessage, ctx, sess, request) case *packet.RequestMessages: diff --git a/query/users_networks.sql b/query/users_networks.sql index f467753..76cecf4 100644 --- a/query/users_networks.sql +++ b/query/users_networks.sql @@ -22,6 +22,10 @@ JOIN users_networks ON networks.id = users_networks.network_id WHERE users_networks.user_id = ? ORDER BY users_networks.position; +-- name: GetUserNetwork :one +SELECT * FROM users_networks +WHERE user_id = ? AND network_id = ?; + -- name: SetNetworkUser :one INSERT INTO users_networks ( user_id, network_id, @@ -20,4 +20,5 @@ sql: - column: "*.*_id" go_type: "github.com/kyren223/eko/pkg/snowflake.ID" rename: + users_network: "UserNetwork" is_public_dm: "IsPublicDM" |
