Files
NixMsg/internal/app/identity/lifecycle.go
T

665 lines
18 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package identity
import (
"bytes"
"context"
"database/sql"
"errors"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
const (
reasonEndpointDisabled = "endpoint_disabled"
reasonEndpointDeleted = "endpoint_deleted"
reasonSenderDisabled = "sender_disabled"
reasonSenderDeleted = "sender_deleted"
reasonLeftGroup = "left_group"
reasonGroupDissolved = "group_dissolved"
eventOwnerChanged = "owner_changed"
eventLeft = "left"
eventMemberRemoved = "member_removed"
eventDissolved = "dissolved"
)
type revokeItem struct {
endpointID string
msgID string
fromID string
reason string
}
type groupNotify struct {
recipients []string
groupID string
event string
endpointID string
atMs int64
}
// Disable 停用端:清令牌、作废相关消息/投递,并踢连接(DEVELOPMENT 7.6)。
func (a *App) Disable(ctx context.Context, endpointID string) error {
return a.disableOrDelete(ctx, endpointID, false)
}
// Enable 仅恢复 enabled=1;已作废消息不恢复。
func (a *App) Enable(ctx context.Context, endpointID string) error {
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
res, e := tx.ExecContext(ctx, `UPDATE endpoints SET enabled = 1 WHERE id = ?`, endpointID)
if e != nil {
return e
}
n, _ := res.RowsAffected()
if n == 0 {
return errCode(protocol.CodeNotFound, "endpoint not found")
}
return nil
})
return err
}
// Delete 删除端:停用效果(原因改为 deleted)+ 退群/转让群主 + 清授权与发出记录。
func (a *App) Delete(ctx context.Context, endpointID string) error {
return a.disableOrDelete(ctx, endpointID, true)
}
func (a *App) disableOrDelete(ctx context.Context, endpointID string, hardDelete bool) error {
nowMs := a.now().UnixMilli()
recvReason := reasonEndpointDisabled
sendReason := reasonSenderDisabled
if hardDelete {
recvReason = reasonEndpointDeleted
sendReason = reasonSenderDeleted
}
var revokes []revokeItem
var notifies []groupNotify
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
var one int
if e := tx.QueryRowContext(ctx, `SELECT 1 FROM endpoints WHERE id = ?`, endpointID).Scan(&one); e != nil {
if errors.Is(e, sql.ErrNoRows) {
return errCode(protocol.CodeNotFound, "endpoint not found")
}
return e
}
if _, e := tx.ExecContext(ctx, `
UPDATE endpoints SET enabled = 0,
session_hash = NULL, session_issued_at = NULL, session_used_at = NULL
WHERE id = ?`, endpointID); e != nil {
return e
}
if e := voidEndpointMessagesTx(tx, endpointID, recvReason, sendReason, nowMs, &revokes); e != nil {
return e
}
if !hardDelete {
return nil
}
if e := leaveAllGroupsTx(tx, endpointID, nowMs, &revokes, &notifies); e != nil {
return e
}
if _, e := tx.ExecContext(ctx, `DELETE FROM talk_grants WHERE sender_id = ? OR target_id = ?`, endpointID, endpointID); e != nil {
return e
}
if _, e := tx.ExecContext(ctx, `DELETE FROM receipts WHERE sender_id = ?`, endpointID); e != nil {
return e
}
if _, e := tx.ExecContext(ctx, `DELETE FROM send_keys WHERE sender_id = ?`, endpointID); e != nil {
return e
}
if _, e := tx.ExecContext(ctx, `DELETE FROM messages WHERE sender_id = ?`, endpointID); e != nil {
return e
}
_, e := tx.ExecContext(ctx, `DELETE FROM endpoints WHERE id = ?`, endpointID)
return e
})
if err != nil {
return err
}
a.publishRevokes(ctx, revokes)
a.publishGroupEvents(ctx, notifies)
if a.connCtrl != nil {
_ = a.connCtrl.Disconnect(ctx, endpointID, "", port.DisconnectFatal)
}
return nil
}
func voidEndpointMessagesTx(tx *sql.Tx, endpointID, recvReason, sendReason string, nowMs int64, revokes *[]revokeItem) error {
// 发给 X 的 pending → rejected
rows, err := tx.Query(`
SELECT d.seq, d.pushed_at, m.id, m.sender_id, m.receipt
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.endpoint_id = ? AND d.state = 'pending'`, endpointID)
if err != nil {
return err
}
type pendRow struct {
seq int64
pushed sql.NullInt64
msgID string
senderID string
receipt int
}
var pending []pendRow
for rows.Next() {
var r pendRow
if scanErr := rows.Scan(&r.seq, &r.pushed, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
_ = rows.Close()
return scanErr
}
pending = append(pending, r)
}
_ = rows.Close()
if err = rows.Err(); err != nil {
return err
}
finalSeqs := map[int64]struct{}{}
for _, r := range pending {
if _, execErr := tx.Exec(`
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
recvReason, nowMs, r.seq, endpointID); execErr != nil {
return execErr
}
if r.receipt != 0 {
if e := insertReceiptIfWantedTx(tx, r.senderID, r.msgID, endpointID, "rejected", recvReason, nowMs, true); e != nil {
return e
}
}
if r.pushed.Valid && revokes != nil {
*revokes = append(*revokes, revokeItem{
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: recvReason,
})
}
finalSeqs[r.seq] = struct{}{}
}
// 发给 X 的 scheduled 单聊 → completed,要回执则写
srows, err := tx.Query(`
SELECT seq, id, sender_id, receipt FROM messages
WHERE dest_kind = 'endpoint' AND dest_id = ? AND state = 'scheduled'`, endpointID)
if err != nil {
return err
}
type schedRow struct {
seq int64
msgID string
senderID string
receipt int
}
var scheduledTo []schedRow
for srows.Next() {
var r schedRow
if scanErr := srows.Scan(&r.seq, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
_ = srows.Close()
return scanErr
}
scheduledTo = append(scheduledTo, r)
}
_ = srows.Close()
if err = srows.Err(); err != nil {
return err
}
for _, r := range scheduledTo {
if _, execErr := tx.Exec(`
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
recvReason, r.seq); execErr != nil {
return execErr
}
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, r.seq); execErr != nil {
return execErr
}
if r.receipt != 0 {
// 消息级作废:endpoint_id 空,state=rejected(DEVELOPMENT 6.4)
if e := insertReceiptIfWantedTx(tx, r.senderID, r.msgID, "", "rejected", recvReason, nowMs, true); e != nil {
return e
}
}
}
// X 发出的 scheduled → completed(sender_*),不写回执
outSched, err := tx.Query(`SELECT seq FROM messages WHERE sender_id = ? AND state = 'scheduled'`, endpointID)
if err != nil {
return err
}
var outSeqs []int64
for outSched.Next() {
var seq int64
if scanErr := outSched.Scan(&seq); scanErr != nil {
_ = outSched.Close()
return scanErr
}
outSeqs = append(outSeqs, seq)
}
_ = outSched.Close()
if err = outSched.Err(); err != nil {
return err
}
for _, seq := range outSeqs {
if _, execErr := tx.Exec(`
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
sendReason, seq); execErr != nil {
return execErr
}
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq); execErr != nil {
return execErr
}
}
// X 发出的消息的 pending 投递 → rejected(sender_*),不写回执
drows, err := tx.Query(`
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE m.sender_id = ? AND d.state = 'pending'`, endpointID)
if err != nil {
return err
}
type outPend struct {
seq int64
endpointID string
pushed sql.NullInt64
msgID string
senderID string
}
var outPending []outPend
for drows.Next() {
var r outPend
if scanErr := drows.Scan(&r.seq, &r.endpointID, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
_ = drows.Close()
return scanErr
}
outPending = append(outPending, r)
}
_ = drows.Close()
if err = drows.Err(); err != nil {
return err
}
for _, r := range outPending {
if _, execErr := tx.Exec(`
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
sendReason, nowMs, r.seq, r.endpointID); execErr != nil {
return execErr
}
if r.pushed.Valid && revokes != nil {
*revokes = append(*revokes, revokeItem{
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: sendReason,
})
}
finalSeqs[r.seq] = struct{}{}
}
for seq := range finalSeqs {
if e := tryFinalizeTx(tx, seq); e != nil {
return e
}
}
return nil
}
func leaveAllGroupsTx(tx *sql.Tx, endpointID string, nowMs int64, revokes *[]revokeItem, notifies *[]groupNotify) error {
grows, err := tx.Query(`
SELECT g.id, g.owner_id
FROM groups g
JOIN group_members gm ON gm.group_id = g.id
WHERE gm.endpoint_id = ?`, endpointID)
if err != nil {
return err
}
type grow struct {
id string
owner string
}
var groups []grow
for grows.Next() {
var g grow
if scanErr := grows.Scan(&g.id, &g.owner); scanErr != nil {
_ = grows.Close()
return scanErr
}
groups = append(groups, g)
}
_ = grows.Close()
if err = grows.Err(); err != nil {
return err
}
for _, g := range groups {
members, memErr := loadMembersTx(tx, g.id)
if memErr != nil {
return memErr
}
if g.owner == endpointID {
others := withoutMember(members, endpointID)
if len(others) == 0 {
if e := voidGroupAllTx(tx, g.id, nowMs, revokes); e != nil {
return e
}
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ?`, g.id); e != nil {
return e
}
if _, e := tx.Exec(`DELETE FROM groups WHERE id = ?`, g.id); e != nil {
return e
}
if notifies != nil {
*notifies = append(*notifies, groupNotify{
recipients: members, groupID: g.id, event: eventDissolved, atMs: nowMs,
})
}
continue
}
newOwner, ownErr := earliestOtherMemberTx(tx, g.id, endpointID)
if ownErr != nil {
return ownErr
}
if _, e := tx.Exec(`UPDATE groups SET owner_id = ? WHERE id = ?`, newOwner, g.id); e != nil {
return e
}
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`, g.id, endpointID); e != nil {
return e
}
if e := voidMemberDeliveriesTx(tx, g.id, endpointID, reasonLeftGroup, nowMs, revokes); e != nil {
return e
}
left := withoutMember(members, endpointID)
if notifies != nil {
*notifies = append(*notifies,
groupNotify{recipients: left, groupID: g.id, event: eventOwnerChanged, endpointID: newOwner, atMs: nowMs},
groupNotify{recipients: append(append([]string{}, left...), endpointID), groupID: g.id, event: eventMemberRemoved, endpointID: endpointID, atMs: nowMs},
)
}
continue
}
if _, e := tx.Exec(`DELETE FROM group_members WHERE group_id = ? AND endpoint_id = ?`, g.id, endpointID); e != nil {
return e
}
if e := voidMemberDeliveriesTx(tx, g.id, endpointID, reasonLeftGroup, nowMs, revokes); e != nil {
return e
}
left := withoutMember(members, endpointID)
if notifies != nil {
*notifies = append(*notifies, groupNotify{
recipients: append(append([]string{}, left...), endpointID),
groupID: g.id, event: eventLeft, endpointID: endpointID, atMs: nowMs,
})
}
}
return nil
}
func loadMembersTx(tx *sql.Tx, groupID string) ([]string, error) {
rows, err := tx.Query(`SELECT endpoint_id FROM group_members WHERE group_id = ? ORDER BY joined_at ASC, endpoint_id ASC`, groupID)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
var out []string
for rows.Next() {
var id string
if e := rows.Scan(&id); e != nil {
return nil, e
}
out = append(out, id)
}
return out, rows.Err()
}
func earliestOtherMemberTx(tx *sql.Tx, groupID, exceptID string) (string, error) {
var id string
err := tx.QueryRow(`
SELECT endpoint_id FROM group_members
WHERE group_id = ? AND endpoint_id != ?
ORDER BY joined_at ASC, endpoint_id ASC
LIMIT 1`, groupID, exceptID).Scan(&id)
return id, err
}
func withoutMember(ids []string, drop string) []string {
out := make([]string, 0, len(ids))
for _, id := range ids {
if id != drop {
out = append(out, id)
}
}
return out
}
// voidMemberDeliveriesTx 与 group 包同语义:退群成员的 pending 群投递改 rejected。
func voidMemberDeliveriesTx(tx *sql.Tx, groupID, endpointID, reason string, nowMs int64, revokes *[]revokeItem) error {
rows, err := tx.Query(`
SELECT d.seq, d.pushed_at, m.id, m.sender_id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.endpoint_id = ? AND d.state = 'pending'
AND m.dest_kind = 'group' AND m.dest_id = ?`, endpointID, groupID)
if err != nil {
return err
}
type row struct {
seq int64
pushed sql.NullInt64
msgID string
senderID string
}
var list []row
for rows.Next() {
var r row
if scanErr := rows.Scan(&r.seq, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
_ = rows.Close()
return scanErr
}
list = append(list, r)
}
_ = rows.Close()
if err = rows.Err(); err != nil {
return err
}
for _, r := range list {
if _, execErr := tx.Exec(`
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
reason, nowMs, r.seq, endpointID); execErr != nil {
return execErr
}
if r.pushed.Valid && revokes != nil {
*revokes = append(*revokes, revokeItem{
endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason,
})
}
if e := tryFinalizeTx(tx, r.seq); e != nil {
return e
}
}
return nil
}
func voidGroupAllTx(tx *sql.Tx, groupID string, nowMs int64, revokes *[]revokeItem) error {
rows, err := tx.Query(`
SELECT d.seq, d.endpoint_id, d.pushed_at, m.id, m.sender_id
FROM deliveries d
JOIN messages m ON m.seq = d.seq
WHERE d.state = 'pending' AND m.dest_kind = 'group' AND m.dest_id = ?`, groupID)
if err != nil {
return err
}
type drow struct {
seq int64
endpointID string
pushed sql.NullInt64
msgID string
senderID string
}
var dlist []drow
for rows.Next() {
var r drow
if scanErr := rows.Scan(&r.seq, &r.endpointID, &r.pushed, &r.msgID, &r.senderID); scanErr != nil {
_ = rows.Close()
return scanErr
}
dlist = append(dlist, r)
}
_ = rows.Close()
if err = rows.Err(); err != nil {
return err
}
for _, r := range dlist {
if _, execErr := tx.Exec(`
UPDATE deliveries SET state = 'rejected', reason = ?, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
reasonGroupDissolved, nowMs, r.seq, r.endpointID); execErr != nil {
return execErr
}
if r.pushed.Valid && revokes != nil {
*revokes = append(*revokes, revokeItem{
endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved,
})
}
if e := tryFinalizeTx(tx, r.seq); e != nil {
return e
}
}
srows, err := tx.Query(`
SELECT seq, id, sender_id, receipt FROM messages
WHERE dest_kind = 'group' AND dest_id = ? AND state = 'scheduled'`, groupID)
if err != nil {
return err
}
type srow struct {
seq int64
msgID string
senderID string
receipt int
}
var slist []srow
for srows.Next() {
var r srow
if scanErr := srows.Scan(&r.seq, &r.msgID, &r.senderID, &r.receipt); scanErr != nil {
_ = srows.Close()
return scanErr
}
slist = append(slist, r)
}
_ = srows.Close()
if err = srows.Err(); err != nil {
return err
}
for _, r := range slist {
if _, execErr := tx.Exec(`
UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 'scheduled'`,
reasonGroupDissolved, r.seq); execErr != nil {
return execErr
}
if _, execErr := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, r.seq); execErr != nil {
return execErr
}
if r.receipt != 0 {
if e := insertReceiptIfWantedTx(tx, r.senderID, r.msgID, "", "rejected", reasonGroupDissolved, nowMs, true); e != nil {
return e
}
}
}
return nil
}
func insertReceiptIfWantedTx(tx *sql.Tx, senderID, msgID, endpointID, state, reason string, nowMs int64, alreadyWanted bool) error {
if !alreadyWanted {
return nil
}
var one int
err := tx.QueryRow(`SELECT 1 FROM endpoints WHERE id = ?`, senderID).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return nil
}
if err != nil {
return err
}
_, err = tx.Exec(`
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
VALUES(?,?,?,?,?,?,0)`, senderID, msgID, endpointID, state, reason, nowMs)
return err
}
func tryFinalizeTx(tx *sql.Tx, seq int64) error {
var n int
if err := tx.QueryRow(`SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state = 'pending'`, seq).Scan(&n); err != nil {
return err
}
if n > 0 {
return nil
}
var state string
if err := tx.QueryRow(`SELECT state FROM messages WHERE seq = ?`, seq).Scan(&state); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil
}
return err
}
if state == "completed" {
return nil
}
if _, err := tx.Exec(`UPDATE messages SET state = 'completed' WHERE seq = ?`, seq); err != nil {
return err
}
_, err := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq)
return err
}
func (a *App) publishRevokes(ctx context.Context, items []revokeItem) {
if a.down == nil || len(items) == 0 {
return
}
for _, it := range items {
frame := protocol.Revoked{
V: protocol.Version, Type: protocol.TypeRevoked,
ID: it.msgID, From: it.fromID, Reason: it.reason,
}
payload, encErr := encodeFrame(frame)
if encErr != nil {
continue
}
_ = a.down.PublishDown(ctx, it.endpointID, "", payload, port.PublishOpts{QoS: 1})
}
}
func (a *App) publishGroupEvents(ctx context.Context, items []groupNotify) {
if a.down == nil || len(items) == 0 {
return
}
for _, it := range items {
frame := protocol.GroupEvent{
V: protocol.Version, Type: protocol.TypeGroupEvent,
GroupID: it.groupID, Event: it.event, EndpointID: it.endpointID, AtMs: it.atMs,
}
payload, encErr := encodeFrame(frame)
if encErr != nil {
continue
}
seen := map[string]struct{}{}
for _, id := range it.recipients {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
_ = a.down.PublishDown(ctx, id, "", payload, port.PublishOpts{QoS: 0})
}
}
}
func encodeFrame(v any) ([]byte, error) {
var buf bytes.Buffer
if err := protocol.Encode(&buf, v); err != nil {
return nil, err
}
return buf.Bytes(), nil
}