fix: delete user cleanup (#5981)
This commit is contained in:
parent
d1208a68e9
commit
e53b7d96e7
20 changed files with 1985 additions and 747 deletions
|
|
@ -189,7 +189,7 @@ func (s *APIV1Service) resolveSSOUser(ctx context.Context, currentUser *store.Us
|
|||
ExternUID: externUID,
|
||||
}); err != nil {
|
||||
// Best-effort cleanup: the provisional user row has no linkage and should not remain.
|
||||
_ = s.Store.DeleteUser(ctx, &store.DeleteUser{ID: user.ID})
|
||||
_, _ = s.Store.DeleteUser(ctx, &store.DeleteUser{ID: user.ID})
|
||||
if isUniqueConstraintViolation(err) {
|
||||
// Concurrent first login won the race; load the winning linkage's user.
|
||||
winner, getErr := s.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
|
||||
|
|
|
|||
|
|
@ -402,7 +402,7 @@ func TestListMemosSkipsReactionsWithMissingCreators(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemos(ownerCtx, &apiv1.ListMemosRequest{PageSize: 10})
|
||||
|
|
@ -442,7 +442,7 @@ func TestListMemosSkipsMemosWithMissingCreators(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: orphanCreator.ID})
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: orphanCreator.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemos(ownerCtx, &apiv1.ListMemosRequest{PageSize: 10})
|
||||
|
|
@ -482,7 +482,7 @@ func TestListMemoCommentsSkipsCommentsWithMissingCreators(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: commenter.ID})
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: commenter.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemoComments(ownerCtx, &apiv1.ListMemoCommentsRequest{Name: memo.Name})
|
||||
|
|
|
|||
|
|
@ -147,7 +147,7 @@ func TestGetMemoByShare_SkipsReactionsWithMissingCreators(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
shareToken := share.Name[strings.LastIndex(share.Name, "/")+1:]
|
||||
|
|
|
|||
|
|
@ -226,7 +226,7 @@ func TestListMemoReactionsSkipsMissingCreators(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: reactor.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListMemoReactions(ctx, &apiv1.ListMemoReactionsRequest{Name: memo.Name})
|
||||
|
|
|
|||
|
|
@ -360,7 +360,7 @@ func TestListUserNotificationsSkipsNotificationsWithMissingUsers(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: commenter.ID})
|
||||
_, err = ts.Store.DeleteUser(ctx, &store.DeleteUser{ID: commenter.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := ts.Service.ListUserNotifications(ownerCtx, &apiv1.ListUserNotificationsRequest{
|
||||
|
|
|
|||
|
|
@ -373,12 +373,13 @@ func (s *APIV1Service) DeleteUser(ctx context.Context, request *v1pb.DeleteUserR
|
|||
}
|
||||
isSelfDelete := currentUser.ID == userID
|
||||
|
||||
attachments, err := s.Store.DeleteUserCompletely(ctx, &store.DeleteUser{
|
||||
deleteResult, err := s.Store.DeleteUser(ctx, &store.DeleteUser{
|
||||
ID: user.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to delete user: %v", err)
|
||||
}
|
||||
attachments := deleteResult.Attachments
|
||||
var attachmentCleanupErr error
|
||||
failedAttachmentIDs := make([]int32, 0)
|
||||
attachmentStorageSetting, attachmentStorageSettingErr := getDeleteUserAttachmentStorageSetting(ctx, s.Store, attachments)
|
||||
|
|
|
|||
|
|
@ -1,64 +0,0 @@
|
|||
package store
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
)
|
||||
|
||||
func TestAttachmentNeedsInstanceStorageSetting(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
attachment *Attachment
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "nil attachment",
|
||||
},
|
||||
{
|
||||
name: "local attachment",
|
||||
attachment: &Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_LOCAL,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "s3 attachment without payload",
|
||||
attachment: &Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_S3,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "s3 attachment with embedded config",
|
||||
attachment: &Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_S3,
|
||||
Payload: &storepb.AttachmentPayload{
|
||||
Payload: &storepb.AttachmentPayload_S3Object_{
|
||||
S3Object: &storepb.AttachmentPayload_S3Object{
|
||||
S3Config: &storepb.StorageS3Config{},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "s3 attachment without embedded config",
|
||||
attachment: &Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_S3,
|
||||
Payload: &storepb.AttachmentPayload{
|
||||
Payload: &storepb.AttachmentPayload_S3Object_{
|
||||
S3Object: &storepb.AttachmentPayload_S3Object{},
|
||||
},
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if got := AttachmentNeedsInstanceStorageSetting(test.attachment); got != test.want {
|
||||
t.Fatalf("AttachmentNeedsInstanceStorageSetting() = %v, want %v", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -192,14 +192,3 @@ func (d *DB) GetUser(ctx context.Context, find *store.FindUser) (*store.User, er
|
|||
}
|
||||
return list[0], nil
|
||||
}
|
||||
|
||||
func (d *DB) DeleteUser(ctx context.Context, delete *store.DeleteUser) error {
|
||||
result, err := d.db.ExecContext(ctx, "DELETE FROM `user` WHERE `id` = ?", delete.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := result.RowsAffected(); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
569
store/db/mysql/user_delete.go
Normal file
569
store/db/mysql/user_delete.go
Normal file
|
|
@ -0,0 +1,569 @@
|
|||
package mysql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const deleteUserBatchSize = 500
|
||||
|
||||
type deleteUserMemoRef struct {
|
||||
ID int32
|
||||
UID string
|
||||
}
|
||||
|
||||
type deleteUserTargetSet struct {
|
||||
memos []deleteUserMemoRef
|
||||
attachments []*store.Attachment
|
||||
attachmentIDs []int32
|
||||
userSettingKeys []storepb.UserSetting_Key
|
||||
inboxIDs []int32
|
||||
}
|
||||
|
||||
func (d *DB) DeleteUser(ctx context.Context, delete *store.DeleteUser) (*store.DeleteUserResult, error) {
|
||||
tx, err := d.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to begin delete user transaction")
|
||||
}
|
||||
defer func() {
|
||||
_ = tx.Rollback()
|
||||
}()
|
||||
|
||||
targets, err := collectDeleteUserTargets(ctx, tx, delete.ID)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to collect delete user targets")
|
||||
}
|
||||
|
||||
if err := deleteUserTargetsTx(ctx, tx, delete.ID, targets); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to delete user targets")
|
||||
}
|
||||
|
||||
if store.GetDeleteUserFailpoint(ctx) == store.DeleteUserFailpointBeforeCommit {
|
||||
return nil, errors.New("delete user failpoint before commit")
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to commit delete user transaction")
|
||||
}
|
||||
|
||||
return &store.DeleteUserResult{
|
||||
Attachments: targets.attachments,
|
||||
UserSettingKeys: targets.userSettingKeys,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func collectDeleteUserTargets(ctx context.Context, tx *sql.Tx, userID int32) (*deleteUserTargetSet, error) {
|
||||
targets := &deleteUserTargetSet{}
|
||||
|
||||
memos, err := listDeleteUserMemoTree(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.memos = memos
|
||||
|
||||
attachments, err := listDeleteUserAttachments(ctx, tx, userID, memoIDsFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.attachments = attachments
|
||||
targets.attachmentIDs = attachmentIDsFromList(attachments)
|
||||
|
||||
userSettingKeys, err := listDeleteUserSettingKeys(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.userSettingKeys = userSettingKeys
|
||||
|
||||
inboxIDs, err := listDeleteUserInboxIDs(ctx, tx, userID, memoIDSetFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.inboxIDs = inboxIDs
|
||||
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
func deleteUserTargetsTx(ctx context.Context, tx *sql.Tx, userID int32, targets *deleteUserTargetSet) error {
|
||||
memoIDs := memoIDsFromRefs(targets.memos)
|
||||
contentIDs := memoContentIDsFromRefs(targets.memos)
|
||||
|
||||
if err := deleteReactionsByContentIDsTx(ctx, tx, contentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteAttachmentsByIDsTx(ctx, tx, targets.attachmentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteReactionsByCreatorTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoSharesTx(ctx, tx, userID, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteInboxesByIDsTx(ctx, tx, targets.inboxIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserIdentitiesTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserSettingsTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoRelationsTx(ctx, tx, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemosTx(ctx, tx, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserRowTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTree(ctx context.Context, tx *sql.Tx, userID int32) ([]deleteUserMemoRef, error) {
|
||||
return listDeleteUserMemoTreeIterative(ctx, tx, userID)
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTreeIterative(ctx context.Context, tx *sql.Tx, userID int32) ([]deleteUserMemoRef, error) {
|
||||
roots, err := queryDeleteUserMemoRefs(ctx, tx, `
|
||||
SELECT id, uid
|
||||
FROM memo
|
||||
WHERE creator_id = `+deleteUserPlaceholder(1), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0, len(roots))
|
||||
seen := make(map[int32]struct{})
|
||||
frontier := make([]int32, 0, len(roots))
|
||||
for _, memo := range roots {
|
||||
if _, exists := seen[memo.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[memo.ID] = struct{}{}
|
||||
memos = append(memos, memo)
|
||||
frontier = append(frontier, memo.ID)
|
||||
}
|
||||
|
||||
for len(frontier) > 0 {
|
||||
currentFrontier := frontier
|
||||
nextFrontier := make([]int32, 0)
|
||||
for _, batch := range deleteUserBatches(currentFrontier, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
children, err := queryDeleteUserMemoRefs(ctx, tx, `
|
||||
SELECT child.id, child.uid
|
||||
FROM memo child
|
||||
JOIN memo_relation rel ON rel.memo_id = child.id AND rel.type = 'COMMENT'
|
||||
WHERE rel.related_memo_id IN `+clause, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, child := range children {
|
||||
if _, exists := seen[child.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[child.ID] = struct{}{}
|
||||
memos = append(memos, child)
|
||||
nextFrontier = append(nextFrontier, child.ID)
|
||||
}
|
||||
}
|
||||
frontier = nextFrontier
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func queryDeleteUserMemoRefs(ctx context.Context, tx *sql.Tx, query string, args ...any) ([]deleteUserMemoRef, error) {
|
||||
rows, err := tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0)
|
||||
for rows.Next() {
|
||||
var memo deleteUserMemoRef
|
||||
if err := rows.Scan(&memo.ID, &memo.UID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
memos = append(memos, memo)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func listDeleteUserAttachments(ctx context.Context, tx *sql.Tx, userID int32, memoIDs []int32) ([]*store.Attachment, error) {
|
||||
attachments := make([]*store.Attachment, 0)
|
||||
seen := make(map[int32]struct{})
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE creator_id = `+deleteUserPlaceholder(1), []any{userID}, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE memo_id IN `+clause, args, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return attachments, nil
|
||||
}
|
||||
|
||||
func appendDeleteUserAttachments(ctx context.Context, tx *sql.Tx, query string, args []any, seen map[int32]struct{}, attachments *[]*store.Attachment) error {
|
||||
rows, err := tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
attachment := &store.Attachment{}
|
||||
var memoID sql.NullInt32
|
||||
var storageType string
|
||||
var payloadBytes []byte
|
||||
if err := rows.Scan(&attachment.ID, &attachment.UID, &attachment.CreatorID, &memoID, &storageType, &attachment.Reference, &payloadBytes); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, exists := seen[attachment.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[attachment.ID] = struct{}{}
|
||||
if memoID.Valid {
|
||||
attachment.MemoID = &memoID.Int32
|
||||
}
|
||||
attachment.StorageType = storepb.AttachmentStorageType(storepb.AttachmentStorageType_value[storageType])
|
||||
payload := &storepb.AttachmentPayload{}
|
||||
if len(payloadBytes) > 0 {
|
||||
if err := protojsonUnmarshaler.Unmarshal(payloadBytes, payload); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
attachment.Payload = payload
|
||||
*attachments = append(*attachments, attachment)
|
||||
}
|
||||
return rows.Err()
|
||||
}
|
||||
|
||||
func listDeleteUserSettingKeys(ctx context.Context, tx *sql.Tx, userID int32) ([]storepb.UserSetting_Key, error) {
|
||||
rows, err := tx.QueryContext(ctx, deleteUserSettingKeysQuery(), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
keys := make([]storepb.UserSetting_Key, 0)
|
||||
for rows.Next() {
|
||||
var keyString string
|
||||
if err := rows.Scan(&keyString); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
key := storepb.UserSetting_Key(storepb.UserSetting_Key_value[keyString])
|
||||
keys = append(keys, key)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func deleteUserSettingKeysQuery() string {
|
||||
return "SELECT `key` FROM `user_setting` WHERE user_id = " + deleteUserPlaceholder(1)
|
||||
}
|
||||
|
||||
func listDeleteUserInboxIDs(ctx context.Context, tx *sql.Tx, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
directIDs, err := listDeleteUserDirectInboxIDs(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs := append([]int32{}, directIDs...)
|
||||
if len(memoIDSet) == 0 {
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
memoIDs, err := listDeleteUserMemoReferencedInboxIDs(ctx, tx, userID, memoIDSet)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return append(inboxIDs, memoIDs...), nil
|
||||
}
|
||||
|
||||
func listDeleteUserDirectInboxIDs(ctx context.Context, tx *sql.Tx, userID int32) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id
|
||||
FROM inbox
|
||||
WHERE sender_id = `+deleteUserPlaceholder(1)+`
|
||||
OR receiver_id = `+deleteUserPlaceholder(2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var inboxID int32
|
||||
if err := rows.Scan(&inboxID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoReferencedInboxIDs(ctx context.Context, tx *sql.Tx, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id, message
|
||||
FROM inbox
|
||||
WHERE sender_id <> `+deleteUserPlaceholder(1)+`
|
||||
AND receiver_id <> `+deleteUserPlaceholder(2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var (
|
||||
inboxID int32
|
||||
messageRaw []byte
|
||||
)
|
||||
if err := rows.Scan(&inboxID, &messageRaw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(messageRaw) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
message := &storepb.InboxMessage{}
|
||||
if err := protojsonUnmarshaler.Unmarshal(messageRaw, message); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if inboxMessageTouchesMemoSet(message, memoIDSet) {
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func inboxMessageTouchesMemoSet(message *storepb.InboxMessage, memoIDSet map[int32]struct{}) bool {
|
||||
if message == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
switch message.Type {
|
||||
case storepb.InboxMessage_MEMO_COMMENT:
|
||||
payload := message.GetMemoComment()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
case storepb.InboxMessage_MEMO_MENTION:
|
||||
payload := message.GetMemoMention()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func memoIDInSet(id int32, memoIDSet map[int32]struct{}) bool {
|
||||
if id == 0 {
|
||||
return false
|
||||
}
|
||||
_, exists := memoIDSet[id]
|
||||
return exists
|
||||
}
|
||||
|
||||
func deleteReactionsByContentIDsTx(ctx context.Context, tx *sql.Tx, contentIDs []string) error {
|
||||
for _, batch := range deleteUserBatches(contentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE content_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteAttachmentsByIDsTx(ctx context.Context, tx *sql.Tx, attachmentIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(attachmentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM attachment WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteReactionsByCreatorTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE creator_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoSharesTx(ctx context.Context, tx *sql.Tx, userID int32, memoIDs []int32) error {
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE creator_id = `+deleteUserPlaceholder(1), userID); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE memo_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteInboxesByIDsTx(ctx context.Context, tx *sql.Tx, inboxIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(inboxIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM inbox WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserIdentitiesTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_identity WHERE user_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserSettingsTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, "DELETE FROM `user_setting` WHERE user_id = "+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoRelationsTx(ctx context.Context, tx *sql.Tx, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
memoClause, args := deleteUserInClause(1, batch)
|
||||
relatedClause, relatedArgs := deleteUserInClause(len(args)+1, batch)
|
||||
query := `DELETE FROM memo_relation WHERE memo_id IN ` + memoClause + ` OR related_memo_id IN ` + relatedClause
|
||||
args = append(args, relatedArgs...)
|
||||
if _, err := tx.ExecContext(ctx, query, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteMemosTx(ctx context.Context, tx *sql.Tx, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserRowTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, "DELETE FROM `user` WHERE id = "+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserPlaceholder(_ int) string {
|
||||
return "?"
|
||||
}
|
||||
|
||||
func deleteUserInClause[T any](start int, values []T) (string, []any) {
|
||||
placeholders := make([]string, 0, len(values))
|
||||
args := make([]any, 0, len(values))
|
||||
for i, value := range values {
|
||||
placeholders = append(placeholders, deleteUserPlaceholder(start+i))
|
||||
args = append(args, value)
|
||||
}
|
||||
return "(" + strings.Join(placeholders, ", ") + ")", args
|
||||
}
|
||||
|
||||
func deleteUserBatches[T any](values []T, size int) [][]T {
|
||||
if len(values) == 0 {
|
||||
return nil
|
||||
}
|
||||
if size <= 0 {
|
||||
size = len(values)
|
||||
}
|
||||
|
||||
batches := make([][]T, 0, (len(values)+size-1)/size)
|
||||
for start := 0; start < len(values); start += size {
|
||||
end := start + size
|
||||
if end > len(values) {
|
||||
end = len(values)
|
||||
}
|
||||
batches = append(batches, values[start:end])
|
||||
}
|
||||
return batches
|
||||
}
|
||||
|
||||
func memoIDsFromRefs(memos []deleteUserMemoRef) []int32 {
|
||||
ids := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
ids = append(ids, memo.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func memoIDSetFromRefs(memos []deleteUserMemoRef) map[int32]struct{} {
|
||||
idSet := make(map[int32]struct{}, len(memos))
|
||||
for _, memo := range memos {
|
||||
idSet[memo.ID] = struct{}{}
|
||||
}
|
||||
return idSet
|
||||
}
|
||||
|
||||
func memoContentIDsFromRefs(memos []deleteUserMemoRef) []string {
|
||||
contentIDs := make([]string, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
contentIDs = append(contentIDs, "memos/"+memo.UID)
|
||||
}
|
||||
return contentIDs
|
||||
}
|
||||
|
||||
func attachmentIDsFromList(attachments []*store.Attachment) []int32 {
|
||||
ids := make([]int32, 0, len(attachments))
|
||||
for _, attachment := range attachments {
|
||||
if attachment == nil {
|
||||
continue
|
||||
}
|
||||
ids = append(ids, attachment.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
11
store/db/mysql/user_delete_test.go
Normal file
11
store/db/mysql/user_delete_test.go
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
package mysql
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestDeleteUserSettingKeysQueryQuotesReservedKey(t *testing.T) {
|
||||
got := deleteUserSettingKeysQuery()
|
||||
want := "SELECT `key` FROM `user_setting` WHERE user_id = ?"
|
||||
if got != want {
|
||||
t.Fatalf("deleteUserSettingKeysQuery() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
|
@ -190,14 +190,3 @@ func (d *DB) ListUsers(ctx context.Context, find *store.FindUser) ([]*store.User
|
|||
|
||||
return list, nil
|
||||
}
|
||||
|
||||
func (d *DB) DeleteUser(ctx context.Context, delete *store.DeleteUser) error {
|
||||
result, err := d.db.ExecContext(ctx, `DELETE FROM "user" WHERE id = $1`, delete.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := result.RowsAffected(); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
532
store/db/postgres/user_delete.go
Normal file
532
store/db/postgres/user_delete.go
Normal file
|
|
@ -0,0 +1,532 @@
|
|||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const deleteUserBatchSize = 500
|
||||
|
||||
type deleteUserMemoRef struct {
|
||||
ID int32
|
||||
UID string
|
||||
}
|
||||
|
||||
type deleteUserTargetSet struct {
|
||||
memos []deleteUserMemoRef
|
||||
attachments []*store.Attachment
|
||||
attachmentIDs []int32
|
||||
userSettingKeys []storepb.UserSetting_Key
|
||||
inboxIDs []int32
|
||||
}
|
||||
|
||||
func (d *DB) DeleteUser(ctx context.Context, delete *store.DeleteUser) (*store.DeleteUserResult, error) {
|
||||
tx, err := d.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to begin delete user transaction")
|
||||
}
|
||||
defer func() {
|
||||
_ = tx.Rollback()
|
||||
}()
|
||||
|
||||
targets, err := collectDeleteUserTargets(ctx, tx, delete.ID)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to collect delete user targets")
|
||||
}
|
||||
|
||||
if err := deleteUserTargetsTx(ctx, tx, delete.ID, targets); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to delete user targets")
|
||||
}
|
||||
|
||||
if store.GetDeleteUserFailpoint(ctx) == store.DeleteUserFailpointBeforeCommit {
|
||||
return nil, errors.New("delete user failpoint before commit")
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to commit delete user transaction")
|
||||
}
|
||||
|
||||
return &store.DeleteUserResult{
|
||||
Attachments: targets.attachments,
|
||||
UserSettingKeys: targets.userSettingKeys,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func collectDeleteUserTargets(ctx context.Context, tx *sql.Tx, userID int32) (*deleteUserTargetSet, error) {
|
||||
targets := &deleteUserTargetSet{}
|
||||
|
||||
memos, err := listDeleteUserMemoTree(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.memos = memos
|
||||
|
||||
attachments, err := listDeleteUserAttachments(ctx, tx, userID, memoIDsFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.attachments = attachments
|
||||
targets.attachmentIDs = attachmentIDsFromList(attachments)
|
||||
|
||||
userSettingKeys, err := listDeleteUserSettingKeys(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.userSettingKeys = userSettingKeys
|
||||
|
||||
inboxIDs, err := listDeleteUserInboxIDs(ctx, tx, userID, memoIDSetFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.inboxIDs = inboxIDs
|
||||
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
func deleteUserTargetsTx(ctx context.Context, tx *sql.Tx, userID int32, targets *deleteUserTargetSet) error {
|
||||
memoIDs := memoIDsFromRefs(targets.memos)
|
||||
contentIDs := memoContentIDsFromRefs(targets.memos)
|
||||
|
||||
if err := deleteReactionsByContentIDsTx(ctx, tx, contentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteAttachmentsByIDsTx(ctx, tx, targets.attachmentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteReactionsByCreatorTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoSharesTx(ctx, tx, userID, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteInboxesByIDsTx(ctx, tx, targets.inboxIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserIdentitiesTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserSettingsTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoRelationsTx(ctx, tx, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemosTx(ctx, tx, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserRowTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTree(ctx context.Context, tx *sql.Tx, userID int32) ([]deleteUserMemoRef, error) {
|
||||
return listDeleteUserMemoTreeRecursive(ctx, tx, userID)
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTreeRecursive(ctx context.Context, tx *sql.Tx, userID int32) ([]deleteUserMemoRef, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
WITH RECURSIVE memo_tree(id, uid) AS (
|
||||
SELECT id, uid
|
||||
FROM memo
|
||||
WHERE creator_id = `+deleteUserPlaceholder(1)+`
|
||||
UNION
|
||||
SELECT child.id, child.uid
|
||||
FROM memo child
|
||||
JOIN memo_relation rel ON rel.memo_id = child.id AND rel.type = 'COMMENT'
|
||||
JOIN memo_tree parent ON rel.related_memo_id = parent.id
|
||||
)
|
||||
SELECT id, uid
|
||||
FROM memo_tree
|
||||
`, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0)
|
||||
for rows.Next() {
|
||||
var memo deleteUserMemoRef
|
||||
if err := rows.Scan(&memo.ID, &memo.UID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
memos = append(memos, memo)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func listDeleteUserAttachments(ctx context.Context, tx *sql.Tx, userID int32, memoIDs []int32) ([]*store.Attachment, error) {
|
||||
attachments := make([]*store.Attachment, 0)
|
||||
seen := make(map[int32]struct{})
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE creator_id = `+deleteUserPlaceholder(1), []any{userID}, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE memo_id IN `+clause, args, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return attachments, nil
|
||||
}
|
||||
|
||||
func appendDeleteUserAttachments(ctx context.Context, tx *sql.Tx, query string, args []any, seen map[int32]struct{}, attachments *[]*store.Attachment) error {
|
||||
rows, err := tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
attachment := &store.Attachment{}
|
||||
var memoID sql.NullInt32
|
||||
var storageType string
|
||||
var payloadBytes []byte
|
||||
if err := rows.Scan(&attachment.ID, &attachment.UID, &attachment.CreatorID, &memoID, &storageType, &attachment.Reference, &payloadBytes); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, exists := seen[attachment.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[attachment.ID] = struct{}{}
|
||||
if memoID.Valid {
|
||||
attachment.MemoID = &memoID.Int32
|
||||
}
|
||||
attachment.StorageType = storepb.AttachmentStorageType(storepb.AttachmentStorageType_value[storageType])
|
||||
payload := &storepb.AttachmentPayload{}
|
||||
if len(payloadBytes) > 0 {
|
||||
if err := protojsonUnmarshaler.Unmarshal(payloadBytes, payload); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
attachment.Payload = payload
|
||||
*attachments = append(*attachments, attachment)
|
||||
}
|
||||
return rows.Err()
|
||||
}
|
||||
|
||||
func listDeleteUserSettingKeys(ctx context.Context, tx *sql.Tx, userID int32) ([]storepb.UserSetting_Key, error) {
|
||||
rows, err := tx.QueryContext(ctx, deleteUserSettingKeysQuery(), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
keys := make([]storepb.UserSetting_Key, 0)
|
||||
for rows.Next() {
|
||||
var keyString string
|
||||
if err := rows.Scan(&keyString); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
key := storepb.UserSetting_Key(storepb.UserSetting_Key_value[keyString])
|
||||
keys = append(keys, key)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func deleteUserSettingKeysQuery() string {
|
||||
return `SELECT key FROM user_setting WHERE user_id = ` + deleteUserPlaceholder(1)
|
||||
}
|
||||
|
||||
func listDeleteUserInboxIDs(ctx context.Context, tx *sql.Tx, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
directIDs, err := listDeleteUserDirectInboxIDs(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs := append([]int32{}, directIDs...)
|
||||
if len(memoIDSet) == 0 {
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
memoIDs, err := listDeleteUserMemoReferencedInboxIDs(ctx, tx, userID, memoIDSet)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return append(inboxIDs, memoIDs...), nil
|
||||
}
|
||||
|
||||
func listDeleteUserDirectInboxIDs(ctx context.Context, tx *sql.Tx, userID int32) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id
|
||||
FROM inbox
|
||||
WHERE sender_id = `+deleteUserPlaceholder(1)+`
|
||||
OR receiver_id = `+deleteUserPlaceholder(2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var inboxID int32
|
||||
if err := rows.Scan(&inboxID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoReferencedInboxIDs(ctx context.Context, tx *sql.Tx, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id, message
|
||||
FROM inbox
|
||||
WHERE sender_id <> `+deleteUserPlaceholder(1)+`
|
||||
AND receiver_id <> `+deleteUserPlaceholder(2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var (
|
||||
inboxID int32
|
||||
messageRaw []byte
|
||||
)
|
||||
if err := rows.Scan(&inboxID, &messageRaw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(messageRaw) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
message := &storepb.InboxMessage{}
|
||||
if err := protojsonUnmarshaler.Unmarshal(messageRaw, message); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if inboxMessageTouchesMemoSet(message, memoIDSet) {
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func inboxMessageTouchesMemoSet(message *storepb.InboxMessage, memoIDSet map[int32]struct{}) bool {
|
||||
if message == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
switch message.Type {
|
||||
case storepb.InboxMessage_MEMO_COMMENT:
|
||||
payload := message.GetMemoComment()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
case storepb.InboxMessage_MEMO_MENTION:
|
||||
payload := message.GetMemoMention()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func memoIDInSet(id int32, memoIDSet map[int32]struct{}) bool {
|
||||
if id == 0 {
|
||||
return false
|
||||
}
|
||||
_, exists := memoIDSet[id]
|
||||
return exists
|
||||
}
|
||||
|
||||
func deleteReactionsByContentIDsTx(ctx context.Context, tx *sql.Tx, contentIDs []string) error {
|
||||
for _, batch := range deleteUserBatches(contentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE content_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteAttachmentsByIDsTx(ctx context.Context, tx *sql.Tx, attachmentIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(attachmentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM attachment WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteReactionsByCreatorTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE creator_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoSharesTx(ctx context.Context, tx *sql.Tx, userID int32, memoIDs []int32) error {
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE creator_id = `+deleteUserPlaceholder(1), userID); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE memo_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteInboxesByIDsTx(ctx context.Context, tx *sql.Tx, inboxIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(inboxIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM inbox WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserIdentitiesTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_identity WHERE user_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserSettingsTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_setting WHERE user_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoRelationsTx(ctx context.Context, tx *sql.Tx, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
memoClause, args := deleteUserInClause(1, batch)
|
||||
relatedClause, relatedArgs := deleteUserInClause(len(args)+1, batch)
|
||||
query := `DELETE FROM memo_relation WHERE memo_id IN ` + memoClause + ` OR related_memo_id IN ` + relatedClause
|
||||
args = append(args, relatedArgs...)
|
||||
if _, err := tx.ExecContext(ctx, query, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteMemosTx(ctx context.Context, tx *sql.Tx, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserRowTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM "user" WHERE id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserPlaceholder(index int) string {
|
||||
return placeholder(index)
|
||||
}
|
||||
|
||||
func deleteUserInClause[T any](start int, values []T) (string, []any) {
|
||||
placeholders := make([]string, 0, len(values))
|
||||
args := make([]any, 0, len(values))
|
||||
for i, value := range values {
|
||||
placeholders = append(placeholders, deleteUserPlaceholder(start+i))
|
||||
args = append(args, value)
|
||||
}
|
||||
return "(" + strings.Join(placeholders, ", ") + ")", args
|
||||
}
|
||||
|
||||
func deleteUserBatches[T any](values []T, size int) [][]T {
|
||||
if len(values) == 0 {
|
||||
return nil
|
||||
}
|
||||
if size <= 0 {
|
||||
size = len(values)
|
||||
}
|
||||
|
||||
batches := make([][]T, 0, (len(values)+size-1)/size)
|
||||
for start := 0; start < len(values); start += size {
|
||||
end := start + size
|
||||
if end > len(values) {
|
||||
end = len(values)
|
||||
}
|
||||
batches = append(batches, values[start:end])
|
||||
}
|
||||
return batches
|
||||
}
|
||||
|
||||
func memoIDsFromRefs(memos []deleteUserMemoRef) []int32 {
|
||||
ids := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
ids = append(ids, memo.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func memoIDSetFromRefs(memos []deleteUserMemoRef) map[int32]struct{} {
|
||||
idSet := make(map[int32]struct{}, len(memos))
|
||||
for _, memo := range memos {
|
||||
idSet[memo.ID] = struct{}{}
|
||||
}
|
||||
return idSet
|
||||
}
|
||||
|
||||
func memoContentIDsFromRefs(memos []deleteUserMemoRef) []string {
|
||||
contentIDs := make([]string, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
contentIDs = append(contentIDs, "memos/"+memo.UID)
|
||||
}
|
||||
return contentIDs
|
||||
}
|
||||
|
||||
func attachmentIDsFromList(attachments []*store.Attachment) []int32 {
|
||||
ids := make([]int32, 0, len(attachments))
|
||||
for _, attachment := range attachments {
|
||||
if attachment == nil {
|
||||
continue
|
||||
}
|
||||
ids = append(ids, attachment.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
|
@ -200,16 +200,3 @@ func (d *DB) ListUsers(ctx context.Context, find *store.FindUser) ([]*store.User
|
|||
|
||||
return list, nil
|
||||
}
|
||||
|
||||
func (d *DB) DeleteUser(ctx context.Context, delete *store.DeleteUser) error {
|
||||
result, err := d.db.ExecContext(ctx, `
|
||||
DELETE FROM user WHERE id = ?
|
||||
`, delete.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := result.RowsAffected(); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
|
|||
532
store/db/sqlite/user_delete.go
Normal file
532
store/db/sqlite/user_delete.go
Normal file
|
|
@ -0,0 +1,532 @@
|
|||
package sqlite
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const deleteUserBatchSize = 500
|
||||
|
||||
type deleteUserMemoRef struct {
|
||||
ID int32
|
||||
UID string
|
||||
}
|
||||
|
||||
type deleteUserTargetSet struct {
|
||||
memos []deleteUserMemoRef
|
||||
attachments []*store.Attachment
|
||||
attachmentIDs []int32
|
||||
userSettingKeys []storepb.UserSetting_Key
|
||||
inboxIDs []int32
|
||||
}
|
||||
|
||||
func (d *DB) DeleteUser(ctx context.Context, delete *store.DeleteUser) (*store.DeleteUserResult, error) {
|
||||
tx, err := d.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to begin delete user transaction")
|
||||
}
|
||||
defer func() {
|
||||
_ = tx.Rollback()
|
||||
}()
|
||||
|
||||
targets, err := collectDeleteUserTargets(ctx, tx, delete.ID)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to collect delete user targets")
|
||||
}
|
||||
|
||||
if err := deleteUserTargetsTx(ctx, tx, delete.ID, targets); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to delete user targets")
|
||||
}
|
||||
|
||||
if store.GetDeleteUserFailpoint(ctx) == store.DeleteUserFailpointBeforeCommit {
|
||||
return nil, errors.New("delete user failpoint before commit")
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to commit delete user transaction")
|
||||
}
|
||||
|
||||
return &store.DeleteUserResult{
|
||||
Attachments: targets.attachments,
|
||||
UserSettingKeys: targets.userSettingKeys,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func collectDeleteUserTargets(ctx context.Context, tx *sql.Tx, userID int32) (*deleteUserTargetSet, error) {
|
||||
targets := &deleteUserTargetSet{}
|
||||
|
||||
memos, err := listDeleteUserMemoTree(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.memos = memos
|
||||
|
||||
attachments, err := listDeleteUserAttachments(ctx, tx, userID, memoIDsFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.attachments = attachments
|
||||
targets.attachmentIDs = attachmentIDsFromList(attachments)
|
||||
|
||||
userSettingKeys, err := listDeleteUserSettingKeys(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.userSettingKeys = userSettingKeys
|
||||
|
||||
inboxIDs, err := listDeleteUserInboxIDs(ctx, tx, userID, memoIDSetFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.inboxIDs = inboxIDs
|
||||
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
func deleteUserTargetsTx(ctx context.Context, tx *sql.Tx, userID int32, targets *deleteUserTargetSet) error {
|
||||
memoIDs := memoIDsFromRefs(targets.memos)
|
||||
contentIDs := memoContentIDsFromRefs(targets.memos)
|
||||
|
||||
if err := deleteReactionsByContentIDsTx(ctx, tx, contentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteAttachmentsByIDsTx(ctx, tx, targets.attachmentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteReactionsByCreatorTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoSharesTx(ctx, tx, userID, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteInboxesByIDsTx(ctx, tx, targets.inboxIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserIdentitiesTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserSettingsTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoRelationsTx(ctx, tx, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemosTx(ctx, tx, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserRowTx(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTree(ctx context.Context, tx *sql.Tx, userID int32) ([]deleteUserMemoRef, error) {
|
||||
return listDeleteUserMemoTreeRecursive(ctx, tx, userID)
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTreeRecursive(ctx context.Context, tx *sql.Tx, userID int32) ([]deleteUserMemoRef, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
WITH RECURSIVE memo_tree(id, uid) AS (
|
||||
SELECT id, uid
|
||||
FROM memo
|
||||
WHERE creator_id = `+deleteUserPlaceholder(1)+`
|
||||
UNION
|
||||
SELECT child.id, child.uid
|
||||
FROM memo child
|
||||
JOIN memo_relation rel ON rel.memo_id = child.id AND rel.type = 'COMMENT'
|
||||
JOIN memo_tree parent ON rel.related_memo_id = parent.id
|
||||
)
|
||||
SELECT id, uid
|
||||
FROM memo_tree
|
||||
`, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0)
|
||||
for rows.Next() {
|
||||
var memo deleteUserMemoRef
|
||||
if err := rows.Scan(&memo.ID, &memo.UID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
memos = append(memos, memo)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func listDeleteUserAttachments(ctx context.Context, tx *sql.Tx, userID int32, memoIDs []int32) ([]*store.Attachment, error) {
|
||||
attachments := make([]*store.Attachment, 0)
|
||||
seen := make(map[int32]struct{})
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE creator_id = `+deleteUserPlaceholder(1), []any{userID}, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE memo_id IN `+clause, args, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return attachments, nil
|
||||
}
|
||||
|
||||
func appendDeleteUserAttachments(ctx context.Context, tx *sql.Tx, query string, args []any, seen map[int32]struct{}, attachments *[]*store.Attachment) error {
|
||||
rows, err := tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
attachment := &store.Attachment{}
|
||||
var memoID sql.NullInt32
|
||||
var storageType string
|
||||
var payloadBytes []byte
|
||||
if err := rows.Scan(&attachment.ID, &attachment.UID, &attachment.CreatorID, &memoID, &storageType, &attachment.Reference, &payloadBytes); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, exists := seen[attachment.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[attachment.ID] = struct{}{}
|
||||
if memoID.Valid {
|
||||
attachment.MemoID = &memoID.Int32
|
||||
}
|
||||
attachment.StorageType = storepb.AttachmentStorageType(storepb.AttachmentStorageType_value[storageType])
|
||||
payload := &storepb.AttachmentPayload{}
|
||||
if len(payloadBytes) > 0 {
|
||||
if err := protojsonUnmarshaler.Unmarshal(payloadBytes, payload); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
attachment.Payload = payload
|
||||
*attachments = append(*attachments, attachment)
|
||||
}
|
||||
return rows.Err()
|
||||
}
|
||||
|
||||
func listDeleteUserSettingKeys(ctx context.Context, tx *sql.Tx, userID int32) ([]storepb.UserSetting_Key, error) {
|
||||
rows, err := tx.QueryContext(ctx, deleteUserSettingKeysQuery(), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
keys := make([]storepb.UserSetting_Key, 0)
|
||||
for rows.Next() {
|
||||
var keyString string
|
||||
if err := rows.Scan(&keyString); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
key := storepb.UserSetting_Key(storepb.UserSetting_Key_value[keyString])
|
||||
keys = append(keys, key)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func deleteUserSettingKeysQuery() string {
|
||||
return `SELECT key FROM user_setting WHERE user_id = ` + deleteUserPlaceholder(1)
|
||||
}
|
||||
|
||||
func listDeleteUserInboxIDs(ctx context.Context, tx *sql.Tx, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
directIDs, err := listDeleteUserDirectInboxIDs(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs := append([]int32{}, directIDs...)
|
||||
if len(memoIDSet) == 0 {
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
memoIDs, err := listDeleteUserMemoReferencedInboxIDs(ctx, tx, userID, memoIDSet)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return append(inboxIDs, memoIDs...), nil
|
||||
}
|
||||
|
||||
func listDeleteUserDirectInboxIDs(ctx context.Context, tx *sql.Tx, userID int32) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id
|
||||
FROM inbox
|
||||
WHERE sender_id = `+deleteUserPlaceholder(1)+`
|
||||
OR receiver_id = `+deleteUserPlaceholder(2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var inboxID int32
|
||||
if err := rows.Scan(&inboxID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoReferencedInboxIDs(ctx context.Context, tx *sql.Tx, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id, message
|
||||
FROM inbox
|
||||
WHERE sender_id <> `+deleteUserPlaceholder(1)+`
|
||||
AND receiver_id <> `+deleteUserPlaceholder(2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var (
|
||||
inboxID int32
|
||||
messageRaw []byte
|
||||
)
|
||||
if err := rows.Scan(&inboxID, &messageRaw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(messageRaw) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
message := &storepb.InboxMessage{}
|
||||
if err := protojsonUnmarshaler.Unmarshal(messageRaw, message); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if inboxMessageTouchesMemoSet(message, memoIDSet) {
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func inboxMessageTouchesMemoSet(message *storepb.InboxMessage, memoIDSet map[int32]struct{}) bool {
|
||||
if message == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
switch message.Type {
|
||||
case storepb.InboxMessage_MEMO_COMMENT:
|
||||
payload := message.GetMemoComment()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
case storepb.InboxMessage_MEMO_MENTION:
|
||||
payload := message.GetMemoMention()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func memoIDInSet(id int32, memoIDSet map[int32]struct{}) bool {
|
||||
if id == 0 {
|
||||
return false
|
||||
}
|
||||
_, exists := memoIDSet[id]
|
||||
return exists
|
||||
}
|
||||
|
||||
func deleteReactionsByContentIDsTx(ctx context.Context, tx *sql.Tx, contentIDs []string) error {
|
||||
for _, batch := range deleteUserBatches(contentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE content_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteAttachmentsByIDsTx(ctx context.Context, tx *sql.Tx, attachmentIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(attachmentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM attachment WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteReactionsByCreatorTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE creator_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoSharesTx(ctx context.Context, tx *sql.Tx, userID int32, memoIDs []int32) error {
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE creator_id = `+deleteUserPlaceholder(1), userID); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE memo_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteInboxesByIDsTx(ctx context.Context, tx *sql.Tx, inboxIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(inboxIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM inbox WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserIdentitiesTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_identity WHERE user_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserSettingsTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_setting WHERE user_id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoRelationsTx(ctx context.Context, tx *sql.Tx, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
memoClause, args := deleteUserInClause(1, batch)
|
||||
relatedClause, relatedArgs := deleteUserInClause(len(args)+1, batch)
|
||||
query := `DELETE FROM memo_relation WHERE memo_id IN ` + memoClause + ` OR related_memo_id IN ` + relatedClause
|
||||
args = append(args, relatedArgs...)
|
||||
if _, err := tx.ExecContext(ctx, query, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteMemosTx(ctx context.Context, tx *sql.Tx, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserRowTx(ctx context.Context, tx *sql.Tx, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user WHERE id = `+deleteUserPlaceholder(1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserPlaceholder(_ int) string {
|
||||
return "?"
|
||||
}
|
||||
|
||||
func deleteUserInClause[T any](start int, values []T) (string, []any) {
|
||||
placeholders := make([]string, 0, len(values))
|
||||
args := make([]any, 0, len(values))
|
||||
for i, value := range values {
|
||||
placeholders = append(placeholders, deleteUserPlaceholder(start+i))
|
||||
args = append(args, value)
|
||||
}
|
||||
return "(" + strings.Join(placeholders, ", ") + ")", args
|
||||
}
|
||||
|
||||
func deleteUserBatches[T any](values []T, size int) [][]T {
|
||||
if len(values) == 0 {
|
||||
return nil
|
||||
}
|
||||
if size <= 0 {
|
||||
size = len(values)
|
||||
}
|
||||
|
||||
batches := make([][]T, 0, (len(values)+size-1)/size)
|
||||
for start := 0; start < len(values); start += size {
|
||||
end := start + size
|
||||
if end > len(values) {
|
||||
end = len(values)
|
||||
}
|
||||
batches = append(batches, values[start:end])
|
||||
}
|
||||
return batches
|
||||
}
|
||||
|
||||
func memoIDsFromRefs(memos []deleteUserMemoRef) []int32 {
|
||||
ids := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
ids = append(ids, memo.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func memoIDSetFromRefs(memos []deleteUserMemoRef) map[int32]struct{} {
|
||||
idSet := make(map[int32]struct{}, len(memos))
|
||||
for _, memo := range memos {
|
||||
idSet[memo.ID] = struct{}{}
|
||||
}
|
||||
return idSet
|
||||
}
|
||||
|
||||
func memoContentIDsFromRefs(memos []deleteUserMemoRef) []string {
|
||||
contentIDs := make([]string, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
contentIDs = append(contentIDs, "memos/"+memo.UID)
|
||||
}
|
||||
return contentIDs
|
||||
}
|
||||
|
||||
func attachmentIDsFromList(attachments []*store.Attachment) []int32 {
|
||||
ids := make([]int32, 0, len(attachments))
|
||||
for _, attachment := range attachments {
|
||||
if attachment == nil {
|
||||
continue
|
||||
}
|
||||
ids = append(ids, attachment.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
|
@ -45,7 +45,7 @@ type Driver interface {
|
|||
CreateUser(ctx context.Context, create *User) (*User, error)
|
||||
UpdateUser(ctx context.Context, update *UpdateUser) (*User, error)
|
||||
ListUsers(ctx context.Context, find *FindUser) ([]*User, error)
|
||||
DeleteUser(ctx context.Context, delete *DeleteUser) error
|
||||
DeleteUser(ctx context.Context, delete *DeleteUser) (*DeleteUserResult, error)
|
||||
|
||||
// UserSetting model related methods.
|
||||
UpsertUserSetting(ctx context.Context, upsert *UserSetting) (*UserSetting, error)
|
||||
|
|
|
|||
|
|
@ -8,9 +8,67 @@ import (
|
|||
"github.com/lithammer/shortuuid/v4"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestAttachmentNeedsInstanceStorageSetting(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
attachment *store.Attachment
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "nil attachment",
|
||||
},
|
||||
{
|
||||
name: "local attachment",
|
||||
attachment: &store.Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_LOCAL,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "s3 attachment without payload",
|
||||
attachment: &store.Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_S3,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "s3 attachment with embedded config",
|
||||
attachment: &store.Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_S3,
|
||||
Payload: &storepb.AttachmentPayload{
|
||||
Payload: &storepb.AttachmentPayload_S3Object_{
|
||||
S3Object: &storepb.AttachmentPayload_S3Object{
|
||||
S3Config: &storepb.StorageS3Config{},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "s3 attachment without embedded config",
|
||||
attachment: &store.Attachment{
|
||||
StorageType: storepb.AttachmentStorageType_S3,
|
||||
Payload: &storepb.AttachmentPayload{
|
||||
Payload: &storepb.AttachmentPayload_S3Object_{
|
||||
S3Object: &storepb.AttachmentPayload_S3Object{},
|
||||
},
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if got := store.AttachmentNeedsInstanceStorageSetting(test.attachment); got != test.want {
|
||||
t.Fatalf("AttachmentNeedsInstanceStorageSetting() = %v, want %v", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAttachmentStore(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
|
|
|||
248
store/test/user_delete_test.go
Normal file
248
store/test/user_delete_test.go
Normal file
|
|
@ -0,0 +1,248 @@
|
|||
package test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func TestDeleteUserCleansRelatedData(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := context.Background()
|
||||
ts := NewTestingStore(ctx, t)
|
||||
defer ts.Close()
|
||||
|
||||
user, err := createTestingHostUser(ctx, ts)
|
||||
require.NoError(t, err)
|
||||
peer, err := createTestingUserWithRole(ctx, ts, "delete-peer", store.RoleUser)
|
||||
require.NoError(t, err)
|
||||
|
||||
ownMemo, err := ts.CreateMemo(ctx, &store.Memo{
|
||||
UID: "delete-own-memo",
|
||||
CreatorID: user.ID,
|
||||
Content: "owner memo",
|
||||
Visibility: store.Public,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerMemo, err := ts.CreateMemo(ctx, &store.Memo{
|
||||
UID: "delete-peer-memo",
|
||||
CreatorID: peer.ID,
|
||||
Content: "peer memo",
|
||||
Visibility: store.Public,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerCommentOnOwnMemo, err := ts.CreateMemo(ctx, &store.Memo{
|
||||
UID: "delete-peer-comment",
|
||||
CreatorID: peer.ID,
|
||||
Content: "peer comment on owner memo",
|
||||
Visibility: store.Public,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.UpsertMemoRelation(ctx, &store.MemoRelation{
|
||||
MemoID: peerCommentOnOwnMemo.ID,
|
||||
RelatedMemoID: ownMemo.ID,
|
||||
Type: store.MemoRelationComment,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
userCommentOnPeerMemo, err := ts.CreateMemo(ctx, &store.Memo{
|
||||
UID: "delete-user-comment",
|
||||
CreatorID: user.ID,
|
||||
Content: "owner comment on peer memo",
|
||||
Visibility: store.Public,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.UpsertMemoRelation(ctx, &store.MemoRelation{
|
||||
MemoID: userCommentOnPeerMemo.ID,
|
||||
RelatedMemoID: peerMemo.ID,
|
||||
Type: store.MemoRelationComment,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
ownerAttachment, err := ts.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "delete-owner-attachment",
|
||||
CreatorID: user.ID,
|
||||
Filename: "owner.txt",
|
||||
Type: "text/plain",
|
||||
Size: 5,
|
||||
Blob: []byte("owner"),
|
||||
MemoID: &ownMemo.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerAttachmentOnDeletedMemo, err := ts.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "delete-peer-attachment",
|
||||
CreatorID: peer.ID,
|
||||
Filename: "peer-on-owner.txt",
|
||||
Type: "text/plain",
|
||||
Size: 4,
|
||||
Blob: []byte("peer"),
|
||||
MemoID: &ownMemo.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerAttachmentToKeep, err := ts.CreateAttachment(ctx, &store.Attachment{
|
||||
UID: "keep-peer-attachment",
|
||||
CreatorID: peer.ID,
|
||||
Filename: "peer.txt",
|
||||
Type: "text/plain",
|
||||
Size: 4,
|
||||
Blob: []byte("peer"),
|
||||
MemoID: &peerMemo.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: peer.ID,
|
||||
ContentID: "memos/" + ownMemo.UID,
|
||||
ReactionType: "thumbs-up",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: user.ID,
|
||||
ContentID: "memos/" + peerMemo.UID,
|
||||
ReactionType: "heart",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerReactionToKeep, err := ts.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: peer.ID,
|
||||
ContentID: "memos/" + peerMemo.UID,
|
||||
ReactionType: "sparkle",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "delete-owner-share",
|
||||
MemoID: peerMemo.ID,
|
||||
CreatorID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "delete-memo-share",
|
||||
MemoID: ownMemo.ID,
|
||||
CreatorID: peer.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
peerShareToKeep, err := ts.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: "keep-peer-share",
|
||||
MemoID: peerMemo.ID,
|
||||
CreatorID: peer.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.CreateInbox(ctx, &store.Inbox{
|
||||
SenderID: user.ID,
|
||||
ReceiverID: peer.ID,
|
||||
Status: store.UNREAD,
|
||||
Message: &storepb.InboxMessage{Type: storepb.InboxMessage_MEMO_MENTION},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.CreateInbox(ctx, &store.Inbox{
|
||||
SenderID: peer.ID,
|
||||
ReceiverID: peer.ID,
|
||||
Status: store.UNREAD,
|
||||
Message: &storepb.InboxMessage{
|
||||
Type: storepb.InboxMessage_MEMO_COMMENT,
|
||||
Payload: &storepb.InboxMessage_MemoComment{
|
||||
MemoComment: &storepb.InboxMessage_MemoCommentPayload{
|
||||
MemoId: ownMemo.ID,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
inboxToKeep, err := ts.CreateInbox(ctx, &store.Inbox{
|
||||
SenderID: peer.ID,
|
||||
ReceiverID: peer.ID,
|
||||
Status: store.UNREAD,
|
||||
Message: &storepb.InboxMessage{
|
||||
Type: storepb.InboxMessage_MEMO_COMMENT,
|
||||
Payload: &storepb.InboxMessage_MemoComment{
|
||||
MemoComment: &storepb.InboxMessage_MemoCommentPayload{
|
||||
MemoId: peerMemo.ID,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: user.ID,
|
||||
Provider: "google",
|
||||
ExternUID: "delete-user-sub",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = ts.AddUserPersonalAccessToken(ctx, user.ID, &storepb.PersonalAccessTokensUserSetting_PersonalAccessToken{
|
||||
TokenId: "delete-user-pat",
|
||||
TokenHash: "delete-user-pat-hash",
|
||||
Description: "delete user pat",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ts.DeleteUser(ctx, &store.DeleteUser{ID: user.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
deletedUser, err := ts.GetUser(ctx, &store.FindUser{ID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, deletedUser)
|
||||
keptUser, err := ts.GetUser(ctx, &store.FindUser{ID: &peer.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, keptUser)
|
||||
|
||||
for _, memo := range []*store.Memo{ownMemo, peerCommentOnOwnMemo, userCommentOnPeerMemo} {
|
||||
got, getErr := ts.GetMemo(ctx, &store.FindMemo{ID: &memo.ID})
|
||||
require.NoError(t, getErr)
|
||||
require.Nil(t, got, memo.UID)
|
||||
}
|
||||
keptMemo, err := ts.GetMemo(ctx, &store.FindMemo{ID: &peerMemo.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, keptMemo)
|
||||
|
||||
for _, attachment := range []*store.Attachment{ownerAttachment, peerAttachmentOnDeletedMemo} {
|
||||
got, getErr := ts.GetAttachment(ctx, &store.FindAttachment{ID: &attachment.ID})
|
||||
require.NoError(t, getErr)
|
||||
require.Nil(t, got, attachment.UID)
|
||||
}
|
||||
keptAttachment, err := ts.GetAttachment(ctx, &store.FindAttachment{ID: &peerAttachmentToKeep.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, keptAttachment)
|
||||
|
||||
deletedMemoRelations, err := ts.ListMemoRelations(ctx, &store.FindMemoRelation{MemoIDList: []int32{ownMemo.ID, peerCommentOnOwnMemo.ID, userCommentOnPeerMemo.ID}})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, deletedMemoRelations)
|
||||
|
||||
peerMemoContentID := "memos/" + peerMemo.UID
|
||||
keptReactions, err := ts.ListReactions(ctx, &store.FindReaction{ContentID: &peerMemoContentID})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, keptReactions, 1)
|
||||
require.Equal(t, peerReactionToKeep.ID, keptReactions[0].ID)
|
||||
|
||||
deletedOwnerShares, err := ts.ListMemoShares(ctx, &store.FindMemoShare{CreatorID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, deletedOwnerShares)
|
||||
keptShare, err := ts.GetMemoShare(ctx, &store.FindMemoShare{ID: &peerShareToKeep.ID})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, keptShare)
|
||||
|
||||
deletedSentInboxes, err := ts.ListInboxes(ctx, &store.FindInbox{SenderID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, deletedSentInboxes)
|
||||
deletedReceivedInboxes, err := ts.ListInboxes(ctx, &store.FindInbox{ReceiverID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, deletedReceivedInboxes)
|
||||
keptInboxes, err := ts.ListInboxes(ctx, &store.FindInbox{ID: &inboxToKeep.ID})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, keptInboxes, 1)
|
||||
|
||||
identities, err := ts.ListUserIdentities(ctx, &store.FindUserIdentity{UserID: &user.ID})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, identities)
|
||||
setting, err := ts.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &user.ID,
|
||||
Key: storepb.UserSetting_PERSONAL_ACCESS_TOKENS,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, setting)
|
||||
}
|
||||
|
|
@ -30,7 +30,7 @@ func TestUserStore(t *testing.T) {
|
|||
user, err = ts.UpdateUser(ctx, userPatch)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, userPatchNickname, user.Nickname)
|
||||
err = ts.DeleteUser(ctx, &store.DeleteUser{
|
||||
_, err = ts.DeleteUser(ctx, &store.DeleteUser{
|
||||
ID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ package store
|
|||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// Role is the type of a role.
|
||||
|
|
@ -164,11 +166,14 @@ func (s *Store) GetUser(ctx context.Context, find *FindUser) (*User, error) {
|
|||
return user, nil
|
||||
}
|
||||
|
||||
func (s *Store) DeleteUser(ctx context.Context, delete *DeleteUser) error {
|
||||
err := s.driver.DeleteUser(ctx, delete)
|
||||
func (s *Store) DeleteUser(ctx context.Context, delete *DeleteUser) (*DeleteUserResult, error) {
|
||||
result, err := s.driver.DeleteUser(ctx, delete)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
s.userCache.Delete(ctx, userCacheKey(delete.ID))
|
||||
return nil
|
||||
if result == nil {
|
||||
return nil, errors.New("unexpected nil delete user result")
|
||||
}
|
||||
s.deleteUserCache(ctx, delete.ID, result)
|
||||
return result, nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,11 +2,6 @@ package store
|
|||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
)
|
||||
|
|
@ -21,75 +16,19 @@ const (
|
|||
|
||||
type deleteUserFailpointKey struct{}
|
||||
|
||||
type deleteUserDialect string
|
||||
|
||||
const (
|
||||
deleteUserDialectSQLite deleteUserDialect = "sqlite"
|
||||
deleteUserDialectMySQL deleteUserDialect = "mysql"
|
||||
deleteUserDialectPostgres deleteUserDialect = "postgres"
|
||||
deleteUserBatchSize int = 500
|
||||
)
|
||||
|
||||
type deleteUserMemoRef struct {
|
||||
ID int32
|
||||
UID string
|
||||
// DeleteUserResult contains resources collected while deleting a user.
|
||||
type DeleteUserResult struct {
|
||||
Attachments []*Attachment
|
||||
UserSettingKeys []storepb.UserSetting_Key
|
||||
}
|
||||
|
||||
type deleteUserTargetSet struct {
|
||||
memos []deleteUserMemoRef
|
||||
attachments []*Attachment
|
||||
attachmentIDs []int32
|
||||
userSettingKeys []storepb.UserSetting_Key
|
||||
inboxIDs []int32
|
||||
}
|
||||
|
||||
// WithDeleteUserFailpoint is a test-only helper that forces DeleteUserCompletely to roll back.
|
||||
// WithDeleteUserFailpoint is a test-only helper that forces DeleteUser to roll back.
|
||||
func WithDeleteUserFailpoint(ctx context.Context, failpoint DeleteUserFailpoint) context.Context {
|
||||
return context.WithValue(ctx, deleteUserFailpointKey{}, failpoint)
|
||||
}
|
||||
|
||||
// DeleteUserCompletely deletes the user and all directly associated database resources in one transaction.
|
||||
// Attachment file/object cleanup must happen after commit because external storage cannot participate in SQL transactions.
|
||||
func (s *Store) DeleteUserCompletely(ctx context.Context, delete *DeleteUser) ([]*Attachment, error) {
|
||||
dialect, err := getDeleteUserDialect(s.profile.Driver)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
tx, err := s.driver.GetDB().BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to begin delete user transaction")
|
||||
}
|
||||
defer func() {
|
||||
_ = tx.Rollback()
|
||||
}()
|
||||
|
||||
targets, err := collectDeleteUserTargets(ctx, tx, dialect, delete.ID)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to collect delete user targets")
|
||||
}
|
||||
|
||||
if err := deleteUserTargetsTx(ctx, tx, dialect, delete.ID, targets); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to delete user targets")
|
||||
}
|
||||
|
||||
if getDeleteUserFailpoint(ctx) == DeleteUserFailpointBeforeCommit {
|
||||
return nil, errors.New("delete user failpoint before commit")
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to commit delete user transaction")
|
||||
}
|
||||
|
||||
s.userCache.Delete(ctx, userCacheKey(delete.ID))
|
||||
for _, key := range targets.userSettingKeys {
|
||||
s.userSettingCache.Delete(ctx, getUserSettingCacheKey(delete.ID, key.String()))
|
||||
}
|
||||
|
||||
return targets.attachments, nil
|
||||
}
|
||||
|
||||
func getDeleteUserFailpoint(ctx context.Context) DeleteUserFailpoint {
|
||||
// GetDeleteUserFailpoint returns the delete-user failpoint attached to ctx, if any.
|
||||
func GetDeleteUserFailpoint(ctx context.Context) DeleteUserFailpoint {
|
||||
failpoint, ok := ctx.Value(deleteUserFailpointKey{}).(DeleteUserFailpoint)
|
||||
if !ok {
|
||||
return ""
|
||||
|
|
@ -97,570 +36,12 @@ func getDeleteUserFailpoint(ctx context.Context) DeleteUserFailpoint {
|
|||
return failpoint
|
||||
}
|
||||
|
||||
func getDeleteUserDialect(driver string) (deleteUserDialect, error) {
|
||||
switch driver {
|
||||
case "sqlite":
|
||||
return deleteUserDialectSQLite, nil
|
||||
case "mysql":
|
||||
return deleteUserDialectMySQL, nil
|
||||
case "postgres":
|
||||
return deleteUserDialectPostgres, nil
|
||||
default:
|
||||
return "", errors.Errorf("unsupported delete user dialect: %s", driver)
|
||||
func (s *Store) deleteUserCache(ctx context.Context, userID int32, result *DeleteUserResult) {
|
||||
s.userCache.Delete(ctx, userCacheKey(userID))
|
||||
if result == nil {
|
||||
return
|
||||
}
|
||||
for _, key := range result.UserSettingKeys {
|
||||
s.userSettingCache.Delete(ctx, getUserSettingCacheKey(userID, key.String()))
|
||||
}
|
||||
}
|
||||
|
||||
func collectDeleteUserTargets(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) (*deleteUserTargetSet, error) {
|
||||
targets := &deleteUserTargetSet{}
|
||||
|
||||
memos, err := listDeleteUserMemoTree(ctx, tx, dialect, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.memos = memos
|
||||
|
||||
attachments, err := listDeleteUserAttachments(ctx, tx, dialect, userID, memoIDsFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.attachments = attachments
|
||||
targets.attachmentIDs = attachmentIDsFromList(attachments)
|
||||
|
||||
userSettingKeys, err := listDeleteUserSettingKeys(ctx, tx, dialect, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.userSettingKeys = userSettingKeys
|
||||
|
||||
inboxIDs, err := listDeleteUserInboxIDs(ctx, tx, dialect, userID, memoIDSetFromRefs(memos))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targets.inboxIDs = inboxIDs
|
||||
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
func deleteUserTargetsTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32, targets *deleteUserTargetSet) error {
|
||||
memoIDs := memoIDsFromRefs(targets.memos)
|
||||
contentIDs := memoContentIDsFromRefs(targets.memos)
|
||||
|
||||
if err := deleteReactionsByContentIDsTx(ctx, tx, dialect, contentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteAttachmentsByIDsTx(ctx, tx, dialect, targets.attachmentIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteReactionsByCreatorTx(ctx, tx, dialect, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoSharesTx(ctx, tx, dialect, userID, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteInboxesByIDsTx(ctx, tx, dialect, targets.inboxIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserIdentitiesTx(ctx, tx, dialect, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserSettingsTx(ctx, tx, dialect, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemoRelationsTx(ctx, tx, dialect, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteMemosTx(ctx, tx, dialect, memoIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteUserRowTx(ctx, tx, dialect, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTree(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) ([]deleteUserMemoRef, error) {
|
||||
if dialect == deleteUserDialectMySQL {
|
||||
return listDeleteUserMemoTreeIterative(ctx, tx, dialect, userID)
|
||||
}
|
||||
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
WITH RECURSIVE memo_tree(id, uid) AS (
|
||||
SELECT id, uid
|
||||
FROM memo
|
||||
WHERE creator_id = `+deleteUserPlaceholder(dialect, 1)+`
|
||||
UNION
|
||||
SELECT child.id, child.uid
|
||||
FROM memo child
|
||||
JOIN memo_relation rel ON rel.memo_id = child.id AND rel.type = 'COMMENT'
|
||||
JOIN memo_tree parent ON rel.related_memo_id = parent.id
|
||||
)
|
||||
SELECT id, uid
|
||||
FROM memo_tree
|
||||
`, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0)
|
||||
for rows.Next() {
|
||||
var memo deleteUserMemoRef
|
||||
if err := rows.Scan(&memo.ID, &memo.UID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
memos = append(memos, memo)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoTreeIterative(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) ([]deleteUserMemoRef, error) {
|
||||
roots, err := queryDeleteUserMemoRefs(ctx, tx, `
|
||||
SELECT id, uid
|
||||
FROM memo
|
||||
WHERE creator_id = `+deleteUserPlaceholder(dialect, 1), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0, len(roots))
|
||||
seen := make(map[int32]struct{})
|
||||
frontier := make([]int32, 0, len(roots))
|
||||
for _, memo := range roots {
|
||||
if _, exists := seen[memo.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[memo.ID] = struct{}{}
|
||||
memos = append(memos, memo)
|
||||
frontier = append(frontier, memo.ID)
|
||||
}
|
||||
|
||||
for len(frontier) > 0 {
|
||||
currentFrontier := frontier
|
||||
nextFrontier := make([]int32, 0)
|
||||
for _, batch := range deleteUserBatches(currentFrontier, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
children, err := queryDeleteUserMemoRefs(ctx, tx, `
|
||||
SELECT child.id, child.uid
|
||||
FROM memo child
|
||||
JOIN memo_relation rel ON rel.memo_id = child.id AND rel.type = 'COMMENT'
|
||||
WHERE rel.related_memo_id IN `+clause, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, child := range children {
|
||||
if _, exists := seen[child.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[child.ID] = struct{}{}
|
||||
memos = append(memos, child)
|
||||
nextFrontier = append(nextFrontier, child.ID)
|
||||
}
|
||||
}
|
||||
frontier = nextFrontier
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func queryDeleteUserMemoRefs(ctx context.Context, tx *sql.Tx, query string, args ...any) ([]deleteUserMemoRef, error) {
|
||||
rows, err := tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
memos := make([]deleteUserMemoRef, 0)
|
||||
for rows.Next() {
|
||||
var memo deleteUserMemoRef
|
||||
if err := rows.Scan(&memo.ID, &memo.UID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
memos = append(memos, memo)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return memos, nil
|
||||
}
|
||||
|
||||
func listDeleteUserAttachments(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32, memoIDs []int32) ([]*Attachment, error) {
|
||||
attachments := make([]*Attachment, 0)
|
||||
seen := make(map[int32]struct{})
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE creator_id = `+deleteUserPlaceholder(dialect, 1), []any{userID}, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
if err := appendDeleteUserAttachments(ctx, tx, `
|
||||
SELECT
|
||||
id,
|
||||
uid,
|
||||
creator_id,
|
||||
memo_id,
|
||||
storage_type,
|
||||
reference,
|
||||
payload
|
||||
FROM attachment
|
||||
WHERE memo_id IN `+clause, args, seen, &attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return attachments, nil
|
||||
}
|
||||
|
||||
func appendDeleteUserAttachments(ctx context.Context, tx *sql.Tx, query string, args []any, seen map[int32]struct{}, attachments *[]*Attachment) error {
|
||||
rows, err := tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
attachment := &Attachment{}
|
||||
var memoID sql.NullInt32
|
||||
var storageType string
|
||||
var payloadBytes []byte
|
||||
if err := rows.Scan(&attachment.ID, &attachment.UID, &attachment.CreatorID, &memoID, &storageType, &attachment.Reference, &payloadBytes); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, exists := seen[attachment.ID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[attachment.ID] = struct{}{}
|
||||
if memoID.Valid {
|
||||
attachment.MemoID = &memoID.Int32
|
||||
}
|
||||
attachment.StorageType = storepb.AttachmentStorageType(storepb.AttachmentStorageType_value[storageType])
|
||||
payload := &storepb.AttachmentPayload{}
|
||||
if len(payloadBytes) > 0 {
|
||||
if err := protojsonUnmarshaler.Unmarshal(payloadBytes, payload); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
attachment.Payload = payload
|
||||
*attachments = append(*attachments, attachment)
|
||||
}
|
||||
return rows.Err()
|
||||
}
|
||||
|
||||
func listDeleteUserSettingKeys(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) ([]storepb.UserSetting_Key, error) {
|
||||
rows, err := tx.QueryContext(ctx, `SELECT key FROM user_setting WHERE user_id = `+deleteUserPlaceholder(dialect, 1), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
keys := make([]storepb.UserSetting_Key, 0)
|
||||
for rows.Next() {
|
||||
var keyString string
|
||||
if err := rows.Scan(&keyString); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
key := storepb.UserSetting_Key(storepb.UserSetting_Key_value[keyString])
|
||||
keys = append(keys, key)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func listDeleteUserInboxIDs(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
directIDs, err := listDeleteUserDirectInboxIDs(ctx, tx, dialect, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs := append([]int32{}, directIDs...)
|
||||
if len(memoIDSet) == 0 {
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
memoIDs, err := listDeleteUserMemoReferencedInboxIDs(ctx, tx, dialect, userID, memoIDSet)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return append(inboxIDs, memoIDs...), nil
|
||||
}
|
||||
|
||||
func listDeleteUserDirectInboxIDs(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id
|
||||
FROM inbox
|
||||
WHERE sender_id = `+deleteUserPlaceholder(dialect, 1)+`
|
||||
OR receiver_id = `+deleteUserPlaceholder(dialect, 2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var inboxID int32
|
||||
if err := rows.Scan(&inboxID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func listDeleteUserMemoReferencedInboxIDs(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32, memoIDSet map[int32]struct{}) ([]int32, error) {
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
SELECT id, message
|
||||
FROM inbox
|
||||
WHERE sender_id <> `+deleteUserPlaceholder(dialect, 1)+`
|
||||
AND receiver_id <> `+deleteUserPlaceholder(dialect, 2), userID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
inboxIDs := make([]int32, 0)
|
||||
for rows.Next() {
|
||||
var (
|
||||
inboxID int32
|
||||
messageRaw []byte
|
||||
)
|
||||
if err := rows.Scan(&inboxID, &messageRaw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(messageRaw) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
message := &storepb.InboxMessage{}
|
||||
if err := protojsonUnmarshaler.Unmarshal(messageRaw, message); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if inboxMessageTouchesMemoSet(message, memoIDSet) {
|
||||
inboxIDs = append(inboxIDs, inboxID)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return inboxIDs, nil
|
||||
}
|
||||
|
||||
func inboxMessageTouchesMemoSet(message *storepb.InboxMessage, memoIDSet map[int32]struct{}) bool {
|
||||
if message == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
switch message.Type {
|
||||
case storepb.InboxMessage_MEMO_COMMENT:
|
||||
payload := message.GetMemoComment()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
case storepb.InboxMessage_MEMO_MENTION:
|
||||
payload := message.GetMemoMention()
|
||||
if payload == nil {
|
||||
return false
|
||||
}
|
||||
return memoIDInSet(payload.MemoId, memoIDSet) || memoIDInSet(payload.RelatedMemoId, memoIDSet)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func memoIDInSet(id int32, memoIDSet map[int32]struct{}) bool {
|
||||
if id == 0 {
|
||||
return false
|
||||
}
|
||||
_, exists := memoIDSet[id]
|
||||
return exists
|
||||
}
|
||||
|
||||
func deleteReactionsByContentIDsTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, contentIDs []string) error {
|
||||
for _, batch := range deleteUserBatches(contentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE content_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteAttachmentsByIDsTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, attachmentIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(attachmentIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM attachment WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteReactionsByCreatorTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM reaction WHERE creator_id = `+deleteUserPlaceholder(dialect, 1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoSharesTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32, memoIDs []int32) error {
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE creator_id = `+deleteUserPlaceholder(dialect, 1), userID); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo_share WHERE memo_id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteInboxesByIDsTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, inboxIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(inboxIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM inbox WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserIdentitiesTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_identity WHERE user_id = `+deleteUserPlaceholder(dialect, 1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserSettingsTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM user_setting WHERE user_id = `+deleteUserPlaceholder(dialect, 1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteMemoRelationsTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
memoClause, args := deleteUserInClause(dialect, 1, batch)
|
||||
relatedClause, relatedArgs := deleteUserInClause(dialect, len(args)+1, batch)
|
||||
query := `DELETE FROM memo_relation WHERE memo_id IN ` + memoClause + ` OR related_memo_id IN ` + relatedClause
|
||||
args = append(args, relatedArgs...)
|
||||
if _, err := tx.ExecContext(ctx, query, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteMemosTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, memoIDs []int32) error {
|
||||
for _, batch := range deleteUserBatches(memoIDs, deleteUserBatchSize) {
|
||||
clause, args := deleteUserInClause(dialect, 1, batch)
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM memo WHERE id IN `+clause, args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteUserRowTx(ctx context.Context, tx *sql.Tx, dialect deleteUserDialect, userID int32) error {
|
||||
_, err := tx.ExecContext(ctx, `DELETE FROM `+deleteUserTableName(dialect, "user")+` WHERE id = `+deleteUserPlaceholder(dialect, 1), userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUserTableName(dialect deleteUserDialect, table string) string {
|
||||
switch dialect {
|
||||
case deleteUserDialectMySQL:
|
||||
return "`" + table + "`"
|
||||
case deleteUserDialectPostgres:
|
||||
return `"` + table + `"`
|
||||
default:
|
||||
return table
|
||||
}
|
||||
}
|
||||
|
||||
func deleteUserPlaceholder(dialect deleteUserDialect, index int) string {
|
||||
if dialect == deleteUserDialectPostgres {
|
||||
return fmt.Sprintf("$%d", index)
|
||||
}
|
||||
return "?"
|
||||
}
|
||||
|
||||
func deleteUserInClause[T any](dialect deleteUserDialect, start int, values []T) (string, []any) {
|
||||
placeholders := make([]string, 0, len(values))
|
||||
args := make([]any, 0, len(values))
|
||||
for i, value := range values {
|
||||
placeholders = append(placeholders, deleteUserPlaceholder(dialect, start+i))
|
||||
args = append(args, value)
|
||||
}
|
||||
return "(" + strings.Join(placeholders, ", ") + ")", args
|
||||
}
|
||||
|
||||
func deleteUserBatches[T any](values []T, size int) [][]T {
|
||||
if len(values) == 0 {
|
||||
return nil
|
||||
}
|
||||
if size <= 0 {
|
||||
size = len(values)
|
||||
}
|
||||
|
||||
batches := make([][]T, 0, (len(values)+size-1)/size)
|
||||
for start := 0; start < len(values); start += size {
|
||||
end := start + size
|
||||
if end > len(values) {
|
||||
end = len(values)
|
||||
}
|
||||
batches = append(batches, values[start:end])
|
||||
}
|
||||
return batches
|
||||
}
|
||||
|
||||
func memoIDsFromRefs(memos []deleteUserMemoRef) []int32 {
|
||||
ids := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
ids = append(ids, memo.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func memoIDSetFromRefs(memos []deleteUserMemoRef) map[int32]struct{} {
|
||||
idSet := make(map[int32]struct{}, len(memos))
|
||||
for _, memo := range memos {
|
||||
idSet[memo.ID] = struct{}{}
|
||||
}
|
||||
return idSet
|
||||
}
|
||||
|
||||
func memoContentIDsFromRefs(memos []deleteUserMemoRef) []string {
|
||||
contentIDs := make([]string, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
contentIDs = append(contentIDs, "memos/"+memo.UID)
|
||||
}
|
||||
return contentIDs
|
||||
}
|
||||
|
||||
func attachmentIDsFromList(attachments []*Attachment) []int32 {
|
||||
ids := make([]int32, 0, len(attachments))
|
||||
for _, attachment := range attachments {
|
||||
if attachment == nil {
|
||||
continue
|
||||
}
|
||||
ids = append(ids, attachment.ID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue