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
8 changes: 8 additions & 0 deletions controller/custom_oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ type CustomOAuthProviderResponse struct {
EmailField string `json:"email_field"`
WellKnown string `json:"well_known"`
AuthStyle int `json:"auth_style"`
AutoLinkPolicy string `json:"auto_link_policy"`
AccessPolicy string `json:"access_policy"`
AccessDeniedMessage string `json:"access_denied_message"`
}
Expand Down Expand Up @@ -64,6 +65,7 @@ func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthPro
EmailField: p.EmailField,
WellKnown: p.WellKnown,
AuthStyle: p.AuthStyle,
AutoLinkPolicy: p.AutoLinkPolicy,
AccessPolicy: p.AccessPolicy,
AccessDeniedMessage: p.AccessDeniedMessage,
}
Expand Down Expand Up @@ -129,6 +131,7 @@ type CreateCustomOAuthProviderRequest struct {
EmailField string `json:"email_field"`
WellKnown string `json:"well_known"`
AuthStyle int `json:"auth_style"`
AutoLinkPolicy string `json:"auto_link_policy"`
AccessPolicy string `json:"access_policy"`
AccessDeniedMessage string `json:"access_denied_message"`
}
Expand Down Expand Up @@ -247,6 +250,7 @@ func CreateCustomOAuthProvider(c *gin.Context) {
EmailField: req.EmailField,
WellKnown: req.WellKnown,
AuthStyle: req.AuthStyle,
AutoLinkPolicy: req.AutoLinkPolicy,
AccessPolicy: req.AccessPolicy,
AccessDeniedMessage: req.AccessDeniedMessage,
}
Expand Down Expand Up @@ -284,6 +288,7 @@ type UpdateCustomOAuthProviderRequest struct {
EmailField string `json:"email_field"`
WellKnown *string `json:"well_known"` // Optional: if nil, keep existing
AuthStyle *int `json:"auth_style"` // Optional: if nil, keep existing
AutoLinkPolicy *string `json:"auto_link_policy"` // Optional: if nil, keep existing
AccessPolicy *string `json:"access_policy"` // Optional: if nil, keep existing
AccessDeniedMessage *string `json:"access_denied_message"` // Optional: if nil, keep existing
}
Expand Down Expand Up @@ -374,6 +379,9 @@ func UpdateCustomOAuthProvider(c *gin.Context) {
if req.AuthStyle != nil {
provider.AuthStyle = *req.AuthStyle
}
if req.AutoLinkPolicy != nil {
provider.AutoLinkPolicy = *req.AutoLinkPolicy
}
if req.AccessPolicy != nil {
provider.AccessPolicy = *req.AccessPolicy
}
Expand Down
54 changes: 54 additions & 0 deletions controller/oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -325,6 +325,17 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
}
}

// Custom providers may explicitly link to an existing local user before creating a new one.
if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok {
linkedUser, linked, err := tryAutoLinkCustomOAuthUser(genericProvider, oauthUser)
if err != nil {
return nil, err
}
if linked {
return linkedUser, nil
}
}

// User doesn't exist, create new user if registration is enabled
if !common.RegisterEnabled {
return nil, &OAuthRegistrationDisabledError{}
Expand Down Expand Up @@ -428,6 +439,49 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
return user, nil
}

