summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/data/db.go2
-rw-r--r--internal/data/frequencies.sql.go2
-rw-r--r--internal/data/members.sql.go2
-rw-r--r--internal/data/messages.sql.go2
-rw-r--r--internal/data/models.go8
-rw-r--r--internal/data/networks.sql.go2
-rw-r--r--internal/data/notifications.sql.go32
-rw-r--r--internal/data/trusted_and_blocked_users.sql.go2
-rw-r--r--internal/data/users.sql.go2
-rw-r--r--internal/packet/packet.go6
-rw-r--r--internal/packet/types.go10
-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
-rw-r--r--internal/server/server.go56
15 files changed, 219 insertions, 86 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
+}