Skip to content
Closed
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
2 changes: 2 additions & 0 deletions controller/misc.go
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,8 @@ func GetStatus(c *gin.Context) {
"email_verification": common.EmailVerificationEnabled,
"github_oauth": common.GitHubOAuthEnabled,
"github_client_id": common.GitHubClientId,
"google_oauth": system_setting.GetGoogleSettings().Enabled,
"google_client_id": system_setting.GetGoogleSettings().ClientId,
"discord_oauth": system_setting.GetDiscordSettings().Enabled,
"discord_client_id": system_setting.GetDiscordSettings().ClientId,
"linuxdo_oauth": common.LinuxDOOAuthEnabled,
Expand Down
1 change: 1 addition & 0 deletions controller/oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -308,6 +308,7 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
provider.SetProviderUserID(user, oauthUser.ProviderUserID)
if err := tx.Model(user).Updates(map[string]interface{}{
"github_id": user.GitHubId,
"google_id": user.GoogleId,
"discord_id": user.DiscordId,
"oidc_id": user.OidcId,
"linux_do_id": user.LinuxDOId,
Expand Down
8 changes: 8 additions & 0 deletions controller/option.go
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,14 @@ func UpdateOption(c *gin.Context) {
})
return
}
case "google.enabled":
if option.Value == "true" && (system_setting.GetGoogleSettings().ClientId == "" || system_setting.GetGoogleSettings().ClientSecret == "") {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "无法启用 Google OAuth,请先填入 Google Client Id 以及 Google Client Secret!",
})
return
}
case "oidc.enabled":
if option.Value == "true" && system_setting.GetOIDCSettings().ClientId == "" {
c.JSON(http.StatusOK, gin.H{
Expand Down
14 changes: 14 additions & 0 deletions model/user.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ type User struct {
Status int `json:"status" gorm:"type:int;default:1"` // enabled, disabled
Email string `json:"email" gorm:"index" validate:"max=50"`
GitHubId string `json:"github_id" gorm:"column:github_id;index"`
GoogleId string `json:"google_id" gorm:"column:google_id;index"`
DiscordId string `json:"discord_id" gorm:"column:discord_id;index"`
OidcId string `json:"oidc_id" gorm:"column:oidc_id;index"`
WeChatId string `json:"wechat_id" gorm:"column:wechat_id;index"`
Expand Down Expand Up @@ -547,6 +548,7 @@ func (user *User) ClearBinding(bindingType string) error {
bindingColumnMap := map[string]string{
"email": "email",
"github": "github_id",
"google": "google_id",
"discord": "discord_id",
"oidc": "oidc_id",
"wechat": "wechat_id",
Expand Down Expand Up @@ -633,6 +635,14 @@ func (user *User) FillUserByGitHubId() error {
return nil
}

func (user *User) FillUserByGoogleId() error {
if user.GoogleId == "" {
return errors.New("Google id 为空!")
}
DB.Where(User{GoogleId: user.GoogleId}).First(user)
return nil
}

// UpdateGitHubId updates the user's GitHub ID (used for migration from login to numeric ID)
func (user *User) UpdateGitHubId(newGitHubId string) error {
if user.Id == 0 {
Expand Down Expand Up @@ -688,6 +698,10 @@ func IsGitHubIdAlreadyTaken(githubId string) bool {
return DB.Unscoped().Where("github_id = ?", githubId).Find(&User{}).RowsAffected == 1
}

func IsGoogleIdAlreadyTaken(googleId string) bool {
return DB.Unscoped().Where("google_id = ?", googleId).Find(&User{}).RowsAffected == 1
}

func IsDiscordIdAlreadyTaken(discordId string) bool {
return DB.Unscoped().Where("discord_id = ?", discordId).Find(&User{}).RowsAffected == 1
}
Expand Down
159 changes: 159 additions & 0 deletions oauth/google.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
package oauth

import (
"context"
"fmt"
"net/http"
"net/url"
"strings"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/i18n"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting/system_setting"
"github.com/gin-gonic/gin"
)

func init() {
Register("google", &GoogleProvider{})
}

// GoogleProvider implements OAuth for Google
type GoogleProvider struct{}

type googleOAuthResponse struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
Scope string `json:"scope"`
IDToken string `json:"id_token"`
}

type googleUser struct {
Sub string `json:"sub"`
Name string `json:"name"`
Email string `json:"email"`
}

func (p *GoogleProvider) GetName() string {
return "Google"
}

func (p *GoogleProvider) IsEnabled() bool {
return system_setting.GetGoogleSettings().Enabled
}

func (p *GoogleProvider) ExchangeToken(ctx context.Context, code string, c *gin.Context) (*OAuthToken, error) {
_ = c
if code == "" {
return nil, NewOAuthError(i18n.MsgOAuthInvalidCode, nil)
}

logger.LogDebug(ctx, "[OAuth-Google] ExchangeToken: code=%s...", code[:min(len(code), 10)])

settings := system_setting.GetGoogleSettings()
redirectUri := fmt.Sprintf("%s/oauth/google", system_setting.ServerAddress)
values := url.Values{}
values.Set("client_id", settings.ClientId)
values.Set("client_secret", settings.ClientSecret)
values.Set("code", code)
values.Set("grant_type", "authorization_code")
values.Set("redirect_uri", redirectUri)

logger.LogDebug(ctx, "[OAuth-Google] ExchangeToken: redirect_uri=%s", redirectUri)

req, err := http.NewRequestWithContext(ctx, "POST", "https://oauth2.googleapis.com/token", strings.NewReader(values.Encode()))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Accept", "application/json")

client := http.Client{Timeout: 5 * time.Second}
res, err := client.Do(req)
if err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-Google] ExchangeToken error: %s", err.Error()))
return nil, NewOAuthErrorWithRaw(i18n.MsgOAuthConnectFailed, map[string]any{"Provider": "Google"}, err.Error())
}
defer res.Body.Close()

logger.LogDebug(ctx, "[OAuth-Google] ExchangeToken response status: %d", res.StatusCode)

var googleResponse googleOAuthResponse
if err = common.DecodeJson(res.Body, &googleResponse); err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-Google] ExchangeToken decode error: %s", err.Error()))
return nil, err
}

if googleResponse.AccessToken == "" {
logger.LogError(ctx, "[OAuth-Google] ExchangeToken failed: empty access token")
return nil, NewOAuthError(i18n.MsgOAuthTokenFailed, map[string]any{"Provider": "Google"})
}

return &OAuthToken{
AccessToken: googleResponse.AccessToken,
TokenType: googleResponse.TokenType,
ExpiresIn: googleResponse.ExpiresIn,
Scope: googleResponse.Scope,
IDToken: googleResponse.IDToken,
}, nil
}

func (p *GoogleProvider) GetUserInfo(ctx context.Context, token *OAuthToken) (*OAuthUser, error) {
req, err := http.NewRequestWithContext(ctx, "GET", "https://openidconnect.googleapis.com/v1/userinfo", nil)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", "Bearer "+token.AccessToken)

client := http.Client{Timeout: 5 * time.Second}
res, err := client.Do(req)
if err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-Google] GetUserInfo error: %s", err.Error()))
return nil, NewOAuthErrorWithRaw(i18n.MsgOAuthConnectFailed, map[string]any{"Provider": "Google"}, err.Error())
}
defer res.Body.Close()

logger.LogDebug(ctx, "[OAuth-Google] GetUserInfo response status: %d", res.StatusCode)

if res.StatusCode != http.StatusOK {
return nil, NewOAuthError(i18n.MsgOAuthGetUserErr, nil)
}

var userInfo googleUser
if err = common.DecodeJson(res.Body, &userInfo); err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-Google] GetUserInfo decode error: %s", err.Error()))
return nil, err
}

