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