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 | |
| parent | 42a1ee12c1e64955e6ddd38c34b122c7c10a93f2 (diff) | |
Refactored the way notifications work on the server-side, next commit
fixes the client to use this new system
| -rw-r--r-- | internal/data/db.go | 2 | ||||
| -rw-r--r-- | internal/data/frequencies.sql.go | 2 | ||||
| -rw-r--r-- | internal/data/members.sql.go | 2 | ||||
| -rw-r--r-- | internal/data/messages.sql.go | 2 | ||||
| -rw-r--r-- | internal/data/models.go | 8 | ||||
| -rw-r--r-- | internal/data/networks.sql.go | 2 | ||||
| -rw-r--r-- | internal/data/notifications.sql.go | 32 | ||||
| -rw-r--r-- | internal/data/trusted_and_blocked_users.sql.go | 2 | ||||
| -rw-r--r-- | internal/data/users.sql.go | 2 | ||||
| -rw-r--r-- | internal/packet/packet.go | 6 | ||||
| -rw-r--r-- | internal/packet/types.go | 10 | ||||
| -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 | ||||
| -rw-r--r-- | internal/server/server.go | 56 | ||||
| -rw-r--r-- | query/notifications.sql | 21 |
16 files changed, 226 insertions, 100 deletions
diff --git a/internal/data/db.go b/internal/data/db.go index 02c70ab..15d00a9 100644 --- a/internal/data/db.go +++ b/internal/data/db.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 package data diff --git a/internal/data/frequencies.sql.go b/internal/data/frequencies.sql.go index 171943f..4971238 100644 --- a/internal/data/frequencies.sql.go +++ b/internal/data/frequencies.sql.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 // source: frequencies.sql package data diff --git a/internal/data/members.sql.go b/internal/data/members.sql.go index 98a6a1e..38dfd3b 100644 --- a/internal/data/members.sql.go +++ b/internal/data/members.sql.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 // source: members.sql package data diff --git a/internal/data/messages.sql.go b/internal/data/messages.sql.go index 82909e3..ed94b95 100644 --- a/internal/data/messages.sql.go +++ b/internal/data/messages.sql.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 // source: messages.sql package data diff --git a/internal/data/models.go b/internal/data/models.go index 8470e2a..02949d0 100644 --- a/internal/data/models.go +++ b/internal/data/models.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 package data @@ -23,6 +23,12 @@ type Frequency struct { Position int64 } +type LastReadMessage struct { + UserID snowflake.ID + SourceID snowflake.ID + LastRead int64 +} + type Member struct { UserID snowflake.ID NetworkID snowflake.ID diff --git a/internal/data/networks.sql.go b/internal/data/networks.sql.go index af3c273..0439caf 100644 --- a/internal/data/networks.sql.go +++ b/internal/data/networks.sql.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 // source: networks.sql package data diff --git a/internal/data/notifications.sql.go b/internal/data/notifications.sql.go new file mode 100644 index 0000000..5819438 --- /dev/null +++ b/internal/data/notifications.sql.go @@ -0,0 +1,32 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.28.0 +// source: notifications.sql + +package data + +import ( + "context" + + "github.com/kyren223/eko/pkg/snowflake" +) + +const setLastReadMessage = `-- name: SetLastReadMessage :exec +INSERT INTO last_read_messages ( + user_id, source_id, last_read +) VALUES (?, ?, ?) +ON CONFLICT DO +UPDATE SET last_read = EXCLUDED.last_read +WHERE user_id = EXCLUDED.user_id AND source_id = EXCLUDED.source_id +` + +type SetLastReadMessageParams struct { + UserID snowflake.ID + SourceID snowflake.ID + LastRead int64 +} + +func (q *Queries) SetLastReadMessage(ctx context.Context, arg SetLastReadMessageParams) error { + _, err := q.db.ExecContext(ctx, setLastReadMessage, arg.UserID, arg.SourceID, arg.LastRead) + return err +} diff --git a/internal/data/trusted_and_blocked_users.sql.go b/internal/data/trusted_and_blocked_users.sql.go index 436b8d9..daeed9e 100644 --- a/internal/data/trusted_and_blocked_users.sql.go +++ b/internal/data/trusted_and_blocked_users.sql.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 // source: trusted_and_blocked_users.sql package data diff --git a/internal/data/users.sql.go b/internal/data/users.sql.go index 2e20cec..b68a207 100644 --- a/internal/data/users.sql.go +++ b/internal/data/users.sql.go @@ -1,6 +1,6 @@ // Code generated by sqlc. DO NOT EDIT. // versions: -// sqlc v1.27.0 +// sqlc v1.28.0 // source: users.sql package data diff --git a/internal/packet/packet.go b/internal/packet/packet.go index 72ff4b9..c8a8e64 100644 --- a/internal/packet/packet.go +++ b/internal/packet/packet.go @@ -80,7 +80,7 @@ const ( PacketTrustUser PacketTrustInfo - PacketGetNotifications + PacketSetLastReadMessages PacketNotificationsInfo PacketMax @@ -248,8 +248,8 @@ func (p Packet) DecodedPayload() (Payload, error) { case PacketTrustInfo: payload = &TrustInfo{} - case PacketGetNotifications: - payload = &GetNotifications{} + case PacketSetLastReadMessages: + payload = &SetLastReadMessages{} case PacketNotificationsInfo: payload = &NotificationsInfo{} diff --git a/internal/packet/types.go b/internal/packet/types.go index 27fd028..bb2b013 100644 --- a/internal/packet/types.go +++ b/internal/packet/types.go @@ -234,13 +234,13 @@ func (m *GetBannedMembers) Type() PacketType { return PacketGetBannedMembers } -type GetNotifications struct { - Source []snowflake.ID - LastReadId []snowflake.ID +type SetLastReadMessages struct { + Source []snowflake.ID + LastRead []int64 } -func (m *GetNotifications) Type() PacketType { - return PacketGetNotifications +func (m *SetLastReadMessages) Type() PacketType { + return PacketSetLastReadMessages } type NotificationsInfo struct { 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; diff --git a/internal/server/server.go b/internal/server/server.go index 5055539..e83fcb4 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -240,27 +240,9 @@ func (server *server) handleConnection(conn net.Conn) { } }() - // Send initial packets - payload := api.GetUserData(ctx, sess, &packet.GetUserData{}) - if payload == &api.ErrInternalError { + if ok := server.sendInitialPackets(ctx, sess); !ok { return // closes the connection } - pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload)) - sess.Write(ctx, pkt) - - payload = api.GetUserTrusteds(ctx, sess) - if payload == &api.ErrInternalError { - return // closes the connection - } - pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload)) - sess.Write(ctx, pkt) - - payload, err = api.GetNetworksInfo(ctx, sess) - if err != nil { - return // closes the connection - } - pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload)) - sess.Write(ctx, pkt) // Infinite read loop buffer := make([]byte, 512) @@ -386,8 +368,8 @@ 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) + case *packet.SetLastReadMessages: + response = timeout(100*time.Millisecond, api.SetLastReadMessages, ctx, sess, request) default: response = &packet.Error{Error: "use of disallowed packet type for request"} @@ -422,3 +404,35 @@ func timeout[T packet.Payload]( return &packet.Error{Error: "request timeout"} } } + +func (server *server) sendInitialPackets(ctx context.Context, sess *session.Session) bool { + payload := api.GetUserData(ctx, sess, &packet.GetUserData{}) + if payload == &api.ErrInternalError { + return false + } + pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload)) + sess.Write(ctx, pkt) + + payload = api.GetUserTrusteds(ctx, sess) + if payload == &api.ErrInternalError { + return false + } + pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload)) + sess.Write(ctx, pkt) + + payload, err := api.GetNetworksInfo(ctx, sess) + if err != nil { + return false + } + pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload)) + sess.Write(ctx, pkt) + + payload = api.GetNotifications(ctx, sess) + if payload == &api.ErrInternalError { + return false + } + pkt = packet.NewPacket(packet.NewMsgPackEncoder(payload)) + sess.Write(ctx, pkt) + + return true +} diff --git a/query/notifications.sql b/query/notifications.sql index 5522fe8..5dd4d5a 100644 --- a/query/notifications.sql +++ b/query/notifications.sql @@ -1,14 +1,7 @@ --- 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; +-- name: SetLastReadMessage :exec +INSERT INTO last_read_messages ( + user_id, source_id, last_read +) VALUES (?, ?, ?) +ON CONFLICT DO +UPDATE SET last_read = EXCLUDED.last_read +WHERE user_id = EXCLUDED.user_id AND source_id = EXCLUDED.source_id; |
