summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--internal/client/gateway/gateway.go2
-rw-r--r--internal/client/ui/core/core.go15
-rw-r--r--internal/client/ui/core/state/state.go49
-rw-r--r--internal/packet/packet.go8
-rw-r--r--internal/packet/types.go18
-rw-r--r--internal/server/api/api.go21
-rw-r--r--internal/server/api/helpers.go74
-rw-r--r--internal/server/server.go3
-rw-r--r--query/notifications.sql14
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;