diff options
| -rw-r--r-- | internal/client/gateway/gateway.go | 2 | ||||
| -rw-r--r-- | internal/client/ui/core/core.go | 15 | ||||
| -rw-r--r-- | internal/client/ui/core/state/state.go | 49 | ||||
| -rw-r--r-- | internal/packet/packet.go | 8 | ||||
| -rw-r--r-- | internal/packet/types.go | 18 | ||||
| -rw-r--r-- | internal/server/api/api.go | 21 | ||||
| -rw-r--r-- | internal/server/api/helpers.go | 74 | ||||
| -rw-r--r-- | internal/server/server.go | 3 | ||||
| -rw-r--r-- | query/notifications.sql | 14 |
9 files changed, 203 insertions, 1 deletions
diff --git a/internal/client/gateway/gateway.go b/internal/client/gateway/gateway.go index c8824fd..178dd92 100644 --- a/internal/client/gateway/gateway.go +++ b/internal/client/gateway/gateway.go @@ -189,7 +189,7 @@ func handlePacketStream() { payload, err := pkt.DecodedPayload() assert.NoError(err, "server should always provide a decodeable packet") - log.Println("received streamed packet:", payload) + log.Printf("received streamed packet %v: %v\n", payload.Type(), payload) ui.Program.Send(payload) } } diff --git a/internal/client/ui/core/core.go b/internal/client/ui/core/core.go index 36c2b0f..c7b3ca4 100644 --- a/internal/client/ui/core/core.go +++ b/internal/client/ui/core/core.go @@ -189,6 +189,9 @@ func (m *Model) updateNotConnected(msg tea.Msg) tea.Cmd { switch msg := msg.(type) { case gateway.ConnectionEstablished: state.UserID = (*snowflake.ID)(&msg) + state.ReceivedData = false + state.ReceivedNetworks = false + state.AskedForNotifs = false m.connected = true m.timeout = initialTimeout @@ -257,6 +260,11 @@ func (m *Model) updateConnected(msg tea.Msg) tea.Cmd { m.memberList.SetWidth(sidebarWidth) m.chat.SetWidth(chatWidth) + if state.ReceivedData && state.ReceivedNetworks && !state.AskedForNotifs { + state.AskedForNotifs = true + state.AskForNotifs() + } + switch msg := msg.(type) { case ui.QuitMsg: data := state.JsonUserData() @@ -335,6 +343,9 @@ func (m *Model) updateConnected(msg tea.Msg) tea.Cmd { case *packet.TrustInfo: state.UpdateTrusteds(msg) + case *packet.NotificationsInfo: + state.UpdateNotifications(msg) + case ui.BanReasonPopupMsg: popup := banreason.New(msg.User, msg.Network) m.banReasonPopup = &popup @@ -693,6 +704,10 @@ func (m *Model) HasPopup() bool { func calculateNotifications() { for networkId := range state.State.Networks { for _, frequency := range state.State.Frequencies[networkId] { + if _, ok := state.State.Messages[frequency.ID]; !ok { + continue + } + pings, hasNotif := getFrequencyNotification(networkId, frequency.ID) if hasNotif { state.State.Notifications[frequency.ID] = pings diff --git a/internal/client/ui/core/state/state.go b/internal/client/ui/core/state/state.go index ee61a2c..ba9e233 100644 --- a/internal/client/ui/core/state/state.go +++ b/internal/client/ui/core/state/state.go @@ -3,6 +3,7 @@ package state import ( "crypto/ed25519" "encoding/json" + "log" "slices" "github.com/google/btree" @@ -59,6 +60,12 @@ var Data UserData = UserData{ var UserID *snowflake.ID = nil +var ( + ReceivedData = false + ReceivedNetworks = false + AskedForNotifs = false +) + func UpdateNetworks(info *packet.NetworksInfo) { networks := State.Networks @@ -111,6 +118,8 @@ func UpdateNetworks(info *packet.NetworksInfo) { Data: &data, User: nil, }) + + ReceivedNetworks = true } func UpdateFrequencies(info *packet.FrequenciesInfo) { @@ -215,6 +224,8 @@ func FromJsonUserData(s string) { if Data.LastReadMessage == nil { Data.LastReadMessage = map[snowflake.ID]*snowflake.ID{} } + + ReceivedData = true } func UpdateTrusteds(info *packet.TrustInfo) { @@ -246,3 +257,41 @@ func SendUserDatUpdate() { User: nil, }) } + +func UpdateNotifications(info *packet.NotificationsInfo) { + for i := 0; i < len(info.Source); i++ { + ping := info.Pings[i] + if ping != nil { + State.Notifications[info.Source[i]] = int(*ping) + } else { + delete(State.Notifications, info.Source[i]) + } + } +} + +func AskForNotifs() { + source := []snowflake.ID{} + + source = append(source, Data.Peers...) + for networkId := range State.Networks { + for _, frequency := range State.Frequencies[networkId] { + source = append(source, frequency.ID) + } + } + + lastReadId := make([]snowflake.ID, 0, len(source)) + + for _, source := range source { + id := Data.LastReadMessage[source] + if id == nil { + lastReadId = append(lastReadId, 0) + } else { + lastReadId = append(lastReadId, *id) + } + } + + gateway.SendAsync(&packet.GetNotifications{ + Source: source, + LastReadId: lastReadId, + }) +} diff --git a/internal/packet/packet.go b/internal/packet/packet.go index d107fb7..72ff4b9 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -80,6 +80,9 @@ const ( PacketTrustUser PacketTrustInfo + PacketGetNotifications + PacketNotificationsInfo + PacketMax ) @@ -245,6 +248,11 @@ func (p Packet) DecodedPayload() (Payload, error) { case PacketTrustInfo: payload = &TrustInfo{} + case PacketGetNotifications: + payload = &GetNotifications{} + case PacketNotificationsInfo: + payload = &NotificationsInfo{} + default: assert.Never("unexpected packet.PacketType", "type", p.Type()) } diff --git a/internal/packet/types.go b/internal/packet/types.go index 6d0b1dd..27fd028 100644 --- a/internal/packet/types.go +++ b/internal/packet/types.go @@ -233,3 +233,21 @@ type GetBannedMembers struct { func (m *GetBannedMembers) Type() PacketType { return PacketGetBannedMembers } + +type GetNotifications struct { + Source []snowflake.ID + LastReadId []snowflake.ID +} + +func (m *GetNotifications) Type() PacketType { + return PacketGetNotifications +} + +type NotificationsInfo struct { + Source []snowflake.ID + Pings []*int64 +} + +func (m *NotificationsInfo) Type() PacketType { + return PacketNotificationsInfo +} diff --git a/internal/server/api/api.go b/internal/server/api/api.go index 3fff656..60e873e 100644 --- a/internal/server/api/api.go +++ b/internal/server/api/api.go @@ -1246,3 +1246,24 @@ func GetBannedMembers(ctx context.Context, sess *session.Session, request *packe Network: request.Network, } } + +func GetNotifications(ctx context.Context, sess *session.Session, request *packet.GetNotifications) packet.Payload { + if len(request.Source) != len(request.LastReadId) { + return &packet.Error{Error: "source length and lastReadId length must match"} + } + + if len(request.Source) == 0 { + return &packet.NotificationsInfo{ + Source: []snowflake.ID{}, + Pings: []*int64{}, + } + } + + info, err := getNotifications(ctx, request, sess.ID()) + if err != nil { + log.Println("database error:", err) + return &ErrInternalError + } + + return &info +} diff --git a/internal/server/api/helpers.go b/internal/server/api/helpers.go index e25d63c..321bfd2 100644 --- a/internal/server/api/helpers.go +++ b/internal/server/api/helpers.go @@ -9,6 +9,7 @@ import ( "github.com/kyren223/eko/internal/data" "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/internal/server/session" + "github.com/kyren223/eko/pkg/assert" "github.com/kyren223/eko/pkg/snowflake" ) @@ -130,3 +131,76 @@ func UserPropagate( return payload } + +const getNotificationsQuery = `-- name: GetNotifications :many +WITH +entries(source, lastId) AS ( + VALUES /*SLICE:pair*/? +), +permitted_frequencies AS ( + SELECT f.id, m.is_admin + FROM frequencies f + JOIN entries e ON f.id = e.source + LEFT JOIN members m + ON m.user_id = ? + AND m.network_id = f.network_id + WHERE m.is_member = true AND (f.perms != 0 OR m.is_admin = true) +) +SELECT + e.source, + CASE + WHEN COUNT(m.id) = 0 THEN NULL + ELSE SUM(CASE WHEN (m.ping = 0 OR (m.ping = 1 AND pf.is_admin = true) OR m.ping = ?) THEN 1 ELSE 0 END) + -- 0 is @everyone, 1 is @admins, otherwise it's user_id + END AS pings +FROM entries e +LEFT JOIN messages m ON m.id > e.lastId + AND (m.frequency_id = e.source OR + (m.receiver_id = e.source AND m.sender_id = ?) OR + (m.sender_id = e.source AND m.receiver_id = ?)) +JOIN permitted_frequencies pf ON e.source = pf.id +GROUP BY e.source, e.lastId; +` + +func getNotifications(ctx context.Context, arg *packet.GetNotifications, ping snowflake.ID) (packet.NotificationsInfo, error) { + assert.Assert( + len(arg.Source) == len(arg.LastReadId), + "Source and LastReadID must match", + "source", arg.Source, "last read id", arg.LastReadId, + ) + + query := getNotificationsQuery + var queryParams []interface{} + if len(arg.Source) > 0 { + for i := 0; i < len(arg.Source); i++ { + queryParams = append(queryParams, arg.Source[i]) + queryParams = append(queryParams, arg.LastReadId[i]) + } + query = strings.Replace(query, "/*SLICE:pair*/?", strings.Repeat(",(?,?)", len(arg.Source))[1:], 1) + } else { + assert.Never("source and id must not be empty") + } + queryParams = append(queryParams, ping, ping, ping, ping) + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return packet.NotificationsInfo{}, err + } + defer rows.Close() + var items packet.NotificationsInfo + for rows.Next() { + var source *snowflake.ID + var pings *int64 + if err := rows.Scan(&source, &pings); err != nil { + return packet.NotificationsInfo{}, err + } + items.Source = append(items.Source, *source) + items.Pings = append(items.Pings, pings) + } + if err := rows.Close(); err != nil { + return packet.NotificationsInfo{}, err + } + if err := rows.Err(); err != nil { + return packet.NotificationsInfo{}, err + } + return items, nil +} diff --git a/internal/server/server.go b/internal/server/server.go index cf7dfcf..5055539 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -386,6 +386,9 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P case *packet.TrustUser: response = timeout(10*time.Millisecond, api.TrustUser, ctx, sess, request) + case *packet.GetNotifications: + response = timeout(100*time.Millisecond, api.GetNotifications, ctx, sess, request) + default: response = &packet.Error{Error: "use of disallowed packet type for request"} } diff --git a/query/notifications.sql b/query/notifications.sql new file mode 100644 index 0000000..5522fe8 --- /dev/null +++ b/query/notifications.sql @@ -0,0 +1,14 @@ +-- name: GetNotifications :many +WITH entries(source, lastId) AS ( + VALUES + ('source1', 12345), + ('source2', 54321) +) +SELECT + e.source, + CASE WHEN COUNT(m.id) > 0 THEN 1 ELSE 0 END AS hasNotif, + COALESCE(SUM(CASE WHEN m.ping = ? THEN 1 ELSE 0 END), 0) AS pings +FROM entries e +LEFT JOIN messages m ON m.id > e.lastId + AND (m.s1 = e.source OR m.s2 = e.source OR m.s3 = e.source) +GROUP BY e.source, e.lastId; |