if userInfo.Sub == "" {
logger.LogError(ctx, "[OAuth-Google] GetUserInfo failed: empty sub")
return nil, NewOAuthError(i18n.MsgOAuthUserInfoEmpty, map[string]any{"Provider": "Google"})
}

return &OAuthUser{
ProviderUserID: userInfo.Sub,
Username: userInfo.Email,
DisplayName: userInfo.Name,
Email: userInfo.Email,
}, nil
}

func (p *GoogleProvider) IsUserIDTaken(providerUserID string) bool {
return model.IsGoogleIdAlreadyTaken(providerUserID)
}

func (p *GoogleProvider) FillUserByProviderID(user *model.User, providerUserID string) error {
user.GoogleId = providerUserID
return user.FillUserByGoogleId()
}

func (p *GoogleProvider) SetProviderUserID(user *model.User, providerUserID string) {
user.GoogleId = providerUserID
}

func (p *GoogleProvider) GetProviderPrefix() string {
return "google_"
}
19 changes: 19 additions & 0 deletions setting/system_setting/google.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
package system_setting

import "github.com/QuantumNous/new-api/setting/config"

type GoogleSettings struct {
Enabled bool `json:"enabled"`
ClientId string `json:"client_id"`
ClientSecret string `json:"client_secret"`
}

