145 lines
4.0 KiB
Go
145 lines
4.0 KiB
Go
package repo
|
||
|
||
import (
|
||
"context"
|
||
"testing"
|
||
|
||
"github.com/agentmail/gateway/internal/db"
|
||
"github.com/google/uuid"
|
||
)
|
||
|
||
// seedMailTo 造一封给 recipient 的未读邮件,可选带抄送。
|
||
func seedMailTo(t *testing.T, recipient string, cc string) uuid.UUID {
|
||
t.Helper()
|
||
ctx := context.Background()
|
||
sid, err := CreateSession(ctx, nil, "sender", "t", "")
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
ccJSON := "[]"
|
||
if cc != "" {
|
||
ccJSON = `[{"name":"` + cc + `","path":"","session":"","raw":"` + cc + `"}]`
|
||
}
|
||
var id uuid.UUID
|
||
err = db.DB.QueryRowContext(ctx,
|
||
`INSERT INTO mails (session_id, from_name, to_name, subject, body, cc_list)
|
||
VALUES ($1, 'sender', $2, 's', 'b', $3) RETURNING mail_id`,
|
||
sid, recipient, ccJSON).Scan(&id)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
return id
|
||
}
|
||
|
||
func statusOf(t *testing.T, id uuid.UUID) string {
|
||
t.Helper()
|
||
var s string
|
||
if err := db.DB.QueryRowContext(context.Background(),
|
||
`SELECT status FROM mails WHERE mail_id = $1`, id).Scan(&s); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
return s
|
||
}
|
||
|
||
func TestMarkMailsReadForOnlyOwnMail(t *testing.T) {
|
||
setupTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
mine := seedMailTo(t, "bot", "")
|
||
others := seedMailTo(t, "other", "")
|
||
|
||
// 一次请求里混着别人的邮件:自己的标掉,别人的动不了。
|
||
// 鉴权写在 UPDATE 的 WHERE 里,所以这不是「先查后拒」而是根本改不动。
|
||
n, err := MarkMailsReadFor(ctx, "bot", []uuid.UUID{mine, others})
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if n != 1 {
|
||
t.Fatalf("影响行数 = %d,期望 1(只有自己那封)", n)
|
||
}
|
||
if statusOf(t, mine) != "read" {
|
||
t.Fatal("自己的邮件没被标记")
|
||
}
|
||
if statusOf(t, others) != "unread" {
|
||
t.Fatal("别人的邮件被标记了 —— 鉴权失效")
|
||
}
|
||
}
|
||
|
||
// 重复标记是幂等的:Agent 通常把上一轮列出的 id 原样传回来,
|
||
// 其中混着已读的不该算错误
|
||
func TestMarkMailsReadForIsIdempotent(t *testing.T) {
|
||
setupTestDB(t)
|
||
ctx := context.Background()
|
||
id := seedMailTo(t, "bot", "")
|
||
|
||
if n, _ := MarkMailsReadFor(ctx, "bot", []uuid.UUID{id}); n != 1 {
|
||
t.Fatalf("首次应标掉 1 封,实际 %d", n)
|
||
}
|
||
n, err := MarkMailsReadFor(ctx, "bot", []uuid.UUID{id})
|
||
if err != nil {
|
||
t.Fatalf("重复标记不该报错: %v", err)
|
||
}
|
||
if n != 0 {
|
||
t.Fatalf("重复标记影响行数 = %d,期望 0", n)
|
||
}
|
||
}
|
||
|
||
// 被抄送的邮件也在收件箱里,也该能标掉
|
||
func TestMarkMailsReadForCoversCC(t *testing.T) {
|
||
setupTestDB(t)
|
||
ctx := context.Background()
|
||
id := seedMailTo(t, "other", "bot") // 主收件人是 other,bot 被抄送
|
||
|
||
n, err := MarkMailsReadFor(ctx, "bot", []uuid.UUID{id})
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if n != 1 {
|
||
t.Fatalf("被抄送的邮件应可标记,影响行数 = %d", n)
|
||
}
|
||
}
|
||
|
||
func TestMarkMailsReadForEmptyList(t *testing.T) {
|
||
setupTestDB(t)
|
||
// 空列表直接返回,不该拼出 `IN ()` 这种非法 SQL
|
||
n, err := MarkMailsReadFor(context.Background(), "bot", nil)
|
||
if err != nil {
|
||
t.Fatalf("空列表不该报错: %v", err)
|
||
}
|
||
if n != 0 {
|
||
t.Fatalf("空列表影响行数 = %d", n)
|
||
}
|
||
}
|
||
|
||
func TestMarkAllInboxReadForSkipsArchivedAndOthers(t *testing.T) {
|
||
setupTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
a := seedMailTo(t, "bot", "")
|
||
b := seedMailTo(t, "bot", "")
|
||
others := seedMailTo(t, "other", "")
|
||
|
||
// 把 b 所在会话归档:那封在收件箱里根本看不到,
|
||
// 标掉它只会让「标记了 N 封」与用户看到的对不上
|
||
var sid uuid.UUID
|
||
db.DB.QueryRowContext(ctx, `SELECT session_id FROM mails WHERE mail_id = $1`, b).Scan(&sid)
|
||
db.DB.ExecContext(ctx, `UPDATE sessions SET status = 'archived' WHERE session_id = $1`, sid)
|
||
|
||
n, err := MarkAllInboxReadFor(ctx, "bot")
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if n != 1 {
|
||
t.Fatalf("影响行数 = %d,期望 1(排除归档会话)", n)
|
||
}
|
||
if statusOf(t, a) != "read" {
|
||
t.Fatal("活跃会话里的未读没被标掉")
|
||
}
|
||
if statusOf(t, b) != "unread" {
|
||
t.Fatal("归档会话里的邮件被标掉了")
|
||
}
|
||
if statusOf(t, others) != "unread" {
|
||
t.Fatal("别人的邮件被标掉了")
|
||
}
|
||
}
|