Skip to content

Commit 95a5b3d

Browse files
committed
Fix shamir update partial failures
1 parent dc4554a commit 95a5b3d

2 files changed

Lines changed: 326 additions & 64 deletions

File tree

apis/shamir.go

Lines changed: 124 additions & 63 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ import (
88
"fmt"
99
"runtime"
1010
"strings"
11+
"sync"
1112

1213
"github.com/ProtonMail/gopenpgp/v2/crypto"
1314
"github.com/go-playground/validator/v10"
@@ -434,89 +435,111 @@ func updateShamir() {
434435

435436
var warningMessage strings.Builder
436437

437-
const shamirTableName = "shamir_email"
438-
439438
shamirEmails := make([]ShamirEmail, 0, len(ShamirPublicKeys)*len(userIDs))
440439

441440
err = func() (err error) {
442441
// concurrently compute
443-
taskChan := make(chan func(), 100)
444-
errChan := make(chan error)
445-
warningMessageChan := make(chan string, 1000)
446-
shamirEmailResultChan := make(chan []ShamirEmail, 100)
447-
defer func() {
448-
close(taskChan)
449-
close(errChan)
450-
close(warningMessageChan)
451-
close(shamirEmailResultChan)
452-
}()
442+
type updateResult struct {
443+
userID int
444+
emails []ShamirEmail
445+
warning string
446+
err error
447+
}
448+
449+
taskChan := make(chan int)
450+
resultChan := make(chan updateResult, 100)
453451

454452
// task executor
453+
var wg sync.WaitGroup
455454
for i := 0; i < runtime.NumCPU(); i++ {
455+
wg.Add(1)
456456
go func() {
457-
for task := range taskChan {
458-
task()
459-
}
460-
}()
461-
}
462-
463-
// task sender
464-
go func() {
465-
466-
// main loop
467-
for _, userID := range userIDs {
468-
469-
userID := userID
470-
// get shares
471-
shares := allShares[userID]
472-
if len(shares) < decryptThreshold {
473-
warningMessageChan <- fmt.Sprintf("user %v don't have enough shares\n", userID)
474-
continue
475-
}
476-
477-
taskChan <- func() {
457+
defer wg.Done()
458+
for userID := range taskChan {
459+
// get shares
460+
shares := allShares[userID]
461+
if len(shares) < decryptThreshold {
462+
resultChan <- updateResult{
463+
userID: userID,
464+
warning: fmt.Sprintf("user %v don't have enough shares\n", userID),
465+
}
466+
continue
467+
}
478468

479469
// decrypt email
480470
email, ok := decryptEmailWithThreshold(shares, decryptThreshold)
481471
if !ok {
482-
errChan <- fmt.Errorf("[email decrypt error] invalid shares, user_id = %d", userID)
483-
return
472+
resultChan <- updateResult{
473+
userID: userID,
474+
err: fmt.Errorf("[email decrypt error] invalid shares, user_id = %d", userID),
475+
}
476+
continue
484477
}
485478
if !utils.ValidateEmail(email) {
486479
if !utils.IsEmail(email) {
487480
// decrypt error
488-
errChan <- fmt.Errorf("[email decrypt error] invalid email, user_id = %d, email: %v", userID, email)
489-
return
490-
} else {
491-
// filter invalid emails
492-
warningMessageChan <- fmt.Sprintf("user %v don't have valid email: %v\n", userID, email)
493-
return
481+
resultChan <- updateResult{
482+
userID: userID,
483+
err: fmt.Errorf("[email decrypt error] invalid email, user_id = %d, email: %v", userID, email),
484+
}
485+
continue
494486
}
487+
488+
// filter invalid emails
489+
resultChan <- updateResult{
490+
userID: userID,
491+
warning: fmt.Sprintf("user %v don't have valid email: %v\n", userID, email),
492+
}
493+
continue
495494
}
496495

497496
// generate shamir emails
498-
var innerShamirEmails []ShamirEmail
499-
innerShamirEmails, err = GenerateShamirEmails(userID, email)
500-
if err != nil {
501-
errChan <- err
502-
return
497+
innerShamirEmails, innerErr := GenerateShamirEmails(userID, email)
498+
if innerErr != nil {
499+
resultChan <- updateResult{userID: userID, err: innerErr}
500+
continue
503501
}
504502

505-
shamirEmailResultChan <- innerShamirEmails
503+
resultChan <- updateResult{userID: userID, emails: innerShamirEmails}
506504
}
505+
}()
506+
}
507+
508+
// task sender
509+
go func() {
510+
for _, userID := range userIDs {
511+
taskChan <- userID
507512
}
513+
close(taskChan)
514+
wg.Wait()
515+
close(resultChan)
508516
}()
509517

510518
// receive task results
511519
taskCount := 0
512-
for range userIDs {
513-
select {
514-
case err = <-errChan:
515-
return err
516-
case innerWarningMessage := <-warningMessageChan:
517-
warningMessage.WriteString(innerWarningMessage)
518-
case innerShamirEmails := <-shamirEmailResultChan:
519-
shamirEmails = append(shamirEmails, innerShamirEmails...)
520+
successCount := 0
521+
skippedCount := 0
522+
for result := range resultChan {
523+
if result.err != nil {
524+
skippedCount++
525+
log.Warn().
526+
Err(result.err).
527+
Int("user_id", result.userID).
528+
Str("scope", taskScope).
529+
Msg("skip user during shamir update")
530+
warningMessage.WriteString(result.err.Error())
531+
warningMessage.WriteByte('\n')
532+
} else if result.warning != "" {
533+
skippedCount++
534+
log.Warn().
535+
Int("user_id", result.userID).
536+
Str("warning", result.warning).
537+
Str("scope", taskScope).
538+
Msg("skip user during shamir update")
539+
warningMessage.WriteString(result.warning)
540+
} else if result.emails != nil {
541+
successCount++
542+
shamirEmails = append(shamirEmails, result.emails...)
520543
}
521544
taskCount++
522545
if taskCount%1000 == 0 {
@@ -526,6 +549,30 @@ func updateShamir() {
526549
GlobalUploadShamirStatus.NowUserID = taskCount
527550
GlobalUploadShamirStatus.Unlock()
528551
}
552+
if successCount*100 < len(userIDs)*99 {
553+
err = fmt.Errorf(
554+
"shamir update success rate below 99%%: success=%d skipped=%d total=%d",
555+
successCount,
556+
skippedCount,
557+
len(userIDs),
558+
)
559+
log.Error().
560+
Err(err).
561+
Int("success_count", successCount).
562+
Int("skipped_count", skippedCount).
563+
Int("total_count", len(userIDs)).
564+
Str("scope", taskScope).
565+
Msg("abort shamir update")
566+
return err
567+
}
568+
if skippedCount > 0 {
569+
log.Warn().
570+
Int("success_count", successCount).
571+
Int("skipped_count", skippedCount).
572+
Int("total_count", len(userIDs)).
573+
Str("scope", taskScope).
574+
Msg("continue shamir update with skipped users")
575+
}
529576

530577
return DB.Session(&gorm.Session{
531578
Logger: DB.Logger.LogMode(logger.Warn),
@@ -534,15 +581,28 @@ func updateShamir() {
534581
CreateBatchSize: 1000,
535582
}).Transaction(func(tx *gorm.DB) error {
536583

537-
// delete old table
538-
if tx.Dialector.Name() == "sqlite" {
539-
//goland:noinspection SqlWithoutWhere
540-
err = tx.Exec(`DELETE FROM ` + shamirTableName).Error
541-
} else {
542-
err = tx.Exec(`TRUNCATE ` + shamirTableName).Error
584+
// Replace only successful users. Skipped users keep their existing rows
585+
// for manual cleanup or later regeneration.
586+
successfulUserIDs := make([]int, 0, successCount)
587+
successfulUserIDSet := make(map[int]struct{}, successCount)
588+
for _, shamirEmail := range shamirEmails {
589+
if _, ok := successfulUserIDSet[shamirEmail.UserID]; !ok {
590+
successfulUserIDSet[shamirEmail.UserID] = struct{}{}
591+
successfulUserIDs = append(successfulUserIDs, shamirEmail.UserID)
592+
}
543593
}
544-
if err != nil {
545-
return err
594+
595+
if len(successfulUserIDs) > 0 {
596+
for start := 0; start < len(successfulUserIDs); start += 1000 {
597+
end := min(start+1000, len(successfulUserIDs))
598+
err = tx.
599+
Where("user_id IN ?", successfulUserIDs[start:end]).
600+
Delete(&ShamirEmail{}).
601+
Error
602+
if err != nil {
603+
return err
604+
}
605+
}
546606
}
547607

548608
// insert new shamir emails
@@ -591,6 +651,7 @@ func updateShamir() {
591651

592652
subject = "shamir update failed"
593653
} else {
654+
status.FailMessage = ""
594655
subject = "shamir update success"
595656
}
596657

0 commit comments

Comments
 (0)