var defaultGoogleSettings = GoogleSettings{}

func init() {
config.GlobalConfig.Register("google", &defaultGoogleSettings)
}

func GetGoogleSettings() *GoogleSettings {
return &defaultGoogleSettings
}
40 changes: 39 additions & 1 deletion web/src/components/auth/LoginForm.jsx
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ import {
getOAuthProviderIcon,
setUserData,
onGitHubOAuthClicked,
onGoogleOAuthClicked,
onDiscordOAuthClicked,
onOIDCClicked,
onLinuxDOOAuthClicked,
Expand Down Expand Up @@ -65,7 +66,7 @@ import WeChatIcon from '../common/logo/WeChatIcon';
import LinuxDoIcon from '../common/logo/LinuxDoIcon';
import TwoFAVerification from './TwoFAVerification';
import { useTranslation } from 'react-i18next';
import { SiDiscord } from 'react-icons/si';
import { SiDiscord, SiGoogle } from 'react-icons/si';

const LoginForm = () => {
let navigate = useNavigate();
Expand All @@ -92,6 +93,7 @@ const LoginForm = () => {
const [showEmailLogin, setShowEmailLogin] = useState(false);
const [wechatLoading, setWechatLoading] = useState(false);
const [githubLoading, setGithubLoading] = useState(false);
const [googleLoading, setGoogleLoading] = useState(false);
const [discordLoading, setDiscordLoading] = useState(false);
const [oidcLoading, setOidcLoading] = useState(false);
const [linuxdoLoading, setLinuxdoLoading] = useState(false);
Expand Down Expand Up @@ -135,6 +137,7 @@ const LoginForm = () => {
(status.custom_oauth_providers || []).length > 0;
const hasOAuthLoginOptions = Boolean(
status.github_oauth ||
status.google_oauth ||
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
Expand Down Expand Up @@ -337,6 +340,20 @@ const LoginForm = () => {
}
};

// 包装的Google登录点击处理
const handleGoogleClick = () => {
if ((hasUserAgreement || hasPrivacyPolicy) && !agreedToTerms) {
showInfo(t('请先阅读并同意用户协议和隐私政策'));
return;
}
setGoogleLoading(true);
try {
onGoogleOAuthClicked(status.google_client_id, { shouldLogout: true });
} finally {
setTimeout(() => setGoogleLoading(false), 3000);
}
};

// 包装的Discord登录点击处理
const handleDiscordClick = () => {
if ((hasUserAgreement || hasPrivacyPolicy) && !agreedToTerms) {
Expand Down Expand Up @@ -548,6 +565,27 @@ const LoginForm = () => {
</Button>
)}

{status.google_oauth && (
<Button
theme='outline'
className='w-full h-12 flex items-center justify-center !rounded-full border border-gray-200 hover:bg-gray-50 transition-colors'
type='tertiary'
icon={
<SiGoogle
style={{
color: '#4285F4',
width: '20px',
height: '20px',
}}
/>
}
onClick={handleGoogleClick}
loading={googleLoading}
>
<span className='ml-3'>{t('使用 Google 继续')}</span>
</Button>
)}

{status.discord_oauth && (
<Button
theme='outline'
Expand Down
Loading