Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions internal/handlers/attachments.go
Original file line number Diff line number Diff line change
Expand Up @@ -201,9 +201,9 @@ func (h *AttachmentHandler) Download(c *gin.Context) {
var recipientCount int
if err = h.DB.Pool.QueryRow(ctx,
`SELECT COUNT(*) FROM (
SELECT 1 FROM msg_to WHERE msg_id = $1 AND addr = $2
SELECT 1 FROM msg_to WHERE msg_id = $1 AND lower(addr) = lower($2)
UNION ALL
SELECT 1 FROM msg_add_to WHERE msg_id = $1 AND addr = $2
SELECT 1 FROM msg_add_to WHERE msg_id = $1 AND lower(addr) = lower($2)
) r`, msgID, identity,
).Scan(&recipientCount); err != nil || recipientCount == 0 {
c.JSON(http.StatusForbidden, gin.H{"error": "access denied"})
Expand Down
60 changes: 35 additions & 25 deletions internal/handlers/messages.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,9 @@ func NewMessageHandler(database *db.DB, dataDir string, maxDataSize, maxMsgSize
// or the sub-account named via X-FMSG-Act-As / a sub-account's own API-key
// token). Callers never see another identity's messages in the same request.
func (h *MessageHandler) visibleAddrs(c *gin.Context) ([]string, error) {
return []string{middleware.GetIdentity(c)}, nil
// Lower-cased: fmsg addresses are case-insensitive and the msg tables carry
// lower(addr) expression indexes, so queries compare lower(col) = ANY(addrs).
return []string{strings.ToLower(middleware.GetIdentity(c))}, nil
}

// messageListItem is the JSON shape for each message in the list response.
Expand Down Expand Up @@ -293,12 +295,12 @@ func (h *MessageHandler) List(c *gin.Context) {
rows, err := h.DB.Pool.Query(ctx,
`SELECT m.id, m.version, m.pid, m.no_reply, m.is_important, m.is_deflate, m.is_terminal, m.time_sent, m.from_addr, m.topic, m.type, m.size, m.filepath,
COALESCE(
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = m.id AND mt.addr = ANY($1)),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = m.id AND mat.addr = ANY($1))
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = m.id AND lower(mt.addr) = ANY($1)),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = m.id AND lower(mat.addr) = ANY($1))
) AS time_read
FROM msg m
WHERE EXISTS (SELECT 1 FROM msg_to mt WHERE mt.msg_id = m.id AND mt.addr = ANY($1))
OR EXISTS (SELECT 1 FROM msg_add_to mat WHERE mat.msg_id = m.id AND mat.addr = ANY($1))
WHERE EXISTS (SELECT 1 FROM msg_to mt WHERE mt.msg_id = m.id AND lower(mt.addr) = ANY($1))
OR EXISTS (SELECT 1 FROM msg_add_to mat WHERE mat.msg_id = m.id AND lower(mat.addr) = ANY($1))
ORDER BY m.id DESC
LIMIT $2 OFFSET $3`,
addrs, limit, offset,
Expand Down Expand Up @@ -397,7 +399,7 @@ func (h *MessageHandler) Sent(c *gin.Context) {
rows, err := h.DB.Pool.Query(ctx,
`SELECT m.id, m.version, m.pid, m.no_reply, m.is_important, m.is_deflate, m.is_terminal, m.time_sent, m.from_addr, m.topic, m.type, m.size, m.filepath
FROM msg m
WHERE m.from_addr = ANY($1)
WHERE lower(m.from_addr) = ANY($1)
ORDER BY m.id DESC
LIMIT $2 OFFSET $3`,
addrs, limit, offset,
Expand Down Expand Up @@ -481,7 +483,7 @@ func (h *MessageHandler) Create(c *gin.Context) {
}

// Enforce ownership: from must match the JWT identity.
if msg.From != identity {
if !sameAddr(msg.From, identity) {
c.JSON(http.StatusForbidden, gin.H{"error": "from address must match authenticated user"})
return
}
Expand Down Expand Up @@ -608,8 +610,8 @@ func (h *MessageHandler) Get(c *gin.Context) {
var timeRead *float64
err := h.DB.Pool.QueryRow(ctx,
`SELECT COALESCE(
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = $1 AND mt.addr = ANY($2)),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = $1 AND mat.addr = ANY($2))
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = $1 AND lower(mt.addr) = ANY($2)),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = $1 AND lower(mat.addr) = ANY($2))
)`,
msgID, addrs,
).Scan(&timeRead)
Expand Down Expand Up @@ -658,9 +660,9 @@ func (h *MessageHandler) DownloadData(c *gin.Context) {
var recipientCount int
if err = h.DB.Pool.QueryRow(ctx,
`SELECT COUNT(*) FROM (
SELECT 1 FROM msg_to WHERE msg_id = $1 AND addr = ANY($2)
SELECT 1 FROM msg_to WHERE msg_id = $1 AND lower(addr) = ANY($2)
UNION ALL
SELECT 1 FROM msg_add_to WHERE msg_id = $1 AND addr = ANY($2)
SELECT 1 FROM msg_add_to WHERE msg_id = $1 AND lower(addr) = ANY($2)
) r`, msgID, addrs,
).Scan(&recipientCount); err != nil || recipientCount == 0 {
c.JSON(http.StatusForbidden, gin.H{"error": "access denied"})
Expand Down Expand Up @@ -704,7 +706,7 @@ func (h *MessageHandler) Update(c *gin.Context) {
return
}

if existing.From != identity {
if !sameAddr(existing.From, identity) {
c.JSON(http.StatusForbidden, gin.H{"error": "only the owner may update a message"})
return
}
Expand All @@ -718,7 +720,7 @@ func (h *MessageHandler) Update(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if msg.From != identity {
if !sameAddr(msg.From, identity) {
c.JSON(http.StatusForbidden, gin.H{"error": "from address must match authenticated user"})
return
}
Expand Down Expand Up @@ -813,7 +815,7 @@ func (h *MessageHandler) Delete(c *gin.Context) {
return
}

if existing.From != identity {
if !sameAddr(existing.From, identity) {
c.JSON(http.StatusForbidden, gin.H{"error": "only the owner may delete a message"})
return
}
Expand Down Expand Up @@ -885,7 +887,7 @@ func (h *MessageHandler) Send(c *gin.Context) {
return
}

if existing.From != identity {
if !sameAddr(existing.From, identity) {
c.JSON(http.StatusForbidden, gin.H{"error": "only the owner may send a message"})
return
}
Expand Down Expand Up @@ -961,11 +963,11 @@ func (h *MessageHandler) MarkRead(c *gin.Context) {
err := h.DB.Pool.QueryRow(ctx,
`SELECT
COALESCE(
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = $1 AND mt.addr = $2),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = $1 AND mat.addr = $2)
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = $1 AND lower(mt.addr) = lower($2)),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = $1 AND lower(mat.addr) = lower($2))
),
EXISTS (SELECT 1 FROM msg_to mt WHERE mt.msg_id = $1 AND mt.addr = $2)
OR EXISTS (SELECT 1 FROM msg_add_to mat WHERE mat.msg_id = $1 AND mat.addr = $2)`,
EXISTS (SELECT 1 FROM msg_to mt WHERE mt.msg_id = $1 AND lower(mt.addr) = lower($2))
OR EXISTS (SELECT 1 FROM msg_add_to mat WHERE mat.msg_id = $1 AND lower(mat.addr) = lower($2))`,
msgID, identity,
).Scan(&existing, &recipient)
if err != nil {
Expand All @@ -989,7 +991,7 @@ func (h *MessageHandler) MarkRead(c *gin.Context) {
// each table.
if _, err = h.DB.Pool.Exec(ctx,
`UPDATE msg_to SET time_read = $1
WHERE msg_id = $2 AND addr = $3 AND time_read IS NULL`,
WHERE msg_id = $2 AND lower(addr) = lower($3) AND time_read IS NULL`,
now, msgID, identity,
); err != nil {
log.Printf("mark read %d: update msg_to: %v", msgID, err)
Expand All @@ -998,7 +1000,7 @@ func (h *MessageHandler) MarkRead(c *gin.Context) {
}
if _, err = h.DB.Pool.Exec(ctx,
`UPDATE msg_add_to SET time_read = $1
WHERE msg_id = $2 AND addr = $3 AND time_read IS NULL`,
WHERE msg_id = $2 AND lower(addr) = lower($3) AND time_read IS NULL`,
now, msgID, identity,
); err != nil {
log.Printf("mark read %d: update msg_add_to: %v", msgID, err)
Expand Down Expand Up @@ -1068,10 +1070,10 @@ func (h *MessageHandler) AddRecipients(c *gin.Context) {
}

// Verify the requester is an existing participant (from or msg_to).
if fromAddr != identity {
if !sameAddr(fromAddr, identity) {
var recipientCount int
if err = h.DB.Pool.QueryRow(ctx,
"SELECT COUNT(*) FROM msg_to WHERE msg_id = $1 AND addr = $2", msgID, identity,
"SELECT COUNT(*) FROM msg_to WHERE msg_id = $1 AND lower(addr) = lower($2)", msgID, identity,
).Scan(&recipientCount); err != nil || recipientCount == 0 {
c.JSON(http.StatusForbidden, gin.H{"error": "only existing participants may add recipients"})
return
Expand Down Expand Up @@ -1375,8 +1377,8 @@ func (h *MessageHandler) messageItemFor(ctx context.Context, msgID int64, recipi
var timeRead *float64
if err := h.DB.Pool.QueryRow(ctx,
`SELECT COALESCE(
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = $1 AND mt.addr = $2),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = $1 AND mat.addr = $2)
(SELECT mt.time_read FROM msg_to mt WHERE mt.msg_id = $1 AND lower(mt.addr) = lower($2)),
(SELECT mat.time_read FROM msg_add_to mat WHERE mat.msg_id = $1 AND lower(mat.addr) = lower($2))
)`,
msgID, recipient,
).Scan(&timeRead); err == nil {
Expand Down Expand Up @@ -1484,6 +1486,14 @@ func parseLimitOffset(c *gin.Context) (int, int, bool) {
}

// isRecipient checks whether addr appears in the to list (case-insensitive).
// sameAddr reports whether two fmsg addresses are equal. Addresses are
// case-insensitive (SPECIFICATION.md: case folding), and the identity in a
// first-party token keeps the case the sub-account was created with, so
// ownership checks must never compare bytes.
func sameAddr(a, b string) bool {
return strings.EqualFold(a, b)
}

func isRecipient(to []string, addr string) bool {
for _, a := range to {
if strings.EqualFold(a, addr) {
Expand Down
18 changes: 18 additions & 0 deletions internal/handlers/messages_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,24 @@ func TestParseAddr(t *testing.T) {
}
}

func TestSameAddr(t *testing.T) {
cases := []struct {
a, b string
want bool
}{
{"@alice@example.com", "@alice@example.com", true},
{"@alice_ChatGPT@example.com", "@alice_chatgpt@example.com", true},
{"@Alice@Example.COM", "@alice@example.com", true},
{"@alice@example.com", "@alice@example.org", false},
{"@alice@example.com", "@alicia@example.com", false},
}
for _, tc := range cases {
if got := sameAddr(tc.a, tc.b); got != tc.want {
t.Errorf("sameAddr(%q, %q) = %v, want %v", tc.a, tc.b, got, tc.want)
}
}
}

func TestIsRecipient(t *testing.T) {
list := []string{"@alice@example.com", "@bob@example.com"}
if !isRecipient(list, "@ALICE@example.com") {
Expand Down
12 changes: 6 additions & 6 deletions internal/handlers/thread.go
Original file line number Diff line number Diff line change
Expand Up @@ -285,9 +285,9 @@ func (h *MessageHandler) ThreadText(c *gin.Context) {
WHERE c.depth < $3
)
SELECT c.id, c.from_addr, c.time_sent, c.type, c.size, c.filepath,
(c.from_addr = ANY($2)
OR EXISTS (SELECT 1 FROM msg_to t WHERE t.msg_id = c.id AND t.addr = ANY($2))
OR EXISTS (SELECT 1 FROM msg_add_to a WHERE a.msg_id = c.id AND a.addr = ANY($2))) AS readable
(lower(c.from_addr) = ANY($2)
OR EXISTS (SELECT 1 FROM msg_to t WHERE t.msg_id = c.id AND lower(t.addr) = ANY($2))
OR EXISTS (SELECT 1 FROM msg_add_to a WHERE a.msg_id = c.id AND lower(a.addr) = ANY($2))) AS readable
FROM chain c ORDER BY c.depth DESC`,
msgID, addrs, threadMaxHops,
)
Expand Down Expand Up @@ -391,9 +391,9 @@ func (h *MessageHandler) ThreadMessages(c *gin.Context) {
SELECT c.id, c.version, c.pid, c.no_reply, c.is_important, c.is_deflate,
c.is_terminal, c.time_sent, c.from_addr, c.topic, c.type, c.size, c.filepath,
encode(c.sha256, 'hex'),
(c.from_addr = ANY($2)
OR EXISTS (SELECT 1 FROM msg_to t WHERE t.msg_id = c.id AND t.addr = ANY($2))
OR EXISTS (SELECT 1 FROM msg_add_to a WHERE a.msg_id = c.id AND a.addr = ANY($2)))
(lower(c.from_addr) = ANY($2)
OR EXISTS (SELECT 1 FROM msg_to t WHERE t.msg_id = c.id AND lower(t.addr) = ANY($2))
OR EXISTS (SELECT 1 FROM msg_add_to a WHERE a.msg_id = c.id AND lower(a.addr) = ANY($2)))
FROM chain c ORDER BY c.depth DESC`, msgID, addrs, threadMaxHops)
if err != nil {
log.Printf("thread messages: walk %d: %v", msgID, err)
Expand Down
Loading