Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 26 additions & 11 deletions controller/oauth.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package controller

import (
"errors"
"fmt"
"net/http"
"strconv"
Expand Down Expand Up @@ -104,7 +105,7 @@ func HandleOAuth(c *gin.Context) {
}

// 7. Find or create user
user, err := findOrCreateOAuthUser(c, provider, oauthUser, session)
user, created, err := findOrCreateOAuthUser(c, provider, oauthUser, session)
if err != nil {
switch err.(type) {
case *OAuthUserDeletedError:
Expand All @@ -117,6 +118,20 @@ func HandleOAuth(c *gin.Context) {
return
}

if created {
if err := createDefaultTokenForUser(user.Id, user.Username); err != nil {
switch {
case errors.Is(err, errGenerateDefaultTokenKey):
common.ApiErrorI18n(c, i18n.MsgUserDefaultTokenFailed)
case errors.Is(err, errCreateDefaultToken):
common.ApiErrorI18n(c, i18n.MsgCreateDefaultTokenErr)
default:
common.ApiError(c, err)
}
return
}
}

// 8. Check user status
if user.Status != common.UserStatusEnabled {
common.ApiErrorI18n(c, i18n.MsgOAuthUserBanned)
Expand Down Expand Up @@ -196,28 +211,28 @@ func handleOAuthBind(c *gin.Context, provider oauth.Provider) {
}

// findOrCreateOAuthUser finds existing user or creates new user
func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *oauth.OAuthUser, session sessions.Session) (*model.User, error) {
func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *oauth.OAuthUser, session sessions.Session) (*model.User, bool, error) {
user := &model.User{}

// Check if user already exists with new ID
if provider.IsUserIDTaken(oauthUser.ProviderUserID) {
err := provider.FillUserByProviderID(user, oauthUser.ProviderUserID)
if err != nil {
return nil, err
return nil, false, err
}
// Check if user has been deleted
if user.Id == 0 {
return nil, &OAuthUserDeletedError{}
return nil, false, &OAuthUserDeletedError{}
}
return user, nil
return user, false, nil
}

// Try to find user with legacy ID (for GitHub migration from login to numeric ID)
if legacyID, ok := oauthUser.Extra["legacy_id"].(string); ok && legacyID != "" {
if provider.IsUserIDTaken(legacyID) {
err := provider.FillUserByProviderID(user, legacyID)
if err != nil {
return nil, err
return nil, false, err
}
if user.Id != 0 {
// Found user with legacy ID, migrate to new ID
Expand All @@ -227,14 +242,14 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
common.SysError(fmt.Sprintf("[OAuth] Failed to migrate user %d: %s", user.Id, err.Error()))
// Continue with login even if migration fails
}
return user, nil
return user, false, nil
}
}
}

// User doesn't exist, create new user if registration is enabled
if !common.RegisterEnabled {
return nil, &OAuthRegistrationDisabledError{}
return nil, false, &OAuthRegistrationDisabledError{}
}

// Set up new user
Expand Down Expand Up @@ -291,7 +306,7 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
return nil
})
if err != nil {
return nil, err
return nil, false, err
}

// Perform post-transaction tasks (logs, sidebar config, inviter rewards)
Expand Down Expand Up @@ -320,14 +335,14 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
return nil
})
if err != nil {
return nil, err
return nil, false, err
}

// Perform post-transaction tasks
user.FinalizeOAuthUserCreation(inviterId)
}

return user, nil
return user, true, nil
}

// Error types for OAuth
Expand Down
67 changes: 43 additions & 24 deletions controller/user.go
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,42 @@ func setupLogin(user *model.User, c *gin.Context) {
})
}

var (
errGenerateDefaultTokenKey = errors.New("failed to generate default token key")
errCreateDefaultToken = errors.New("failed to create default token")
)

func createDefaultTokenForUser(userId int, username string) error {
if !constant.GenerateDefaultToken {
return nil
}
key, err := common.GenerateKey()
if err != nil {
common.SysLog("failed to generate token key: " + err.Error())
return errGenerateDefaultTokenKey
}
now := common.GetTimestamp()
token := model.Token{
UserId: userId,
Name: username + "的初始令牌",
Key: key,
CreatedTime: now,
AccessedTime: now,
ExpiredTime: -1,
RemainQuota: 500000,
UnlimitedQuota: true,
ModelLimitsEnabled: false,
}
if setting.DefaultUseAutoGroup {
token.Group = "auto"
}
if err := token.Insert(); err != nil {
common.SysLog(fmt.Sprintf("failed to insert default token for user %d: %s", userId, err.Error()))
return fmt.Errorf("insert default token: %w", errCreateDefaultToken)
}
return nil
}

func Logout(c *gin.Context) {
session := sessions.Default(c)
session.Clear()
Expand Down Expand Up @@ -195,33 +231,16 @@ func Register(c *gin.Context) {
common.ApiErrorI18n(c, i18n.MsgUserRegisterFailed)
return
}
// 生成默认令牌
if constant.GenerateDefaultToken {
key, err := common.GenerateKey()
if err != nil {
if err := createDefaultTokenForUser(insertedUser.Id, cleanUser.Username); err != nil {
switch {
case errors.Is(err, errGenerateDefaultTokenKey):
common.ApiErrorI18n(c, i18n.MsgUserDefaultTokenFailed)
common.SysLog("failed to generate token key: " + err.Error())
return
}
// 生成默认令牌
token := model.Token{
UserId: insertedUser.Id, // 使用插入后的用户ID
Name: cleanUser.Username + "的初始令牌",
Key: key,
CreatedTime: common.GetTimestamp(),
AccessedTime: common.GetTimestamp(),
ExpiredTime: -1, // 永不过期
RemainQuota: 500000, // 示例额度
UnlimitedQuota: true,
ModelLimitsEnabled: false,
}
if setting.DefaultUseAutoGroup {
token.Group = "auto"
}
if err := token.Insert(); err != nil {
case errors.Is(err, errCreateDefaultToken):
common.ApiErrorI18n(c, i18n.MsgCreateDefaultTokenErr)
return
default:
common.ApiError(c, err)
}
return
}

c.JSON(http.StatusOK, gin.H{
Expand Down
12 changes: 12 additions & 0 deletions controller/wechat.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/i18n"
"github.com/QuantumNous/new-api/model"

"github.com/gin-contrib/sessions"
Expand Down Expand Up @@ -103,6 +104,17 @@ func WeChatAuth(c *gin.Context) {
})
return
}
if err := createDefaultTokenForUser(user.Id, user.Username); err != nil {
switch {
case errors.Is(err, errGenerateDefaultTokenKey):
common.ApiErrorI18n(c, i18n.MsgUserDefaultTokenFailed)
case errors.Is(err, errCreateDefaultToken):
common.ApiErrorI18n(c, i18n.MsgCreateDefaultTokenErr)
default:
common.ApiError(c, err)
}
return
}
} else {
c.JSON(http.StatusOK, gin.H{
"success": false,
Expand Down