From 547593d3159d8772392e81b40082928094b35b8c Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Mon, 6 Jan 2025 16:17:59 +0200 Subject: Fleshed out and fixed issues in both the server and client for how networks and members work --- internal/client/ui/core/chat/chat.go | 6 +- internal/client/ui/core/core.go | 48 +++++++++++++- internal/client/ui/core/networklist/networklist.go | 34 +++++++++- internal/data/users_networks.sql.go | 45 +++++++++++++ internal/server/api/api.go | 74 +++++++++++++++++++--- internal/server/api/helpers.go | 46 ++++++++++++++ internal/server/server.go | 12 +++- internal/server/session/session.go | 3 +- query/users_networks.sql | 4 ++ 9 files changed, 255 insertions(+), 17 deletions(-) diff --git a/internal/client/ui/core/chat/chat.go b/internal/client/ui/core/chat/chat.go index c94db31..11cb751 100644 --- a/internal/client/ui/core/chat/chat.go +++ b/internal/client/ui/core/chat/chat.go @@ -255,8 +255,10 @@ func (m *Model) SetFrequency(networkIndex, frequencyIndex int) tea.Cmd { } func (m *Model) ResetBeforeSwitch() { - if m.frequencyIndex != -1 && m.networkIndex != -1 { - network := state.State.Networks[m.networkIndex] + networks := state.State.Networks + networkInBounds := 0 <= m.networkIndex && m.networkIndex < len(networks) + if m.frequencyIndex != -1 && networkInBounds { + network := networks[m.networkIndex] frequencyId := network.Frequencies[m.frequencyIndex].ID log.Println("Saving", frequencyId) state.State.FrequencyState[frequencyId] = state.Frequency{ diff --git a/internal/client/ui/core/core.go b/internal/client/ui/core/core.go index 69bffec..50abbe6 100644 --- a/internal/client/ui/core/core.go +++ b/internal/client/ui/core/core.go @@ -212,6 +212,35 @@ func (m *Model) updateConnected(msg tea.Msg) tea.Cmd { state.State.Networks = networks } + case *packet.MembersInfo: + var network *packet.FullNetwork + for i, fullNetwork := range state.State.Networks { + if fullNetwork.ID == msg.Network { + network = &state.State.Networks[i] + } + } + + members := network.Members + for _, updatedMember := range msg.Members { + add := true + for i, existingMember := range members { + if existingMember.User.ID == updatedMember.User.ID { + add = false + members[i] = updatedMember + break + } + } + if add { + members = append(members, updatedMember) + } + } + + members = slices.DeleteFunc(members, func(member data.GetNetworkMembersRow) bool { + return slices.Contains(msg.RemovedMembers, member.User.ID) + }) + log.Println(members) + network.Members = members + case *packet.FrequenciesInfo: var network *packet.FullNetwork for i, fullNetwork := range state.State.Networks { @@ -219,8 +248,25 @@ func (m *Model) updateConnected(msg tea.Msg) tea.Cmd { network = &state.State.Networks[i] } } + frequencies := network.Frequencies - frequencies = append(frequencies, msg.Frequencies...) + for _, newFrequency := range msg.Frequencies { + add := true + for i, existingFrequency := range frequencies { + if existingFrequency.ID == newFrequency.ID { + add = false + if newFrequency.Position == -1 { + newFrequency.Position = existingFrequency.Position + } + frequencies[i] = newFrequency + break + } + } + if add { + frequencies = append(frequencies, newFrequency) + } + } + frequencies = slices.DeleteFunc(frequencies, func(frequency data.Frequency) bool { return slices.Contains(msg.RemovedFrequencies, frequency.ID) }) diff --git a/internal/client/ui/core/networklist/networklist.go b/internal/client/ui/core/networklist/networklist.go index c510e79..3ff73f4 100644 --- a/internal/client/ui/core/networklist/networklist.go +++ b/internal/client/ui/core/networklist/networklist.go @@ -1,6 +1,7 @@ package networklist import ( + "slices" "strings" tea "github.com/charmbracelet/bubbletea" @@ -51,6 +52,8 @@ func IconStyle(icon string, fg, bg lipgloss.Color) lipgloss.Style { return lipgloss.NewStyle().SetString(combined) } +const TrustedIndex = -1 + type Model struct { history []func(m *Model) index int @@ -60,7 +63,7 @@ type Model struct { func New() Model { return Model{ focus: false, - index: -1, + index: TrustedIndex, } } @@ -71,7 +74,7 @@ func (m Model) Init() tea.Cmd { func (m Model) View() string { var builder strings.Builder builder.WriteString("\n") - if m.index == -1 { + if m.index == TrustedIndex { builder.WriteString(trustedUsersButtonSelected) } else { builder.WriteString(trustedUsersButton) @@ -138,6 +141,33 @@ func (m Model) Update(msg tea.Msg) (Model, tea.Cmd) { m.index = max(-1, m.index-1) case "j": m.index = min(len(state.State.Networks)-1, m.index+1) + + case "Q": + if state.State.UserID == nil || m.index == TrustedIndex { + return m, nil + } + + networks := state.State.Networks + network := networks[m.index] + + // Remove network from the list so it's not visible + state.State.Networks = slices.Delete(networks, m.index, m.index+1) + + // Set back to trusted bcz that's always valid + // Where if u were on the network u just left + // it'd be an issue (or if u left all networks) + m.index = TrustedIndex + + no := false + return m, gateway.Send(&packet.SetMember{ + Member: &no, + Admin: nil, + Muted: nil, + Banned: nil, + BanReason: nil, + Network: network.ID, + User: *state.State.UserID, + }) } } return m, nil diff --git a/internal/data/users_networks.sql.go b/internal/data/users_networks.sql.go index 8191fa2..e180b51 100644 --- a/internal/data/users_networks.sql.go +++ b/internal/data/users_networks.sql.go @@ -7,10 +7,55 @@ 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, diff --git a/internal/server/api/api.go b/internal/server/api/api.go index aedc86b..dee49aa 100644 --- a/internal/server/api/api.go +++ b/internal/server/api/api.go @@ -201,6 +201,12 @@ func GetNetworksInfo(ctx context.Context, sess *session.Session) (packet.Payload 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) frequencies, err := qtx.GetNetworkFrequencies(ctx, network.ID) if err != nil { @@ -446,7 +452,7 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe return &ErrInternalError } - return &packet.MembersInfo{ + NetworkPropagate(ctx, sess, request.Network, &packet.MembersInfo{ RemovedMembers: nil, Members: []data.GetNetworkMembersRow{{ User: user, @@ -455,17 +461,40 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe IsMuted: newMember.IsMuted, }}, Network: request.Network, + }) + + frequencies, err := queries.GetNetworkFrequencies(ctx, network.ID) + if err != nil { + log.Println("database error 4:", err) + return &ErrInternalError + } + + members, err := queries.GetNetworkMembers(ctx, network.ID) + if err != nil { + log.Println("database error 5:", err) + return &ErrInternalError + } + + return &packet.NetworksInfo{ + Networks: []packet.FullNetwork{{ + Network: network, + Frequencies: frequencies, + Members: members, + Position: int(*newMember.Position), + }}, + RemovedNetworks: nil, + Set: false, } } if err != nil { - log.Println("database error 4:", err) + log.Println("database error 6:", err) return &ErrInternalError } isSessAdmin, err := IsNetworkAdmin(ctx, queries, sess.ID(), request.Network) if err != nil { - log.Println("database error 5:", err) + log.Println("database error 7:", err) return &ErrInternalError } @@ -516,25 +545,25 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe BanReason: banReason, }) if err != nil { - log.Println("database error 6:", err) + log.Println("database error 8:", err) return &ErrInternalError } if !newMember.IsMember { - return &packet.MembersInfo{ + return NetworkPropagate(ctx, sess, request.Network, &packet.MembersInfo{ RemovedMembers: []snowflake.ID{newMember.UserID}, Members: nil, Network: request.Network, - } + }) } user, err := queries.GetUserById(ctx, newMember.UserID) if err != nil { - log.Println("database error 7:", err) + log.Println("database error 9:", err) return &ErrInternalError } - return &packet.MembersInfo{ + payload := NetworkPropagate(ctx, sess, request.Network, &packet.MembersInfo{ RemovedMembers: nil, Members: []data.GetNetworkMembersRow{{ User: user, @@ -543,5 +572,34 @@ func SetMember(ctx context.Context, sess *session.Session, request *packet.SetMe IsMuted: newMember.IsMuted, }}, Network: request.Network, + }) + + // Joined + if !member.IsMember && newMember.IsMember { + frequencies, err := queries.GetNetworkFrequencies(ctx, network.ID) + if err != nil { + log.Println("database error 10:", err) + return &ErrInternalError + } + + members, err := queries.GetNetworkMembers(ctx, network.ID) + if err != nil { + log.Println("database error 11:", err) + return &ErrInternalError + } + + return &packet.NetworksInfo{ + Networks: []packet.FullNetwork{{ + Network: network, + Frequencies: frequencies, + Members: members, + Position: int(*newMember.Position), + }}, + RemovedNetworks: nil, + Set: false, + } } + + // Normal case, was already in the server + return payload } diff --git a/internal/server/api/helpers.go b/internal/server/api/helpers.go index f82bf6e..05dae0e 100644 --- a/internal/server/api/helpers.go +++ b/internal/server/api/helpers.go @@ -2,9 +2,13 @@ package api import ( "context" + "log" "strings" + "time" "github.com/kyren223/eko/internal/data" + "github.com/kyren223/eko/internal/packet" + "github.com/kyren223/eko/internal/server/session" "github.com/kyren223/eko/pkg/snowflake" ) @@ -41,3 +45,45 @@ func IsNetworkAdmin(ctx context.Context, queries *data.Queries, userId, networkI return isAdmin, nil } +func NetworkPropagate( + ctx context.Context, sess *session.Session, + network snowflake.ID, payload packet.Payload, +) packet.Payload { + var sessions []snowflake.ID + sess.Manager().UseSessions(func(s map[snowflake.ID]*session.Session) { + sessions = make([]snowflake.ID, 0, len(s)-1) + for key := range s { + if key != sess.ID() { + sessions = append(sessions, key) + } + } + }) + + queries := data.New(db) + sessions, err := queries.FilterUsersInNetwork(ctx, data.FilterUsersInNetworkParams{ + NetworkID: network, + Users: sessions, + }) + if err != nil { + log.Println("database error in propagate:", err) + return &ErrInternalError + } + + for _, sessionId := range sessions { + session := sess.Manager().Session(sessionId) + if session == nil { + continue + } + timeout := 1 * time.Second + context, cancel := context.WithTimeout(context.Background(), timeout) + go func() { + defer cancel() + pkt := packet.NewPacket(packet.NewJsonEncoder(payload)) + if ok := session.Write(context, pkt); !ok { + log.Println(sess.Addr(), "propagation to", session.Addr(), "failed") + } + }() + } + + return payload +} diff --git a/internal/server/server.go b/internal/server/server.go index 2b1003f..0534023 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -80,11 +80,17 @@ func (s *server) RemoveSession(id snowflake.ID) { delete(s.sessions, id) } -func (s *server) Session(id snowflake.ID) (*session.Session, bool) { +func (s *server) Session(id snowflake.ID) *session.Session { s.sessMu.RLock() defer s.sessMu.RUnlock() - session, ok := s.sessions[id] - return session, ok + session := s.sessions[id] + return session +} + +func (s *server) UseSessions(f func(map[snowflake.ID]*session.Session)) { + s.sessMu.RLock() + defer s.sessMu.RUnlock() + f(s.sessions) } func (s *server) Node() *snowflake.Node { diff --git a/internal/server/session/session.go b/internal/server/session/session.go index dea9af2..2b7121a 100644 --- a/internal/server/session/session.go +++ b/internal/server/session/session.go @@ -16,7 +16,8 @@ import ( type SessionManager interface { AddSession(session *Session) RemoveSession(id snowflake.ID) - Session(id snowflake.ID) (session *Session, ok bool) + Session(id snowflake.ID) *Session + UseSessions(f func(map[snowflake.ID]*Session)) Node() *snowflake.Node } diff --git a/query/users_networks.sql b/query/users_networks.sql index 79cfd41..bde3c3b 100644 --- a/query/users_networks.sql +++ b/query/users_networks.sql @@ -59,3 +59,7 @@ UPDATE users_networks SET WHEN position = @pos2 THEN @pos1 END WHERE user_id = @user_id AND position IN (@pos1, @pos2); + +-- name: FilterUsersInNetwork :many +SELECT user_id FROM users_networks +WHERE network_id = ? AND user_id IN (sqlc.slice('users')); -- cgit v1.3.1