diff options
Diffstat (limited to 'internal/server')
| -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 |
3 files changed, 98 insertions, 0 deletions
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"} } |