func tryAutoLinkCustomOAuthUser(provider *oauth.GenericOAuthProvider, oauthUser *oauth.OAuthUser) (*model.User, bool, error) {
policy := provider.GetAutoLinkPolicy()
if policy == "none" {
return nil, false, nil
}

user := &model.User{}
var err error
switch policy {
case "email_verified":
if oauthUser.Email == "" {
return nil, false, nil
}
if verified, ok := oauthUser.Extra["email_verified"].(bool); !ok || !verified {
return nil, false, nil
}
err = model.DB.Where("email = ?", oauthUser.Email).First(user).Error
case "username":
if oauthUser.Username == "" {
return nil, false, nil
}
err = model.DB.Where("username = ?", oauthUser.Username).First(user).Error
default:
return nil, false, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
if err != nil {
return nil, false, err
}
if user.Id == 0 {
return nil, false, nil
}
if user.Status != common.UserStatusEnabled {
return user, true, nil
}
if err := model.UpdateUserOAuthBinding(user.Id, provider.GetProviderId(), oauthUser.ProviderUserID); err != nil {
return nil, false, err
}
return user, true, nil
}

// Error types for OAuth
type OAuthUserDeletedError struct{}

Expand Down
9 changes: 9 additions & 0 deletions model/custom_oauth_provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@ type CustomOAuthProvider struct {
// Advanced options
WellKnown string `json:"well_known" gorm:"type:varchar(512)"` // OIDC discovery endpoint (optional)
AuthStyle int `json:"auth_style" gorm:"default:0"` // 0=auto, 1=params, 2=header (Basic Auth)
AutoLinkPolicy string `json:"auto_link_policy" gorm:"type:varchar(32);default:'none'"`
AccessPolicy string `json:"access_policy" gorm:"type:text"` // JSON policy for access control based on user info
AccessDeniedMessage string `json:"access_denied_message" gorm:"type:varchar(512)"` // Custom error message template when access is denied

Expand Down Expand Up @@ -191,6 +192,14 @@ func validateCustomOAuthProvider(provider *CustomOAuthProvider) error {
if provider.Scopes == "" {
provider.Scopes = "openid profile email"
}
if provider.AutoLinkPolicy == "" {
provider.AutoLinkPolicy = "none"
}
switch provider.AutoLinkPolicy {
case "none", "email_verified", "username":
default:
return errors.New("invalid auto link policy")
}
if strings.TrimSpace(provider.AccessPolicy) != "" {
var policy accessPolicyPayload
if err := common.UnmarshalJsonStr(provider.AccessPolicy, &policy); err != nil {
Expand Down
3 changes: 3 additions & 0 deletions model/user.go
Original file line number Diff line number Diff line change
Expand Up @@ -517,6 +517,9 @@ func DeleteUserById(id int) (err error) {
if id == 0 {
return errors.New("id 为空!")
}
if err = deleteUserOAuthBindingsByUserId(DB, id); err != nil {
return err
}
user := User{Id: id}
return user.Delete()
}
Expand Down
51 changes: 51 additions & 0 deletions model/user_oauth_binding_delete_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
package model

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/stretchr/testify/require"
)

func insertUserWithCustomOAuthBinding(t *testing.T, userId int, providerId int) {
t.Helper()
require.NoError(t, DB.Create(&User{
Id: userId,
Username: "custom_oauth_deleted_user",
Status: common.UserStatusEnabled,
}).Error)
require.NoError(t, CreateUserOAuthBinding(&UserOAuthBinding{
UserId: userId,
ProviderId: providerId,
ProviderUserId: "provider-user-id",
}))
}

func countCustomOAuthBindingsForUser(t *testing.T, userId int) int64 {
t.Helper()
var count int64
require.NoError(t, DB.Model(&UserOAuthBinding{}).Where("user_id = ?", userId).Count(&count).Error)
return count
}

func TestDeleteUserById_RemovesCustomOAuthBindings(t *testing.T) {
truncateTables(t)

insertUserWithCustomOAuthBinding(t, 201, 301)
require.Equal(t, int64(1), countCustomOAuthBindingsForUser(t, 201))

require.NoError(t, DeleteUserById(201))

require.Equal(t, int64(0), countCustomOAuthBindingsForUser(t, 201))
}

func TestHardDeleteUserById_RemovesCustomOAuthBindings(t *testing.T) {
truncateTables(t)

insertUserWithCustomOAuthBinding(t, 202, 302)
require.Equal(t, int64(1), countCustomOAuthBindingsForUser(t, 202))

require.NoError(t, HardDeleteUserById(202))

require.Equal(t, int64(0), countCustomOAuthBindingsForUser(t, 202))
}
17 changes: 14 additions & 3 deletions oauth/generic.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@ import (
"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"
"github.com/samber/lo"
"github.com/tidwall/gjson"
Expand Down Expand Up @@ -94,7 +93,7 @@ func (p *GenericOAuthProvider) ExchangeToken(ctx context.Context, code string, c

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

redirectUri := fmt.Sprintf("%s/oauth/%s", system_setting.ServerAddress, p.config.Slug)
redirectUri := BuildOAuthRedirectURI(c, p.config.Slug)
values := url.Values{}
values.Set("grant_type", "authorization_code")
values.Set("code", code)
Expand Down Expand Up @@ -243,6 +242,10 @@ func (p *GenericOAuthProvider) GetUserInfo(ctx context.Context, token *OAuthToke
username := gjson.Get(bodyStr, p.config.UsernameField).String()
displayName := gjson.Get(bodyStr, p.config.DisplayNameField).String()
email := gjson.Get(bodyStr, p.config.EmailField).String()
emailVerified := false
if emailVerifiedValue := gjson.Get(bodyStr, "email_verified"); emailVerifiedValue.Exists() {
emailVerified = emailVerifiedValue.Bool()
}

// If user ID field returns a number, convert it
if userId == "" {
Expand Down Expand Up @@ -285,7 +288,8 @@ func (p *GenericOAuthProvider) GetUserInfo(ctx context.Context, token *OAuthToke
DisplayName: displayName,
Email: email,
Extra: map[string]any{
"provider": p.config.Slug,
"provider": p.config.Slug,
"email_verified": emailVerified,
},
}, nil
}
Expand Down Expand Up @@ -322,6 +326,13 @@ func (p *GenericOAuthProvider) GetProviderId() int {
return p.config.Id
}

func (p *GenericOAuthProvider) GetAutoLinkPolicy() string {
if p.config.AutoLinkPolicy == "" {
return "none"
}
return p.config.AutoLinkPolicy
}

func normalizeAuthorizationTokenType(tokenType string) string {
tokenType = strings.TrimSpace(tokenType)
if tokenType == "" || strings.EqualFold(tokenType, "Bearer") {
Expand Down
26 changes: 21 additions & 5 deletions oauth/oidc.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,14 @@ package oauth

import (
"context"
"encoding/json"
"fmt"
"io"
"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"
Expand All @@ -30,6 +31,8 @@ type oidcOAuthResponse struct {
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
Scope string `json:"scope"`
Error string `json:"error"`
ErrorDesc string `json:"error_description"`
}

type oidcUser struct {
Expand All @@ -56,7 +59,7 @@ func (p *OIDCProvider) ExchangeToken(ctx context.Context, code string, c *gin.Co
logger.LogDebug(ctx, "[OAuth-OIDC] ExchangeToken: code=%s...", code[:min(len(code), 10)])

settings := system_setting.GetOIDCSettings()
redirectUri := fmt.Sprintf("%s/oauth/oidc", system_setting.ServerAddress)
redirectUri := BuildOAuthRedirectURI(c, "oidc")
values := url.Values{}
values.Set("client_id", settings.ClientId)
values.Set("client_secret", settings.ClientSecret)
Expand Down Expand Up @@ -85,16 +88,29 @@ func (p *OIDCProvider) ExchangeToken(ctx context.Context, code string, c *gin.Co

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

body, err := io.ReadAll(res.Body)
if err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-OIDC] ExchangeToken read body error: %s", err.Error()))
return nil, err
}
bodyStr := string(body)
logger.LogDebug(ctx, "[OAuth-OIDC] ExchangeToken response body: %s", bodyStr[:min(len(bodyStr), 500)])

var oidcResponse oidcOAuthResponse
err = json.NewDecoder(res.Body).Decode(&oidcResponse)
err = common.Unmarshal(body, &oidcResponse)
if err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-OIDC] ExchangeToken decode error: %s", err.Error()))
return nil, err
}

if oidcResponse.Error != "" {
logger.LogError(ctx, fmt.Sprintf("[OAuth-OIDC] ExchangeToken OAuth error: %s - %s", oidcResponse.Error, oidcResponse.ErrorDesc))
return nil, NewOAuthErrorWithRaw(i18n.MsgOAuthTokenFailed, map[string]any{"Provider": "OIDC"}, oidcResponse.ErrorDesc)
}

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

logger.LogDebug(ctx, "[OAuth-OIDC] ExchangeToken success: scope=%s", oidcResponse.Scope)
Expand Down Expand Up @@ -138,7 +154,7 @@ func (p *OIDCProvider) GetUserInfo(ctx context.Context, token *OAuthToken) (*OAu
}

var oidcUser oidcUser
err = json.NewDecoder(res.Body).Decode(&oidcUser)
err = common.DecodeJson(res.Body, &oidcUser)
if err != nil {
logger.LogError(ctx, fmt.Sprintf("[OAuth-OIDC] GetUserInfo decode error: %s", err.Error()))
return nil, err
Expand Down
35 changes: 35 additions & 0 deletions oauth/redirect_uri.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
package oauth

import (
"fmt"
"strings"

"github.com/QuantumNous/new-api/setting/system_setting"
"github.com/gin-gonic/gin"
)

func BuildOAuthRedirectURI(c *gin.Context, provider string) string {
base := ""
if c != nil && c.Request != nil {
proto := strings.TrimSpace(c.GetHeader("X-Forwarded-Proto"))
host := strings.TrimSpace(c.GetHeader("X-Forwarded-Host"))
if host == "" {
host = strings.TrimSpace(c.Request.Host)
}
if proto == "" {
if c.Request.TLS != nil {
proto = "https"
} else {
proto = "http"
}
}
if host != "" {
base = proto + "://" + host
}
}
if base == "" {
base = system_setting.ServerAddress
}
base = strings.TrimRight(base, "/")
return fmt.Sprintf("%s/oauth/%s", base, provider)
}
10 changes: 9 additions & 1 deletion web/src/features/profile/api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -173,6 +173,14 @@ export async function revokeOtherLoginSessions(): Promise<ApiResponse> {
// Custom OAuth Binding APIs
// ============================================================================

export interface CustomOAuthBinding {
provider_id: number
provider_name: string
provider_slug: string
provider_icon: string
provider_user_id: string
}

/**
* Get current user's custom OAuth bindings
*/
Expand All @@ -187,7 +195,7 @@ export async function getSelfOAuthBindings(): Promise<
* Unbind a custom OAuth provider for current user
*/
export async function unbindCustomOAuth(
providerId: number
providerId: number | string
): Promise<ApiResponse> {
const res = await api.delete(`/api/user/oauth/bindings/${providerId}`)
return res.data
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,6 @@ import {
buildOIDCOAuthUrl,
type CustomOAuthBinding,
} from '@/lib/oauth'

import { getSelfOAuthBindings, unbindCustomOAuth } from '../../api'
import type { UserProfile, BindingItem } from '../../types'
import { EmailBindDialog } from '../dialogs/email-bind-dialog'
Expand Down
Loading