@@ -12,6 +12,7 @@ import (
1212 "strings"
1313 "time"
1414
15+ dbent "github.com/Wei-Shaw/sub2api/ent"
1516 "github.com/Wei-Shaw/sub2api/internal/config"
1617 infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
1718 "github.com/Wei-Shaw/sub2api/internal/pkg/logger"
@@ -59,6 +60,7 @@ type JWTClaims struct {
5960
6061// AuthService 认证服务
6162type AuthService struct {
63+ entClient * dbent.Client
6264 userRepo UserRepository
6365 redeemRepo RedeemCodeRepository
6466 refreshTokenCache RefreshTokenCache
@@ -77,6 +79,7 @@ type DefaultSubscriptionAssigner interface {
7779
7880// NewAuthService 创建认证服务实例
7981func NewAuthService (
82+ entClient * dbent.Client ,
8083 userRepo UserRepository ,
8184 redeemRepo RedeemCodeRepository ,
8285 refreshTokenCache RefreshTokenCache ,
@@ -89,6 +92,7 @@ func NewAuthService(
8992 defaultSubAssigner DefaultSubscriptionAssigner ,
9093) * AuthService {
9194 return & AuthService {
95+ entClient : entClient ,
9296 userRepo : userRepo ,
9397 redeemRepo : redeemRepo ,
9498 refreshTokenCache : refreshTokenCache ,
@@ -597,24 +601,52 @@ func (s *AuthService) LoginOrRegisterOAuthWithTokenPair(ctx context.Context, ema
597601 Status : StatusActive ,
598602 }
599603
600- if err := s .userRepo .Create (ctx , newUser ); err != nil {
601- if errors .Is (err , ErrEmailExists ) {
602- user , err = s .userRepo .GetByEmail (ctx , email )
603- if err != nil {
604- logger .LegacyPrintf ("service.auth" , "[Auth] Database error getting user after conflict: %v" , err )
604+ if s .entClient != nil && invitationRedeemCode != nil {
605+ tx , err := s .entClient .Tx (ctx )
606+ if err != nil {
607+ logger .LegacyPrintf ("service.auth" , "[Auth] Failed to begin transaction for oauth registration: %v" , err )
608+ return nil , nil , ErrServiceUnavailable
609+ }
610+ defer func () { _ = tx .Rollback () }()
611+ txCtx := dbent .NewTxContext (ctx , tx )
612+
613+ if err := s .userRepo .Create (txCtx , newUser ); err != nil {
614+ if errors .Is (err , ErrEmailExists ) {
615+ user , err = s .userRepo .GetByEmail (ctx , email )
616+ if err != nil {
617+ logger .LegacyPrintf ("service.auth" , "[Auth] Database error getting user after conflict: %v" , err )
618+ return nil , nil , ErrServiceUnavailable
619+ }
620+ } else {
621+ logger .LegacyPrintf ("service.auth" , "[Auth] Database error creating oauth user: %v" , err )
605622 return nil , nil , ErrServiceUnavailable
606623 }
607624 } else {
608- logger .LegacyPrintf ("service.auth" , "[Auth] Database error creating oauth user: %v" , err )
609- return nil , nil , ErrServiceUnavailable
625+ if err := s .redeemRepo .Use (txCtx , invitationRedeemCode .ID , newUser .ID ); err != nil {
626+ return nil , nil , ErrInvitationCodeInvalid
627+ }
628+ if err := tx .Commit (); err != nil {
629+ logger .LegacyPrintf ("service.auth" , "[Auth] Failed to commit oauth registration transaction: %v" , err )
630+ return nil , nil , ErrServiceUnavailable
631+ }
632+ user = newUser
633+ s .assignDefaultSubscriptions (ctx , user .ID )
610634 }
611635 } else {
612- user = newUser
613- s .assignDefaultSubscriptions (ctx , user .ID )
614- if invitationRedeemCode != nil {
615- if err := s .redeemRepo .Use (ctx , invitationRedeemCode .ID , user .ID ); err != nil {
616- logger .LegacyPrintf ("service.auth" , "[Auth] Failed to mark invitation code as used for oauth user %d: %v" , user .ID , err )
636+ if err := s .userRepo .Create (ctx , newUser ); err != nil {
637+ if errors .Is (err , ErrEmailExists ) {
638+ user , err = s .userRepo .GetByEmail (ctx , email )
639+ if err != nil {
640+ logger .LegacyPrintf ("service.auth" , "[Auth] Database error getting user after conflict: %v" , err )
641+ return nil , nil , ErrServiceUnavailable
642+ }
643+ } else {
644+ logger .LegacyPrintf ("service.auth" , "[Auth] Database error creating oauth user: %v" , err )
645+ return nil , nil , ErrServiceUnavailable
617646 }
647+ } else {
648+ user = newUser
649+ s .assignDefaultSubscriptions (ctx , user .ID )
618650 }
619651 }
620652 } else {
@@ -644,9 +676,13 @@ func (s *AuthService) LoginOrRegisterOAuthWithTokenPair(ctx context.Context, ema
644676// pendingOAuthTokenTTL is the validity period for pending OAuth tokens.
645677const pendingOAuthTokenTTL = 10 * time .Minute
646678
679+ // pendingOAuthPurpose is the purpose claim value for pending OAuth registration tokens.
680+ const pendingOAuthPurpose = "pending_oauth_registration"
681+
647682type pendingOAuthClaims struct {
648683 Email string `json:"email"`
649684 Username string `json:"username"`
685+ Purpose string `json:"purpose"`
650686 jwt.RegisteredClaims
651687}
652688
@@ -657,6 +693,7 @@ func (s *AuthService) CreatePendingOAuthToken(email, username string) (string, e
657693 claims := & pendingOAuthClaims {
658694 Email : email ,
659695 Username : username ,
696+ Purpose : pendingOAuthPurpose ,
660697 RegisteredClaims : jwt.RegisteredClaims {
661698 ExpiresAt : jwt .NewNumericDate (now .Add (pendingOAuthTokenTTL )),
662699 IssuedAt : jwt .NewNumericDate (now ),
@@ -687,6 +724,9 @@ func (s *AuthService) VerifyPendingOAuthToken(tokenStr string) (email, username
687724 if ! ok || ! token .Valid {
688725 return "" , "" , ErrInvalidToken
689726 }
727+ if claims .Purpose != pendingOAuthPurpose {
728+ return "" , "" , ErrInvalidToken
729+ }
690730 return claims .Email , claims .Username , nil
691731}
692732
0 commit comments