package identity import ( "bytes" "context" "database/sql" "errors" "time" "git.asio.asia/nixevol/NixMsg/internal/app/message" "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" // kickFlushDelay 给接线方 Session.Disable/Deleted 留出发 fatal 的窗口。 kickFlushDelay = 20 * time.Millisecond ) 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, ¬ifies); 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) // fatal+断开由 admin DisableKick/DeleteKick(Session.Disable/Deleted)完成。 // 未接 Kick 钩子的单元测试仍可用 ConnControl 兜底断开。 if a.connCtrl != nil { go func() { time.Sleep(kickFlushDelay) _ = a.connCtrl.Disconnect(context.Background(), endpointID, "", port.DisconnectFatal) }() } return nil } func voidEndpointMessagesTx(tx *sql.Tx, endpointID, recvReason, sendReason string, nowMs int64, revokes *[]revokeItem) error { days := message.DefaultVoidRetentionDays // 发给 X 的 pending → rejected rows, err := tx.Query(` SELECT d.seq, m.id, m.sender_id 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 msgID string senderID string } var pending []pendRow for rows.Next() { var r pendRow if scanErr := rows.Scan(&r.seq, &r.msgID, &r.senderID); 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 { pushed, execErr := message.RejectPendingTx(tx, r.seq, endpointID, recvReason, nowMs) if execErr != nil { return execErr } if pushed && 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, 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 senderID string receipt int } var scheduledTo []schedRow for srows.Next() { var r schedRow if scanErr := srows.Scan(&r.seq, &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 e := message.FinalizeMessageTx(tx, r.seq, r.receipt != 0, r.senderID, "", recvReason, nowMs, days); 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 e := message.FinalizeMessageTx(tx, seq, false, endpointID, "", sendReason, nowMs, days); e != nil { return e } } // X 发出的消息的 pending 投递 → rejected(sender_*),不写回执 drows, err := tx.Query(` SELECT d.seq, d.endpoint_id, 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 msgID string senderID string } var outPending []outPend for drows.Next() { var r outPend if scanErr := drows.Scan(&r.seq, &r.endpointID, &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 { pushed, execErr := message.RejectPendingTx(tx, r.seq, r.endpointID, sendReason, nowMs) if execErr != nil { return execErr } if pushed && 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 := message.TryFinalizeTx(tx, seq, nowMs, days); 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, 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 msgID string senderID string } var list []row for rows.Next() { var r row if scanErr := rows.Scan(&r.seq, &r.msgID, &r.senderID); scanErr != nil { _ = rows.Close() return scanErr } list = append(list, r) } _ = rows.Close() if err = rows.Err(); err != nil { return err } days := message.DefaultVoidRetentionDays for _, r := range list { pushed, execErr := message.RejectPendingTx(tx, r.seq, endpointID, reason, nowMs) if execErr != nil { return execErr } if pushed && revokes != nil { *revokes = append(*revokes, revokeItem{ endpointID: endpointID, msgID: r.msgID, fromID: r.senderID, reason: reason, }) } if e := message.TryFinalizeTx(tx, r.seq, nowMs, days); e != nil { return e } } return nil } func voidGroupAllTx(tx *sql.Tx, groupID string, nowMs int64, revokes *[]revokeItem) error { days := message.DefaultVoidRetentionDays rows, err := tx.Query(` SELECT d.seq, d.endpoint_id, 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 msgID string senderID string } var dlist []drow for rows.Next() { var r drow if scanErr := rows.Scan(&r.seq, &r.endpointID, &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 { pushed, execErr := message.RejectPendingTx(tx, r.seq, r.endpointID, reasonGroupDissolved, nowMs) if execErr != nil { return execErr } if pushed && revokes != nil { *revokes = append(*revokes, revokeItem{ endpointID: r.endpointID, msgID: r.msgID, fromID: r.senderID, reason: reasonGroupDissolved, }) } if e := message.TryFinalizeTx(tx, r.seq, nowMs, days); e != nil { return e } } srows, err := tx.Query(` SELECT seq, 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 senderID string receipt int } var slist []srow for srows.Next() { var r srow if scanErr := srows.Scan(&r.seq, &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 e := message.FinalizeMessageTx(tx, r.seq, r.receipt != 0, r.senderID, "", reasonGroupDissolved, nowMs, days); e != nil { return e } } return nil } 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 }