Skip to content

Commit 106b20c

Browse files
committed
fix claudecode review bug
1 parent c069b3b commit 106b20c

4 files changed

Lines changed: 55 additions & 19 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: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -264,11 +264,7 @@ func (h *AuthHandler) CompleteLinuxDoOAuthRegistration(c *gin.Context) {
264264

265265
tokenPair, _, err := h.authService.LoginOrRegisterOAuthWithTokenPair(c.Request.Context(), email, username, req.InvitationCode)
266266
if err != nil {
267-
statusCode := http.StatusBadRequest
268-
c.JSON(statusCode, gin.H{
269-
"error": infraerrors.Reason(err),
270-
"message": infraerrors.Message(err),
271-
})
267+
response.ErrorFrom(c, err)
272268
return
273269
}
274270

backend/internal/service/auth_service.go

Lines changed: 52 additions & 12 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"
@@ -59,6 +60,7 @@ type JWTClaims struct {
5960

6061
// AuthService 认证服务
6162
type 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 创建认证服务实例
7981
func 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.
645677
const pendingOAuthTokenTTL = 10 * time.Minute
646678

679+
// pendingOAuthPurpose is the purpose claim value for pending OAuth registration tokens.
680+
const pendingOAuthPurpose = "pending_oauth_registration"
681+
647682
type 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

Comments
 (0)