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("别人的邮件被标掉了") } }