Commit 29c2c895 by Seefs Committed by GitHub

imporve oauth provider UI/UX (#2983)

* feat: imporve UI/UX

* fix: stabilize provider enabled toggle and polish custom OAuth settings UX

* fix: add access policy/message templates and persist advanced fields reliably

* fix: move template fill actions below fields and keep advanced form flow cleaner
parent 37e4fccb
package controller package controller
import ( import (
"context"
"io"
"net/http" "net/http"
"net/url"
"strconv" "strconv"
"strings"
"time"
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
...@@ -16,6 +21,7 @@ type CustomOAuthProviderResponse struct { ...@@ -16,6 +21,7 @@ type CustomOAuthProviderResponse struct {
Id int `json:"id"` Id int `json:"id"`
Name string `json:"name"` Name string `json:"name"`
Slug string `json:"slug"` Slug string `json:"slug"`
Icon string `json:"icon"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
ClientId string `json:"client_id"` ClientId string `json:"client_id"`
AuthorizationEndpoint string `json:"authorization_endpoint"` AuthorizationEndpoint string `json:"authorization_endpoint"`
...@@ -28,6 +34,8 @@ type CustomOAuthProviderResponse struct { ...@@ -28,6 +34,8 @@ type CustomOAuthProviderResponse struct {
EmailField string `json:"email_field"` EmailField string `json:"email_field"`
WellKnown string `json:"well_known"` WellKnown string `json:"well_known"`
AuthStyle int `json:"auth_style"` AuthStyle int `json:"auth_style"`
AccessPolicy string `json:"access_policy"`
AccessDeniedMessage string `json:"access_denied_message"`
} }
func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthProviderResponse { func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthProviderResponse {
...@@ -35,6 +43,7 @@ func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthPro ...@@ -35,6 +43,7 @@ func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthPro
Id: p.Id, Id: p.Id,
Name: p.Name, Name: p.Name,
Slug: p.Slug, Slug: p.Slug,
Icon: p.Icon,
Enabled: p.Enabled, Enabled: p.Enabled,
ClientId: p.ClientId, ClientId: p.ClientId,
AuthorizationEndpoint: p.AuthorizationEndpoint, AuthorizationEndpoint: p.AuthorizationEndpoint,
...@@ -47,6 +56,8 @@ func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthPro ...@@ -47,6 +56,8 @@ func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthPro
EmailField: p.EmailField, EmailField: p.EmailField,
WellKnown: p.WellKnown, WellKnown: p.WellKnown,
AuthStyle: p.AuthStyle, AuthStyle: p.AuthStyle,
AccessPolicy: p.AccessPolicy,
AccessDeniedMessage: p.AccessDeniedMessage,
} }
} }
...@@ -96,6 +107,7 @@ func GetCustomOAuthProvider(c *gin.Context) { ...@@ -96,6 +107,7 @@ func GetCustomOAuthProvider(c *gin.Context) {
type CreateCustomOAuthProviderRequest struct { type CreateCustomOAuthProviderRequest struct {
Name string `json:"name" binding:"required"` Name string `json:"name" binding:"required"`
Slug string `json:"slug" binding:"required"` Slug string `json:"slug" binding:"required"`
Icon string `json:"icon"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
ClientId string `json:"client_id" binding:"required"` ClientId string `json:"client_id" binding:"required"`
ClientSecret string `json:"client_secret" binding:"required"` ClientSecret string `json:"client_secret" binding:"required"`
...@@ -109,6 +121,85 @@ type CreateCustomOAuthProviderRequest struct { ...@@ -109,6 +121,85 @@ type CreateCustomOAuthProviderRequest struct {
EmailField string `json:"email_field"` EmailField string `json:"email_field"`
WellKnown string `json:"well_known"` WellKnown string `json:"well_known"`
AuthStyle int `json:"auth_style"` AuthStyle int `json:"auth_style"`
AccessPolicy string `json:"access_policy"`
AccessDeniedMessage string `json:"access_denied_message"`
}
type FetchCustomOAuthDiscoveryRequest struct {
WellKnownURL string `json:"well_known_url"`
IssuerURL string `json:"issuer_url"`
}
// FetchCustomOAuthDiscovery fetches OIDC discovery document via backend (root-only route)
func FetchCustomOAuthDiscovery(c *gin.Context) {
var req FetchCustomOAuthDiscoveryRequest
if err := c.ShouldBindJSON(&req); err != nil {
common.ApiErrorMsg(c, "无效的请求参数: "+err.Error())
return
}
wellKnownURL := strings.TrimSpace(req.WellKnownURL)
issuerURL := strings.TrimSpace(req.IssuerURL)
if wellKnownURL == "" && issuerURL == "" {
common.ApiErrorMsg(c, "请先填写 Discovery URL 或 Issuer URL")
return
}
targetURL := wellKnownURL
if targetURL == "" {
targetURL = strings.TrimRight(issuerURL, "/") + "/.well-known/openid-configuration"
}
targetURL = strings.TrimSpace(targetURL)
parsedURL, err := url.Parse(targetURL)
if err != nil || parsedURL.Host == "" || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") {
common.ApiErrorMsg(c, "Discovery URL 无效,仅支持 http/https")
return
}
ctx, cancel := context.WithTimeout(c.Request.Context(), 20*time.Second)
defer cancel()
httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, targetURL, nil)
if err != nil {
common.ApiErrorMsg(c, "创建 Discovery 请求失败: "+err.Error())
return
}
httpReq.Header.Set("Accept", "application/json")
client := &http.Client{Timeout: 20 * time.Second}
resp, err := client.Do(httpReq)
if err != nil {
common.ApiErrorMsg(c, "获取 Discovery 配置失败: "+err.Error())
return
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
message := strings.TrimSpace(string(body))
if message == "" {
message = resp.Status
}
common.ApiErrorMsg(c, "获取 Discovery 配置失败: "+message)
return
}
var discovery map[string]any
if err = common.DecodeJson(resp.Body, &discovery); err != nil {
common.ApiErrorMsg(c, "解析 Discovery 配置失败: "+err.Error())
return
}
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": gin.H{
"well_known_url": targetURL,
"discovery": discovery,
},
})
} }
// CreateCustomOAuthProvider creates a new custom OAuth provider // CreateCustomOAuthProvider creates a new custom OAuth provider
...@@ -134,6 +225,7 @@ func CreateCustomOAuthProvider(c *gin.Context) { ...@@ -134,6 +225,7 @@ func CreateCustomOAuthProvider(c *gin.Context) {
provider := &model.CustomOAuthProvider{ provider := &model.CustomOAuthProvider{
Name: req.Name, Name: req.Name,
Slug: req.Slug, Slug: req.Slug,
Icon: req.Icon,
Enabled: req.Enabled, Enabled: req.Enabled,
ClientId: req.ClientId, ClientId: req.ClientId,
ClientSecret: req.ClientSecret, ClientSecret: req.ClientSecret,
...@@ -147,6 +239,8 @@ func CreateCustomOAuthProvider(c *gin.Context) { ...@@ -147,6 +239,8 @@ func CreateCustomOAuthProvider(c *gin.Context) {
EmailField: req.EmailField, EmailField: req.EmailField,
WellKnown: req.WellKnown, WellKnown: req.WellKnown,
AuthStyle: req.AuthStyle, AuthStyle: req.AuthStyle,
AccessPolicy: req.AccessPolicy,
AccessDeniedMessage: req.AccessDeniedMessage,
} }
if err := model.CreateCustomOAuthProvider(provider); err != nil { if err := model.CreateCustomOAuthProvider(provider); err != nil {
...@@ -168,9 +262,10 @@ func CreateCustomOAuthProvider(c *gin.Context) { ...@@ -168,9 +262,10 @@ func CreateCustomOAuthProvider(c *gin.Context) {
type UpdateCustomOAuthProviderRequest struct { type UpdateCustomOAuthProviderRequest struct {
Name string `json:"name"` Name string `json:"name"`
Slug string `json:"slug"` Slug string `json:"slug"`
Enabled *bool `json:"enabled"` // Optional: if nil, keep existing Icon *string `json:"icon"` // Optional: if nil, keep existing
Enabled *bool `json:"enabled"` // Optional: if nil, keep existing
ClientId string `json:"client_id"` ClientId string `json:"client_id"`
ClientSecret string `json:"client_secret"` // Optional: if empty, keep existing ClientSecret string `json:"client_secret"` // Optional: if empty, keep existing
AuthorizationEndpoint string `json:"authorization_endpoint"` AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"` TokenEndpoint string `json:"token_endpoint"`
UserInfoEndpoint string `json:"user_info_endpoint"` UserInfoEndpoint string `json:"user_info_endpoint"`
...@@ -181,6 +276,8 @@ type UpdateCustomOAuthProviderRequest struct { ...@@ -181,6 +276,8 @@ type UpdateCustomOAuthProviderRequest struct {
EmailField string `json:"email_field"` EmailField string `json:"email_field"`
WellKnown *string `json:"well_known"` // Optional: if nil, keep existing WellKnown *string `json:"well_known"` // Optional: if nil, keep existing
AuthStyle *int `json:"auth_style"` // Optional: if nil, keep existing AuthStyle *int `json:"auth_style"` // 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
} }
// UpdateCustomOAuthProvider updates an existing custom OAuth provider // UpdateCustomOAuthProvider updates an existing custom OAuth provider
...@@ -227,6 +324,9 @@ func UpdateCustomOAuthProvider(c *gin.Context) { ...@@ -227,6 +324,9 @@ func UpdateCustomOAuthProvider(c *gin.Context) {
if req.Slug != "" { if req.Slug != "" {
provider.Slug = req.Slug provider.Slug = req.Slug
} }
if req.Icon != nil {
provider.Icon = *req.Icon
}
if req.Enabled != nil { if req.Enabled != nil {
provider.Enabled = *req.Enabled provider.Enabled = *req.Enabled
} }
...@@ -266,6 +366,12 @@ func UpdateCustomOAuthProvider(c *gin.Context) { ...@@ -266,6 +366,12 @@ func UpdateCustomOAuthProvider(c *gin.Context) {
if req.AuthStyle != nil { if req.AuthStyle != nil {
provider.AuthStyle = *req.AuthStyle provider.AuthStyle = *req.AuthStyle
} }
if req.AccessPolicy != nil {
provider.AccessPolicy = *req.AccessPolicy
}
if req.AccessDeniedMessage != nil {
provider.AccessDeniedMessage = *req.AccessDeniedMessage
}
if err := model.UpdateCustomOAuthProvider(provider); err != nil { if err := model.UpdateCustomOAuthProvider(provider); err != nil {
common.ApiError(c, err) common.ApiError(c, err)
...@@ -346,6 +452,7 @@ func GetUserOAuthBindings(c *gin.Context) { ...@@ -346,6 +452,7 @@ func GetUserOAuthBindings(c *gin.Context) {
ProviderId int `json:"provider_id"` ProviderId int `json:"provider_id"`
ProviderName string `json:"provider_name"` ProviderName string `json:"provider_name"`
ProviderSlug string `json:"provider_slug"` ProviderSlug string `json:"provider_slug"`
ProviderIcon string `json:"provider_icon"`
ProviderUserId string `json:"provider_user_id"` ProviderUserId string `json:"provider_user_id"`
} }
...@@ -359,6 +466,7 @@ func GetUserOAuthBindings(c *gin.Context) { ...@@ -359,6 +466,7 @@ func GetUserOAuthBindings(c *gin.Context) {
ProviderId: binding.ProviderId, ProviderId: binding.ProviderId,
ProviderName: provider.Name, ProviderName: provider.Name,
ProviderSlug: provider.Slug, ProviderSlug: provider.Slug,
ProviderIcon: provider.Icon,
ProviderUserId: binding.ProviderUserId, ProviderUserId: binding.ProviderUserId,
}) })
} }
......
...@@ -134,8 +134,10 @@ func GetStatus(c *gin.Context) { ...@@ -134,8 +134,10 @@ func GetStatus(c *gin.Context) {
customProviders := oauth.GetEnabledCustomProviders() customProviders := oauth.GetEnabledCustomProviders()
if len(customProviders) > 0 { if len(customProviders) > 0 {
type CustomOAuthInfo struct { type CustomOAuthInfo struct {
Id int `json:"id"`
Name string `json:"name"` Name string `json:"name"`
Slug string `json:"slug"` Slug string `json:"slug"`
Icon string `json:"icon"`
ClientId string `json:"client_id"` ClientId string `json:"client_id"`
AuthorizationEndpoint string `json:"authorization_endpoint"` AuthorizationEndpoint string `json:"authorization_endpoint"`
Scopes string `json:"scopes"` Scopes string `json:"scopes"`
...@@ -144,8 +146,10 @@ func GetStatus(c *gin.Context) { ...@@ -144,8 +146,10 @@ func GetStatus(c *gin.Context) {
for _, p := range customProviders { for _, p := range customProviders {
config := p.GetConfig() config := p.GetConfig()
providersInfo = append(providersInfo, CustomOAuthInfo{ providersInfo = append(providersInfo, CustomOAuthInfo{
Id: config.Id,
Name: config.Name, Name: config.Name,
Slug: config.Slug, Slug: config.Slug,
Icon: config.Icon,
ClientId: config.ClientId, ClientId: config.ClientId,
AuthorizationEndpoint: config.AuthorizationEndpoint, AuthorizationEndpoint: config.AuthorizationEndpoint,
Scopes: config.Scopes, Scopes: config.Scopes,
......
...@@ -295,12 +295,12 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o ...@@ -295,12 +295,12 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
// Set the provider user ID on the user model and update // Set the provider user ID on the user model and update
provider.SetProviderUserID(user, oauthUser.ProviderUserID) provider.SetProviderUserID(user, oauthUser.ProviderUserID)
if err := tx.Model(user).Updates(map[string]interface{}{ if err := tx.Model(user).Updates(map[string]interface{}{
"github_id": user.GitHubId, "github_id": user.GitHubId,
"discord_id": user.DiscordId, "discord_id": user.DiscordId,
"oidc_id": user.OidcId, "oidc_id": user.OidcId,
"linux_do_id": user.LinuxDOId, "linux_do_id": user.LinuxDOId,
"wechat_id": user.WeChatId, "wechat_id": user.WeChatId,
"telegram_id": user.TelegramId, "telegram_id": user.TelegramId,
}).Error; err != nil { }).Error; err != nil {
return err return err
} }
...@@ -340,6 +340,8 @@ func handleOAuthError(c *gin.Context, err error) { ...@@ -340,6 +340,8 @@ func handleOAuthError(c *gin.Context, err error) {
} else { } else {
common.ApiErrorI18n(c, e.MsgKey) common.ApiErrorI18n(c, e.MsgKey)
} }
case *oauth.AccessDeniedError:
common.ApiErrorMsg(c, e.Message)
case *oauth.TrustLevelError: case *oauth.TrustLevelError:
common.ApiErrorI18n(c, i18n.MsgOAuthTrustLevelLow) common.ApiErrorI18n(c, i18n.MsgOAuthTrustLevelLow)
default: default:
......
...@@ -2,32 +2,65 @@ package model ...@@ -2,32 +2,65 @@ package model
import ( import (
"errors" "errors"
"fmt"
"strings" "strings"
"time" "time"
"github.com/QuantumNous/new-api/common"
) )
type accessPolicyPayload struct {
Logic string `json:"logic"`
Conditions []accessConditionItem `json:"conditions"`
Groups []accessPolicyPayload `json:"groups"`
}
type accessConditionItem struct {
Field string `json:"field"`
Op string `json:"op"`
Value any `json:"value"`
}
var supportedAccessPolicyOps = map[string]struct{}{
"eq": {},
"ne": {},
"gt": {},
"gte": {},
"lt": {},
"lte": {},
"in": {},
"not_in": {},
"contains": {},
"not_contains": {},
"exists": {},
"not_exists": {},
}
// CustomOAuthProvider stores configuration for custom OAuth providers // CustomOAuthProvider stores configuration for custom OAuth providers
type CustomOAuthProvider struct { type CustomOAuthProvider struct {
Id int `json:"id" gorm:"primaryKey"` Id int `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"type:varchar(64);not null"` // Display name, e.g., "GitHub Enterprise" Name string `json:"name" gorm:"type:varchar(64);not null"` // Display name, e.g., "GitHub Enterprise"
Slug string `json:"slug" gorm:"type:varchar(64);uniqueIndex;not null"` // URL identifier, e.g., "github-enterprise" Slug string `json:"slug" gorm:"type:varchar(64);uniqueIndex;not null"` // URL identifier, e.g., "github-enterprise"
Enabled bool `json:"enabled" gorm:"default:false"` // Whether this provider is enabled Icon string `json:"icon" gorm:"type:varchar(128);default:''"` // Icon name from @lobehub/icons
ClientId string `json:"client_id" gorm:"type:varchar(256)"` // OAuth client ID Enabled bool `json:"enabled" gorm:"default:false"` // Whether this provider is enabled
ClientSecret string `json:"-" gorm:"type:varchar(512)"` // OAuth client secret (not returned to frontend) ClientId string `json:"client_id" gorm:"type:varchar(256)"` // OAuth client ID
AuthorizationEndpoint string `json:"authorization_endpoint" gorm:"type:varchar(512)"` // Authorization URL ClientSecret string `json:"-" gorm:"type:varchar(512)"` // OAuth client secret (not returned to frontend)
TokenEndpoint string `json:"token_endpoint" gorm:"type:varchar(512)"` // Token exchange URL AuthorizationEndpoint string `json:"authorization_endpoint" gorm:"type:varchar(512)"` // Authorization URL
UserInfoEndpoint string `json:"user_info_endpoint" gorm:"type:varchar(512)"` // User info URL TokenEndpoint string `json:"token_endpoint" gorm:"type:varchar(512)"` // Token exchange URL
Scopes string `json:"scopes" gorm:"type:varchar(256);default:'openid profile email'"` // OAuth scopes UserInfoEndpoint string `json:"user_info_endpoint" gorm:"type:varchar(512)"` // User info URL
Scopes string `json:"scopes" gorm:"type:varchar(256);default:'openid profile email'"` // OAuth scopes
// Field mapping configuration (supports JSONPath via gjson) // Field mapping configuration (supports JSONPath via gjson)
UserIdField string `json:"user_id_field" gorm:"type:varchar(128);default:'sub'"` // User ID field path, e.g., "sub", "id", "data.user.id" UserIdField string `json:"user_id_field" gorm:"type:varchar(128);default:'sub'"` // User ID field path, e.g., "sub", "id", "data.user.id"
UsernameField string `json:"username_field" gorm:"type:varchar(128);default:'preferred_username'"` // Username field path UsernameField string `json:"username_field" gorm:"type:varchar(128);default:'preferred_username'"` // Username field path
DisplayNameField string `json:"display_name_field" gorm:"type:varchar(128);default:'name'"` // Display name field path DisplayNameField string `json:"display_name_field" gorm:"type:varchar(128);default:'name'"` // Display name field path
EmailField string `json:"email_field" gorm:"type:varchar(128);default:'email'"` // Email field path EmailField string `json:"email_field" gorm:"type:varchar(128);default:'email'"` // Email field path
// Advanced options // Advanced options
WellKnown string `json:"well_known" gorm:"type:varchar(512)"` // OIDC discovery endpoint (optional) 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) AuthStyle int `json:"auth_style" gorm:"default:0"` // 0=auto, 1=params, 2=header (Basic Auth)
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
CreatedAt time.Time `json:"created_at"` CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"` UpdatedAt time.Time `json:"updated_at"`
...@@ -158,6 +191,57 @@ func validateCustomOAuthProvider(provider *CustomOAuthProvider) error { ...@@ -158,6 +191,57 @@ func validateCustomOAuthProvider(provider *CustomOAuthProvider) error {
if provider.Scopes == "" { if provider.Scopes == "" {
provider.Scopes = "openid profile email" provider.Scopes = "openid profile email"
} }
if strings.TrimSpace(provider.AccessPolicy) != "" {
var policy accessPolicyPayload
if err := common.UnmarshalJsonStr(provider.AccessPolicy, &policy); err != nil {
return errors.New("access_policy must be valid JSON")
}
if err := validateAccessPolicyPayload(&policy); err != nil {
return fmt.Errorf("access_policy is invalid: %w", err)
}
}
return nil
}
func validateAccessPolicyPayload(policy *accessPolicyPayload) error {
if policy == nil {
return errors.New("policy is nil")
}
logic := strings.ToLower(strings.TrimSpace(policy.Logic))
if logic == "" {
logic = "and"
}
if logic != "and" && logic != "or" {
return fmt.Errorf("unsupported logic: %s", logic)
}
if len(policy.Conditions) == 0 && len(policy.Groups) == 0 {
return errors.New("policy requires at least one condition or group")
}
for index, condition := range policy.Conditions {
field := strings.TrimSpace(condition.Field)
if field == "" {
return fmt.Errorf("condition[%d].field is required", index)
}
op := strings.ToLower(strings.TrimSpace(condition.Op))
if _, ok := supportedAccessPolicyOps[op]; !ok {
return fmt.Errorf("condition[%d].op is unsupported: %s", index, op)
}
if op == "in" || op == "not_in" {
if _, ok := condition.Value.([]any); !ok {
return fmt.Errorf("condition[%d].value must be an array for op %s", index, op)
}
}
}
for index := range policy.Groups {
if err := validateAccessPolicyPayload(&policy.Groups[index]); err != nil {
return fmt.Errorf("group[%d]: %w", index, err)
}
}
return nil return nil
} }
...@@ -57,3 +57,12 @@ func NewOAuthErrorWithRaw(msgKey string, params map[string]any, rawError string) ...@@ -57,3 +57,12 @@ func NewOAuthErrorWithRaw(msgKey string, params map[string]any, rawError string)
RawError: rawError, RawError: rawError,
} }
} }
// AccessDeniedError is a direct user-facing access denial message.
type AccessDeniedError struct {
Message string
}
func (e *AccessDeniedError) Error() string {
return e.Message
}
...@@ -170,10 +170,11 @@ func SetApiRouter(router *gin.Engine) { ...@@ -170,10 +170,11 @@ func SetApiRouter(router *gin.Engine) {
optionRoute.POST("/migrate_console_setting", controller.MigrateConsoleSetting) // 用于迁移检测的旧键,下个版本会删除 optionRoute.POST("/migrate_console_setting", controller.MigrateConsoleSetting) // 用于迁移检测的旧键,下个版本会删除
} }
// Custom OAuth provider management (admin only) // Custom OAuth provider management (root only)
customOAuthRoute := apiRouter.Group("/custom-oauth-provider") customOAuthRoute := apiRouter.Group("/custom-oauth-provider")
customOAuthRoute.Use(middleware.RootAuth()) customOAuthRoute.Use(middleware.RootAuth())
{ {
customOAuthRoute.POST("/discovery", controller.FetchCustomOAuthDiscovery)
customOAuthRoute.GET("/", controller.GetCustomOAuthProviders) customOAuthRoute.GET("/", controller.GetCustomOAuthProviders)
customOAuthRoute.GET("/:id", controller.GetCustomOAuthProvider) customOAuthRoute.GET("/:id", controller.GetCustomOAuthProvider)
customOAuthRoute.POST("/", controller.CreateCustomOAuthProvider) customOAuthRoute.POST("/", controller.CreateCustomOAuthProvider)
......
...@@ -29,6 +29,7 @@ import { ...@@ -29,6 +29,7 @@ import {
showSuccess, showSuccess,
updateAPI, updateAPI,
getSystemName, getSystemName,
getOAuthProviderIcon,
setUserData, setUserData,
onGitHubOAuthClicked, onGitHubOAuthClicked,
onDiscordOAuthClicked, onDiscordOAuthClicked,
...@@ -130,6 +131,17 @@ const LoginForm = () => { ...@@ -130,6 +131,17 @@ const LoginForm = () => {
return {}; return {};
} }
}, [statusState?.status]); }, [statusState?.status]);
const hasCustomOAuthProviders =
(status.custom_oauth_providers || []).length > 0;
const hasOAuthLoginOptions = Boolean(
status.github_oauth ||
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
status.linuxdo_oauth ||
status.telegram_oauth ||
hasCustomOAuthProviders,
);
useEffect(() => { useEffect(() => {
if (status?.turnstile_check) { if (status?.turnstile_check) {
...@@ -598,7 +610,7 @@ const LoginForm = () => { ...@@ -598,7 +610,7 @@ const LoginForm = () => {
theme='outline' theme='outline'
className='w-full h-12 flex items-center justify-center !rounded-full border border-gray-200 hover:bg-gray-50 transition-colors' className='w-full h-12 flex items-center justify-center !rounded-full border border-gray-200 hover:bg-gray-50 transition-colors'
type='tertiary' type='tertiary'
icon={<IconLock size='large' />} icon={getOAuthProviderIcon(provider.icon || '', 20)}
onClick={() => handleCustomOAuthClick(provider)} onClick={() => handleCustomOAuthClick(provider)}
loading={customOAuthLoading[provider.slug]} loading={customOAuthLoading[provider.slug]}
> >
...@@ -817,12 +829,7 @@ const LoginForm = () => { ...@@ -817,12 +829,7 @@ const LoginForm = () => {
</div> </div>
</Form> </Form>
{(status.github_oauth || {hasOAuthLoginOptions && (
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
status.linuxdo_oauth ||
status.telegram_oauth) && (
<> <>
<Divider margin='12px' align='center'> <Divider margin='12px' align='center'>
{t('或')} {t('或')}
...@@ -952,14 +959,7 @@ const LoginForm = () => { ...@@ -952,14 +959,7 @@ const LoginForm = () => {
/> />
<div className='w-full max-w-sm mt-[60px]'> <div className='w-full max-w-sm mt-[60px]'>
{showEmailLogin || {showEmailLogin ||
!( !hasOAuthLoginOptions
status.github_oauth ||
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
status.linuxdo_oauth ||
status.telegram_oauth
)
? renderEmailLoginForm() ? renderEmailLoginForm()
: renderOAuthOptions()} : renderOAuthOptions()}
{renderWeChatLoginModal()} {renderWeChatLoginModal()}
......
...@@ -27,8 +27,10 @@ import { ...@@ -27,8 +27,10 @@ import {
showSuccess, showSuccess,
updateAPI, updateAPI,
getSystemName, getSystemName,
getOAuthProviderIcon,
setUserData, setUserData,
onDiscordOAuthClicked, onDiscordOAuthClicked,
onCustomOAuthClicked,
} from '../../helpers'; } from '../../helpers';
import Turnstile from 'react-turnstile'; import Turnstile from 'react-turnstile';
import { import {
...@@ -98,6 +100,7 @@ const RegisterForm = () => { ...@@ -98,6 +100,7 @@ const RegisterForm = () => {
const [otherRegisterOptionsLoading, setOtherRegisterOptionsLoading] = const [otherRegisterOptionsLoading, setOtherRegisterOptionsLoading] =
useState(false); useState(false);
const [wechatCodeSubmitLoading, setWechatCodeSubmitLoading] = useState(false); const [wechatCodeSubmitLoading, setWechatCodeSubmitLoading] = useState(false);
const [customOAuthLoading, setCustomOAuthLoading] = useState({});
const [disableButton, setDisableButton] = useState(false); const [disableButton, setDisableButton] = useState(false);
const [countdown, setCountdown] = useState(30); const [countdown, setCountdown] = useState(30);
const [agreedToTerms, setAgreedToTerms] = useState(false); const [agreedToTerms, setAgreedToTerms] = useState(false);
...@@ -126,6 +129,17 @@ const RegisterForm = () => { ...@@ -126,6 +129,17 @@ const RegisterForm = () => {
return {}; return {};
} }
}, [statusState?.status]); }, [statusState?.status]);
const hasCustomOAuthProviders =
(status.custom_oauth_providers || []).length > 0;
const hasOAuthRegisterOptions = Boolean(
status.github_oauth ||
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
status.linuxdo_oauth ||
status.telegram_oauth ||
hasCustomOAuthProviders,
);
const [showEmailVerification, setShowEmailVerification] = useState(false); const [showEmailVerification, setShowEmailVerification] = useState(false);
...@@ -319,6 +333,17 @@ const RegisterForm = () => { ...@@ -319,6 +333,17 @@ const RegisterForm = () => {
} }
}; };
const handleCustomOAuthClick = (provider) => {
setCustomOAuthLoading((prev) => ({ ...prev, [provider.slug]: true }));
try {
onCustomOAuthClicked(provider, { shouldLogout: true });
} finally {
setTimeout(() => {
setCustomOAuthLoading((prev) => ({ ...prev, [provider.slug]: false }));
}, 3000);
}
};
const handleEmailRegisterClick = () => { const handleEmailRegisterClick = () => {
setEmailRegisterLoading(true); setEmailRegisterLoading(true);
setShowEmailRegister(true); setShowEmailRegister(true);
...@@ -469,6 +494,23 @@ const RegisterForm = () => { ...@@ -469,6 +494,23 @@ const RegisterForm = () => {
</Button> </Button>
)} )}
{status.custom_oauth_providers &&
status.custom_oauth_providers.map((provider) => (
<Button
key={provider.slug}
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={getOAuthProviderIcon(provider.icon || '', 20)}
onClick={() => handleCustomOAuthClick(provider)}
loading={customOAuthLoading[provider.slug]}
>
<span className='ml-3'>
{t('使用 {{name}} 继续', { name: provider.name })}
</span>
</Button>
))}
{status.telegram_oauth && ( {status.telegram_oauth && (
<div className='flex justify-center my-2'> <div className='flex justify-center my-2'>
<TelegramLoginButton <TelegramLoginButton
...@@ -650,12 +692,7 @@ const RegisterForm = () => { ...@@ -650,12 +692,7 @@ const RegisterForm = () => {
</div> </div>
</Form> </Form>
{(status.github_oauth || {hasOAuthRegisterOptions && (
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
status.linuxdo_oauth ||
status.telegram_oauth) && (
<> <>
<Divider margin='12px' align='center'> <Divider margin='12px' align='center'>
{t('或')} {t('或')}
...@@ -745,14 +782,7 @@ const RegisterForm = () => { ...@@ -745,14 +782,7 @@ const RegisterForm = () => {
/> />
<div className='w-full max-w-sm mt-[60px]'> <div className='w-full max-w-sm mt-[60px]'>
{showEmailRegister || {showEmailRegister ||
!( !hasOAuthRegisterOptions
status.github_oauth ||
status.discord_oauth ||
status.oidc_enabled ||
status.wechat_login ||
status.linuxdo_oauth ||
status.telegram_oauth
)
? renderEmailRegisterForm() ? renderEmailRegisterForm()
: renderOAuthOptions()} : renderOAuthOptions()}
{renderWeChatLoginModal()} {renderWeChatLoginModal()}
......
...@@ -50,6 +50,7 @@ import { ...@@ -50,6 +50,7 @@ import {
onLinuxDOOAuthClicked, onLinuxDOOAuthClicked,
onDiscordOAuthClicked, onDiscordOAuthClicked,
onCustomOAuthClicked, onCustomOAuthClicked,
getOAuthProviderIcon,
} from '../../../../helpers'; } from '../../../../helpers';
import TwoFASetting from '../components/TwoFASetting'; import TwoFASetting from '../components/TwoFASetting';
...@@ -148,12 +149,14 @@ const AccountManagement = ({ ...@@ -148,12 +149,14 @@ const AccountManagement = ({
// Check if custom OAuth provider is bound // Check if custom OAuth provider is bound
const isCustomOAuthBound = (providerId) => { const isCustomOAuthBound = (providerId) => {
return customOAuthBindings.some((b) => b.provider_id === providerId); const normalizedId = Number(providerId);
return customOAuthBindings.some((b) => Number(b.provider_id) === normalizedId);
}; };
// Get binding info for a provider // Get binding info for a provider
const getCustomOAuthBinding = (providerId) => { const getCustomOAuthBinding = (providerId) => {
return customOAuthBindings.find((b) => b.provider_id === providerId); const normalizedId = Number(providerId);
return customOAuthBindings.find((b) => Number(b.provider_id) === normalizedId);
}; };
React.useEffect(() => { React.useEffect(() => {
...@@ -524,10 +527,10 @@ const AccountManagement = ({ ...@@ -524,10 +527,10 @@ const AccountManagement = ({
<div className='flex items-center justify-between gap-3'> <div className='flex items-center justify-between gap-3'>
<div className='flex items-center flex-1 min-w-0'> <div className='flex items-center flex-1 min-w-0'>
<div className='w-10 h-10 rounded-full bg-slate-100 dark:bg-slate-700 flex items-center justify-center mr-3 flex-shrink-0'> <div className='w-10 h-10 rounded-full bg-slate-100 dark:bg-slate-700 flex items-center justify-center mr-3 flex-shrink-0'>
<IconLock {getOAuthProviderIcon(
size='default' provider.icon || binding?.provider_icon || '',
className='text-slate-600 dark:text-slate-300' 20,
/> )}
</div> </div>
<div className='flex-1 min-w-0'> <div className='flex-1 min-w-0'>
<div className='font-medium text-gray-900'> <div className='font-medium text-gray-900'>
......
...@@ -76,6 +76,31 @@ import { ...@@ -76,6 +76,31 @@ import {
Server, Server,
CalendarClock, CalendarClock,
} from 'lucide-react'; } from 'lucide-react';
import {
SiAtlassian,
SiAuth0,
SiAuthentik,
SiBitbucket,
SiDiscord,
SiDropbox,
SiFacebook,
SiGitea,
SiGithub,
SiGitlab,
SiGoogle,
SiKeycloak,
SiLinkedin,
SiNextcloud,
SiNotion,
SiOkta,
SiOpenid,
SiReddit,
SiSlack,
SiTelegram,
SiTwitch,
SiWechat,
SiX,
} from 'react-icons/si';
// 获取侧边栏Lucide图标组件 // 获取侧边栏Lucide图标组件
export function getLucideIcon(key, selected = false) { export function getLucideIcon(key, selected = false) {
...@@ -472,6 +497,106 @@ export function getLobeHubIcon(iconName, size = 14) { ...@@ -472,6 +497,106 @@ export function getLobeHubIcon(iconName, size = 14) {
return <IconComponent {...props} />; return <IconComponent {...props} />;
} }
const oauthProviderIconMap = {
github: SiGithub,
gitlab: SiGitlab,
gitea: SiGitea,
google: SiGoogle,
discord: SiDiscord,
facebook: SiFacebook,
linkedin: SiLinkedin,
x: SiX,
twitter: SiX,
slack: SiSlack,
telegram: SiTelegram,
wechat: SiWechat,
keycloak: SiKeycloak,
nextcloud: SiNextcloud,
authentik: SiAuthentik,
openid: SiOpenid,
okta: SiOkta,
auth0: SiAuth0,
atlassian: SiAtlassian,
bitbucket: SiBitbucket,
notion: SiNotion,
twitch: SiTwitch,
reddit: SiReddit,
dropbox: SiDropbox,
};
function isHttpUrl(value) {
return /^https?:\/\//i.test(value || '');
}
function isSimpleEmoji(value) {
if (!value) return false;
const trimmed = String(value).trim();
return trimmed.length > 0 && trimmed.length <= 4 && !isHttpUrl(trimmed);
}
function normalizeOAuthIconKey(raw) {
return raw
.trim()
.toLowerCase()
.replace(/^ri:/, '')
.replace(/^react-icons:/, '')
.replace(/^si:/, '');
}
/**
* Render custom OAuth provider icon with react-icons or URL/emoji fallback.
* Supported formats:
* - react-icons simple key: github / gitlab / google / keycloak
* - prefixed key: ri:github / si:github
* - full URL image: https://example.com/logo.png
* - emoji: 🐱
*/
export function getOAuthProviderIcon(iconName, size = 20) {
const raw = String(iconName || '').trim();
const iconSize = Number(size) > 0 ? Number(size) : 20;
if (!raw) {
return <Layers size={iconSize} color='var(--semi-color-text-2)' />;
}
if (isHttpUrl(raw)) {
return (
<img
src={raw}
alt='provider icon'
width={iconSize}
height={iconSize}
style={{ borderRadius: 4, objectFit: 'cover' }}
/>
);
}
if (isSimpleEmoji(raw)) {
return (
<span
style={{
width: iconSize,
height: iconSize,
lineHeight: `${iconSize}px`,
textAlign: 'center',
display: 'inline-block',
fontSize: Math.max(Math.floor(iconSize * 0.8), 14),
}}
>
{raw}
</span>
);
}
const key = normalizeOAuthIconKey(raw);
const IconComp = oauthProviderIconMap[key];
if (IconComp) {
return <IconComp size={iconSize} />;
}
return <Avatar size='extra-extra-small'>{raw.charAt(0).toUpperCase()}</Avatar>;
}
// 颜色列表 // 颜色列表
const colors = [ const colors = [
'amber', 'amber',
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment