diff options
| author | Kyren223 <Kyren223@proton.me> | 2025-02-04 18:52:57 +0200 |
|---|---|---|
| committer | Kyren223 <Kyren223@proton.me> | 2025-02-04 18:52:57 +0200 |
| commit | eab8cba4bb37491d5c3b497b98faa7029e7b7300 (patch) | |
| tree | 15f9994c423d18f0492f97c43bf71a0793e1e6a4 /internal/server/api | |
| parent | 42a1ee12c1e64955e6ddd38c34b122c7c10a93f2 (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.go | 89 | ||||
| -rw-r--r-- | internal/server/api/helpers.go | 51 | ||||
| -rw-r--r-- | internal/server/api/migrations/20250204150045_last_read_messages.sql | 39 |
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; |
