From eab8cba4bb37491d5c3b497b98faa7029e7b7300 Mon Sep 17 00:00:00 2001 From: Kyren223 Date: Tue, 4 Feb 2025 18:52:57 +0200 Subject: Refactored the way notifications work on the server-side, next commit fixes the client to use this new system --- internal/server/api/api.go | 89 ++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 79 insertions(+), 10 deletions(-) (limited to 'internal/server/api/api.go') 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 } -- cgit v1.3.1