@@ -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"
@@ -21,24 +22,25 @@ import (
2122)
2223
2324var (
24- ErrInvalidCredentials = infraerrors .Unauthorized ("INVALID_CREDENTIALS" , "invalid email or password" )
25- ErrUserNotActive = infraerrors .Forbidden ("USER_NOT_ACTIVE" , "user is not active" )
26- ErrEmailExists = infraerrors .Conflict ("EMAIL_EXISTS" , "email already exists" )
27- ErrEmailReserved = infraerrors .BadRequest ("EMAIL_RESERVED" , "email is reserved" )
28- ErrInvalidToken = infraerrors .Unauthorized ("INVALID_TOKEN" , "invalid token" )
29- ErrTokenExpired = infraerrors .Unauthorized ("TOKEN_EXPIRED" , "token has expired" )
30- ErrAccessTokenExpired = infraerrors .Unauthorized ("ACCESS_TOKEN_EXPIRED" , "access token has expired" )
31- ErrTokenTooLarge = infraerrors .BadRequest ("TOKEN_TOO_LARGE" , "token too large" )
32- ErrTokenRevoked = infraerrors .Unauthorized ("TOKEN_REVOKED" , "token has been revoked" )
33- ErrRefreshTokenInvalid = infraerrors .Unauthorized ("REFRESH_TOKEN_INVALID" , "invalid refresh token" )
34- ErrRefreshTokenExpired = infraerrors .Unauthorized ("REFRESH_TOKEN_EXPIRED" , "refresh token has expired" )
35- ErrRefreshTokenReused = infraerrors .Unauthorized ("REFRESH_TOKEN_REUSED" , "refresh token has been reused" )
36- ErrEmailVerifyRequired = infraerrors .BadRequest ("EMAIL_VERIFY_REQUIRED" , "email verification is required" )
37- ErrEmailSuffixNotAllowed = infraerrors .BadRequest ("EMAIL_SUFFIX_NOT_ALLOWED" , "email suffix is not allowed" )
38- ErrRegDisabled = infraerrors .Forbidden ("REGISTRATION_DISABLED" , "registration is currently disabled" )
39- ErrServiceUnavailable = infraerrors .ServiceUnavailable ("SERVICE_UNAVAILABLE" , "service temporarily unavailable" )
40- ErrInvitationCodeRequired = infraerrors .BadRequest ("INVITATION_CODE_REQUIRED" , "invitation code is required" )
41- ErrInvitationCodeInvalid = infraerrors .BadRequest ("INVITATION_CODE_INVALID" , "invalid or used invitation code" )
25+ ErrInvalidCredentials = infraerrors .Unauthorized ("INVALID_CREDENTIALS" , "invalid email or password" )
26+ ErrUserNotActive = infraerrors .Forbidden ("USER_NOT_ACTIVE" , "user is not active" )
27+ ErrEmailExists = infraerrors .Conflict ("EMAIL_EXISTS" , "email already exists" )
28+ ErrEmailReserved = infraerrors .BadRequest ("EMAIL_RESERVED" , "email is reserved" )
29+ ErrInvalidToken = infraerrors .Unauthorized ("INVALID_TOKEN" , "invalid token" )
30+ ErrTokenExpired = infraerrors .Unauthorized ("TOKEN_EXPIRED" , "token has expired" )
31+ ErrAccessTokenExpired = infraerrors .Unauthorized ("ACCESS_TOKEN_EXPIRED" , "access token has expired" )
32+ ErrTokenTooLarge = infraerrors .BadRequest ("TOKEN_TOO_LARGE" , "token too large" )
33+ ErrTokenRevoked = infraerrors .Unauthorized ("TOKEN_REVOKED" , "token has been revoked" )
34+ ErrRefreshTokenInvalid = infraerrors .Unauthorized ("REFRESH_TOKEN_INVALID" , "invalid refresh token" )
35+ ErrRefreshTokenExpired = infraerrors .Unauthorized ("REFRESH_TOKEN_EXPIRED" , "refresh token has expired" )
36+ ErrRefreshTokenReused = infraerrors .Unauthorized ("REFRESH_TOKEN_REUSED" , "refresh token has been reused" )
37+ ErrEmailVerifyRequired = infraerrors .BadRequest ("EMAIL_VERIFY_REQUIRED" , "email verification is required" )
38+ ErrEmailSuffixNotAllowed = infraerrors .BadRequest ("EMAIL_SUFFIX_NOT_ALLOWED" , "email suffix is not allowed" )
39+ ErrRegDisabled = infraerrors .Forbidden ("REGISTRATION_DISABLED" , "registration is currently disabled" )
40+ ErrServiceUnavailable = infraerrors .ServiceUnavailable ("SERVICE_UNAVAILABLE" , "service temporarily unavailable" )
41+ ErrInvitationCodeRequired = infraerrors .BadRequest ("INVITATION_CODE_REQUIRED" , "invitation code is required" )
42+ ErrInvitationCodeInvalid = infraerrors .BadRequest ("INVITATION_CODE_INVALID" , "invalid or used invitation code" )
43+ ErrOAuthInvitationRequired = infraerrors .Forbidden ("OAUTH_INVITATION_REQUIRED" , "invitation code required to complete oauth registration" )
4244)
4345
4446// maxTokenLength 限制 token 大小,避免超长 header 触发解析时的异常内存分配。
@@ -58,6 +60,7 @@ type JWTClaims struct {
5860
5961// AuthService 认证服务
6062type AuthService struct {
63+ entClient * dbent.Client
6164 userRepo UserRepository
6265 redeemRepo RedeemCodeRepository
6366 refreshTokenCache RefreshTokenCache
@@ -76,6 +79,7 @@ type DefaultSubscriptionAssigner interface {
7679
7780// NewAuthService 创建认证服务实例
7881func NewAuthService (
82+ entClient * dbent.Client ,
7983 userRepo UserRepository ,
8084 redeemRepo RedeemCodeRepository ,
8185 refreshTokenCache RefreshTokenCache ,
@@ -88,6 +92,7 @@ func NewAuthService(
8892 defaultSubAssigner DefaultSubscriptionAssigner ,
8993) * AuthService {
9094 return & AuthService {
95+ entClient : entClient ,
9196 userRepo : userRepo ,
9297 redeemRepo : redeemRepo ,
9398 refreshTokenCache : refreshTokenCache ,
@@ -523,9 +528,10 @@ func (s *AuthService) LoginOrRegisterOAuth(ctx context.Context, email, username
523528 return token , user , nil
524529}
525530
526- // LoginOrRegisterOAuthWithTokenPair 用于第三方 OAuth/SSO 登录,返回完整的 TokenPair
527- // 与 LoginOrRegisterOAuth 功能相同,但返回 TokenPair 而非单个 token
528- func (s * AuthService ) LoginOrRegisterOAuthWithTokenPair (ctx context.Context , email , username string ) (* TokenPair , * User , error ) {
531+ // LoginOrRegisterOAuthWithTokenPair 用于第三方 OAuth/SSO 登录,返回完整的 TokenPair。
532+ // 与 LoginOrRegisterOAuth 功能相同,但返回 TokenPair 而非单个 token。
533+ // invitationCode 仅在邀请码注册模式下新用户注册时使用;已有账号登录时忽略。
534+ func (s * AuthService ) LoginOrRegisterOAuthWithTokenPair (ctx context.Context , email , username , invitationCode string ) (* TokenPair , * User , error ) {
529535 // 检查 refreshTokenCache 是否可用
530536 if s .refreshTokenCache == nil {
531537 return nil , nil , errors .New ("refresh token cache not configured" )
@@ -552,6 +558,22 @@ func (s *AuthService) LoginOrRegisterOAuthWithTokenPair(ctx context.Context, ema
552558 return nil , nil , ErrRegDisabled
553559 }
554560
561+ // 检查是否需要邀请码
562+ var invitationRedeemCode * RedeemCode
563+ if s .settingService != nil && s .settingService .IsInvitationCodeEnabled (ctx ) {
564+ if invitationCode == "" {
565+ return nil , nil , ErrOAuthInvitationRequired
566+ }
567+ redeemCode , err := s .redeemRepo .GetByCode (ctx , invitationCode )
568+ if err != nil {
569+ return nil , nil , ErrInvitationCodeInvalid
570+ }
571+ if redeemCode .Type != RedeemTypeInvitation || redeemCode .Status != StatusUnused {
572+ return nil , nil , ErrInvitationCodeInvalid
573+ }
574+ invitationRedeemCode = redeemCode
575+ }
576+
555577 randomPassword , err := randomHexString (32 )
556578 if err != nil {
557579 logger .LegacyPrintf ("service.auth" , "[Auth] Failed to generate random password for oauth signup: %v" , err )
@@ -579,20 +601,58 @@ func (s *AuthService) LoginOrRegisterOAuthWithTokenPair(ctx context.Context, ema
579601 Status : StatusActive ,
580602 }
581603
582- if err := s .userRepo .Create (ctx , newUser ); err != nil {
583- if errors .Is (err , ErrEmailExists ) {
584- user , err = s .userRepo .GetByEmail (ctx , email )
585- if err != nil {
586- 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 )
587622 return nil , nil , ErrServiceUnavailable
588623 }
589624 } else {
590- logger .LegacyPrintf ("service.auth" , "[Auth] Database error creating oauth user: %v" , err )
591- 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 )
592634 }
593635 } else {
594- user = newUser
595- s .assignDefaultSubscriptions (ctx , user .ID )
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
646+ }
647+ } else {
648+ user = newUser
649+ s .assignDefaultSubscriptions (ctx , user .ID )
650+ if invitationRedeemCode != nil {
651+ if err := s .redeemRepo .Use (ctx , invitationRedeemCode .ID , user .ID ); err != nil {
652+ return nil , nil , ErrInvitationCodeInvalid
653+ }
654+ }
655+ }
596656 }
597657 } else {
598658 logger .LegacyPrintf ("service.auth" , "[Auth] Database error during oauth login: %v" , err )
@@ -618,6 +678,63 @@ func (s *AuthService) LoginOrRegisterOAuthWithTokenPair(ctx context.Context, ema
618678 return tokenPair , user , nil
619679}
620680
681+ // pendingOAuthTokenTTL is the validity period for pending OAuth tokens.
682+ const pendingOAuthTokenTTL = 10 * time .Minute
683+
684+ // pendingOAuthPurpose is the purpose claim value for pending OAuth registration tokens.
685+ const pendingOAuthPurpose = "pending_oauth_registration"
686+
687+ type pendingOAuthClaims struct {
688+ Email string `json:"email"`
689+ Username string `json:"username"`
690+ Purpose string `json:"purpose"`
691+ jwt.RegisteredClaims
692+ }
693+
694+ // CreatePendingOAuthToken generates a short-lived JWT that carries the OAuth identity
695+ // while waiting for the user to supply an invitation code.
696+ func (s * AuthService ) CreatePendingOAuthToken (email , username string ) (string , error ) {
697+ now := time .Now ()
698+ claims := & pendingOAuthClaims {
699+ Email : email ,
700+ Username : username ,
701+ Purpose : pendingOAuthPurpose ,
702+ RegisteredClaims : jwt.RegisteredClaims {
703+ ExpiresAt : jwt .NewNumericDate (now .Add (pendingOAuthTokenTTL )),
704+ IssuedAt : jwt .NewNumericDate (now ),
705+ NotBefore : jwt .NewNumericDate (now ),
706+ },
707+ }
708+ token := jwt .NewWithClaims (jwt .SigningMethodHS256 , claims )
709+ return token .SignedString ([]byte (s .cfg .JWT .Secret ))
710+ }
711+
712+ // VerifyPendingOAuthToken validates a pending OAuth token and returns the embedded identity.
713+ // Returns ErrInvalidToken when the token is invalid or expired.
714+ func (s * AuthService ) VerifyPendingOAuthToken (tokenStr string ) (email , username string , err error ) {
715+ if len (tokenStr ) > maxTokenLength {
716+ return "" , "" , ErrInvalidToken
717+ }
718+ parser := jwt .NewParser (jwt .WithValidMethods ([]string {jwt .SigningMethodHS256 .Name }))
719+ token , parseErr := parser .ParseWithClaims (tokenStr , & pendingOAuthClaims {}, func (t * jwt.Token ) (any , error ) {
720+ if _ , ok := t .Method .(* jwt.SigningMethodHMAC ); ! ok {
721+ return nil , fmt .Errorf ("unexpected signing method: %v" , t .Header ["alg" ])
722+ }
723+ return []byte (s .cfg .JWT .Secret ), nil
724+ })
725+ if parseErr != nil {
726+ return "" , "" , ErrInvalidToken
727+ }
728+ claims , ok := token .Claims .(* pendingOAuthClaims )
729+ if ! ok || ! token .Valid {
730+ return "" , "" , ErrInvalidToken
731+ }
732+ if claims .Purpose != pendingOAuthPurpose {
733+ return "" , "" , ErrInvalidToken
734+ }
735+ return claims .Email , claims .Username , nil
736+ }
737+
621738func (s * AuthService ) assignDefaultSubscriptions (ctx context.Context , userID int64 ) {
622739 if s .settingService == nil || s .defaultSubAssigner == nil || userID <= 0 {
623740 return
0 commit comments