Skip to content

Commit ac6bde7

Browse files
authored
Merge pull request Wei-Shaw#872 from StarryKira/fix/oauth-linuxdo-invitation-required
fix: Linux.do OAuth 注册支持邀请码两步流程 (fix Wei-Shaw#836)
2 parents d2d41d6 + b43ee62 commit ac6bde7

14 files changed

Lines changed: 471 additions & 38 deletions

File tree

backend/cmd/jwtgen/main.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ func main() {
3333
}()
3434

3535
userRepo := repository.NewUserRepository(client, sqlDB)
36-
authService := service.NewAuthService(userRepo, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
36+
authService := service.NewAuthService(client, userRepo, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
3737

3838
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
3939
defer cancel()

backend/cmd/server/wire_gen.go

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

backend/internal/handler/auth_linuxdo_oauth.go

Lines changed: 50 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -211,8 +211,22 @@ func (h *AuthHandler) LinuxDoOAuthCallback(c *gin.Context) {
211211
email = linuxDoSyntheticEmail(subject)
212212
}
213213

214-
tokenPair, _, err := h.authService.LoginOrRegisterOAuthWithTokenPair(c.Request.Context(), email, username)
214+
// 传入空邀请码;如果需要邀请码,服务层返回 ErrOAuthInvitationRequired
215+
tokenPair, _, err := h.authService.LoginOrRegisterOAuthWithTokenPair(c.Request.Context(), email, username, "")
215216
if err != nil {
217+
if errors.Is(err, service.ErrOAuthInvitationRequired) {
218+
pendingToken, tokenErr := h.authService.CreatePendingOAuthToken(email, username)
219+
if tokenErr != nil {
220+
redirectOAuthError(c, frontendCallback, "login_failed", "service_error", "")
221+
return
222+
}
223+
fragment := url.Values{}
224+
fragment.Set("error", "invitation_required")
225+
fragment.Set("pending_oauth_token", pendingToken)
226+
fragment.Set("redirect", redirectTo)
227+
redirectWithFragment(c, frontendCallback, fragment)
228+
return
229+
}
216230
// 避免把内部细节泄露给客户端;给前端保留结构化原因与提示信息即可。
217231
redirectOAuthError(c, frontendCallback, "login_failed", infraerrors.Reason(err), infraerrors.Message(err))
218232
return
@@ -227,6 +241,41 @@ func (h *AuthHandler) LinuxDoOAuthCallback(c *gin.Context) {
227241
redirectWithFragment(c, frontendCallback, fragment)
228242
}
229243

244+
type completeLinuxDoOAuthRequest struct {
245+
PendingOAuthToken string `json:"pending_oauth_token" binding:"required"`
246+
InvitationCode string `json:"invitation_code" binding:"required"`
247+
}
248+
249+
// CompleteLinuxDoOAuthRegistration completes a pending OAuth registration by validating
250+
// the invitation code and creating the user account.
251+
// POST /api/v1/auth/oauth/linuxdo/complete-registration
252+
func (h *AuthHandler) CompleteLinuxDoOAuthRegistration(c *gin.Context) {
253+
var req completeLinuxDoOAuthRequest
254+
if err := c.ShouldBindJSON(&req); err != nil {
255+
c.JSON(http.StatusBadRequest, gin.H{"error": "INVALID_REQUEST", "message": err.Error()})
256+
return
257+
}
258+
259+
email, username, err := h.authService.VerifyPendingOAuthToken(req.PendingOAuthToken)
260+
if err != nil {
261+
c.JSON(http.StatusUnauthorized, gin.H{"error": "INVALID_TOKEN", "message": "invalid or expired registration token"})
262+
return
263+
}
264+
265+
tokenPair, _, err := h.authService.LoginOrRegisterOAuthWithTokenPair(c.Request.Context(), email, username, req.InvitationCode)
266+
if err != nil {
267+
response.ErrorFrom(c, err)
268+
return
269+
}
270+
271+
c.JSON(http.StatusOK, gin.H{
272+
"access_token": tokenPair.AccessToken,
273+
"refresh_token": tokenPair.RefreshToken,
274+
"expires_in": tokenPair.ExpiresIn,
275+
"token_type": "Bearer",
276+
})
277+
}
278+
230279
func (h *AuthHandler) getLinuxDoOAuthConfig(ctx context.Context) (config.LinuxDoConnectConfig, error) {
231280
if h != nil && h.settingSvc != nil {
232281
return h.settingSvc.GetLinuxDoConnectOAuthConfig(ctx)

backend/internal/server/middleware/admin_auth_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@ func TestAdminAuthJWTValidatesTokenVersion(t *testing.T) {
1919
gin.SetMode(gin.TestMode)
2020

2121
cfg := &config.Config{JWT: config.JWTConfig{Secret: "test-secret", ExpireHour: 1}}
22-
authService := service.NewAuthService(nil, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
22+
authService := service.NewAuthService(nil, nil, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
2323

2424
admin := &service.User{
2525
ID: 1,

backend/internal/server/middleware/jwt_auth_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,7 @@ func newJWTTestEnv(users map[int64]*service.User) (*gin.Engine, *service.AuthSer
4040
cfg.JWT.AccessTokenExpireMinutes = 60
4141

4242
userRepo := &stubJWTUserRepo{users: users}
43-
authSvc := service.NewAuthService(userRepo, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
43+
authSvc := service.NewAuthService(nil, userRepo, nil, nil, cfg, nil, nil, nil, nil, nil, nil)
4444
userSvc := service.NewUserService(userRepo, nil, nil)
4545
mw := NewJWTAuthMiddleware(authSvc, userSvc)
4646

backend/internal/server/routes/auth.go

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,12 @@ func RegisterAuthRoutes(
6161
}), h.Auth.ResetPassword)
6262
auth.GET("/oauth/linuxdo/start", h.Auth.LinuxDoOAuthStart)
6363
auth.GET("/oauth/linuxdo/callback", h.Auth.LinuxDoOAuthCallback)
64+
auth.POST("/oauth/linuxdo/complete-registration",
65+
rateLimiter.LimitWithOptions("oauth-linuxdo-complete", 10, time.Minute, middleware.RateLimitOptions{
66+
FailureMode: middleware.RateLimitFailClose,
67+
}),
68+
h.Auth.CompleteLinuxDoOAuthRegistration,
69+
)
6470
}
6571

6672
// 公开设置(无需认证)

backend/internal/service/auth_service.go

Lines changed: 147 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -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

2324
var (
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 认证服务
6062
type 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 创建认证服务实例
7881
func 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+
621738
func (s *AuthService) assignDefaultSubscriptions(ctx context.Context, userID int64) {
622739
if s.settingService == nil || s.defaultSubAssigner == nil || userID <= 0 {
623740
return

0 commit comments

Comments
 (0)