summaryrefslogtreecommitdiff
path: root/internal/server/api
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2025-02-04 18:52:57 +0200
committerKyren223 <Kyren223@proton.me>2025-02-04 18:52:57 +0200
commiteab8cba4bb37491d5c3b497b98faa7029e7b7300 (patch)
tree15f9994c423d18f0492f97c43bf71a0793e1e6a4 /internal/server/api
parent42a1ee12c1e64955e6ddd38c34b122c7c10a93f2 (diff)
Refactored the way notifications work on the server-side, next commit
fixes the client to use this new system
Diffstat (limited to 'internal/server/api')
-rw-r--r--internal/server/api/api.go89
-rw-r--r--internal/server/api/helpers.go51
-rw-r--r--internal/server/api/migrations/20250204150045_last_read_messages.sql39
3 files changed, 130 insertions, 49 deletions
diff --git a/internal/server/api/api.go b/internal/server/api/api.go
index 60e873e..a95065c 100644
--- a/internal/server/api/api.go
+++ b/internal/server/api/api.go
@@ -22,6 +22,7 @@ var (
ErrInternalError = packet.Error{Error: "internal server error"}
ErrPermissionDenied = packet.Error{Error: "permission denied"}
ErrNotImplemented = packet.Error{Error: "not implemented yet"}
+ ErrSuccess = packet.Error{Error: "success"}
DefaultBanReason = ""
)
@@ -1247,23 +1248,91 @@ func GetBannedMembers(ctx context.Context, sess *session.Session, request *packe
}
}
-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"}
+func GetNotifications(ctx context.Context, sess *session.Session) packet.Payload {
+ info, err := getNotifications(ctx, sess.ID())
+ if err != nil {
+ log.Println("database error:", err)
+ return &ErrInternalError
}
- if len(request.Source) == 0 {
- return &packet.NotificationsInfo{
- Source: []snowflake.ID{},
- Pings: []*int64{},
- }
+ return &info
+}
+
+func SetLastReadMessages(ctx context.Context, sess *session.Session, request *packet.SetLastReadMessages) packet.Payload {
+ if len(request.Source) != len(request.LastRead) {
+ return &packet.Error{Error: fmt.Sprintf(
+ "%v sources doesn't match %v last reads",
+ len(request.Source), len(request.LastRead),
+ )}
}
- info, err := getNotifications(ctx, request, sess.ID())
+ tx, err := db.BeginTx(ctx, nil)
if err != nil {
log.Println("database error:", err)
return &ErrInternalError
}
+ defer func() { _ = tx.Rollback() }()
- return &info
+ queries := data.New(db)
+ qtx := queries.WithTx(tx)
+
+ for i := 0; i < len(request.Source); i++ {
+ _, err := qtx.GetUserById(ctx, request.Source[i])
+ if err == nil {
+ err = qtx.SetLastReadMessage(ctx, data.SetLastReadMessageParams{
+ UserID: sess.ID(),
+ SourceID: request.Source[i],
+ LastRead: request.LastRead[i],
+ })
+ if err != nil {
+ log.Println("database error 1:", err)
+ return &ErrInternalError
+ }
+ continue
+ }
+ if err != nil && err != sql.ErrNoRows {
+ log.Println("database error 2:", err)
+ return &ErrInternalError
+ }
+
+ frequency, err := qtx.GetFrequencyById(ctx, request.Source[i])
+ if err == sql.ErrNoRows {
+ return &packet.Error{Error: fmt.Sprintf(
+ "source at %v is not a valid frequency or user id", i,
+ )}
+ }
+ if err != nil {
+ log.Println("database error 3:", err)
+ return &ErrInternalError
+ }
+ if frequency.Perms == packet.PermNoAccess {
+ isAdmin, err := IsNetworkAdmin(ctx, qtx, sess.ID(), frequency.NetworkID)
+ if err != nil {
+ log.Println("database error 4:", err)
+ return &ErrInternalError
+ }
+ if !isAdmin {
+ return &ErrPermissionDenied
+ }
+ }
+
+ err = qtx.SetLastReadMessage(ctx, data.SetLastReadMessageParams{
+ UserID: sess.ID(),
+ SourceID: request.Source[i],
+ LastRead: request.LastRead[i],
+ })
+ if err != nil {
+ log.Println("database error 5:", err)
+ return &ErrInternalError
+ }
+ continue
+ }
+
+ err = tx.Commit()
+ if err != nil {
+ log.Println("database error 6:", err)
+ return &ErrInternalError
+ }
+
+ return &ErrSuccess
}
diff --git a/internal/server/api/helpers.go b/internal/server/api/helpers.go
index 321bfd2..c4deb6c 100644
--- a/internal/server/api/helpers.go
+++ b/internal/server/api/helpers.go
@@ -9,7 +9,6 @@ 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"
)
@@ -133,55 +132,29 @@ func UserPropagate(
}
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)
+WITH entries AS (
+ SELECT source_id, last_read
+ FROM last_read_messages
+ WHERE user_id = ?
)
SELECT
- e.source,
+ e.source_id,
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;
+LEFT JOIN messages m ON m.id > e.last_read
+ AND (m.frequency_id = e.source_id OR
+ (m.receiver_id = e.source_id AND m.sender_id = ?) OR
+ (m.sender_id = e.source_id AND m.receiver_id = ?))
+GROUP BY e.source_id, e.last_read;
`
-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,
- )
-
+func getNotifications(ctx context.Context, userId snowflake.ID) (packet.NotificationsInfo, error) {
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...)
+ rows, err := db.QueryContext(ctx, query, userId, userId, userId, userId, userId)
if err != nil {
return packet.NotificationsInfo{}, err
}
diff --git a/internal/server/api/migrations/20250204150045_last_read_messages.sql b/internal/server/api/migrations/20250204150045_last_read_messages.sql
new file mode 100644
index 0000000..f44436b
--- /dev/null
+++ b/internal/server/api/migrations/20250204150045_last_read_messages.sql
@@ -0,0 +1,39 @@
+-- +goose Up
+CREATE TABLE IF NOT EXISTS last_read_messages (
+ user_id INT NOT NULL REFERENCES users (id),
+ source_id INT NOT NULL, -- frequency_id or receiver_id or sender_id
+ last_read INT NOT NULL,
+ -- can be 0 for "no messages ever read" (null)
+ -- can be between msgs in cases like the message getting deleted
+ PRIMARY KEY (user_id, source_id)
+);
+
+DROP TRIGGER IF EXISTS on_frequency_delete;
+-- +goose StatementBegin
+CREATE TRIGGER IF NOT EXISTS on_frequency_delete
+AFTER DELETE ON frequencies
+BEGIN
+ UPDATE frequencies SET
+ position = position - 1
+ WHERE network_id = OLD.network_id AND position > OLD.position;
+
+ DELETE FROM last_read_messages WHERE source_id = OLD.id;
+END
+-- +goose StatementEnd
+
+DROP TRIGGER IF EXISTS on_user_delete;
+-- +goose StatementBegin
+CREATE TRIGGER IF NOT EXISTS on_user_delete
+AFTER UPDATE OF is_deleted ON users
+WHEN NEW.is_deleted = true
+BEGIN
+ DELETE FROM networks WHERE owner_id = NEW.id;
+ DELETE FROM members WHERE user_id = NEW.id;
+ DELETE FROM trusted_users WHERE trusting_user_id = NEW.id OR trusted_user_id = NEW.id;
+ DELETE FROM blocked_users WHERE blocking_user_id = NEW.id OR blocked_user_id = NEW.id;
+ DELETE FROM last_read_messages WHERE user_id = NEW.id;
+END
+-- +goose StatementEnd
+
+-- +goose Down
+DROP TABLE IF EXISTS last_read_messages;