diff --git a/README.fr.md b/README.fr.md
index 6b4d0cebafee..0036aeb73229 100644
--- a/README.fr.md
+++ b/README.fr.md
@@ -202,6 +202,7 @@ docker run --name new-api -d --restart always \
- 🤖 Connexion par autorisation LinuxDO
- 📱 Connexion par autorisation Telegram
- 🔑 Authentification unifiée OIDC
+- Enterprise SSO (JWT Direct) : validation JWT directe, echange de ticket, UserInfo et flux CAS validate
- 🔍 Requête de quota d'utilisation de clé (avec [neko-api-key-tool](https://github.com/Calcium-Ion/neko-api-key-tool))
### 🚀 Fonctionnalités avancées
diff --git a/README.ja.md b/README.ja.md
index 2b35bdfe9b93..23b70b9b3432 100644
--- a/README.ja.md
+++ b/README.ja.md
@@ -202,6 +202,7 @@ docker run --name new-api -d --restart always \
- 🤖 LinuxDO認証ログイン
- 📱 Telegram認証ログイン
- 🔑 OIDC統一認証
+- Enterprise SSO(JWT Direct): 直接JWT検証、チケット交換、UserInfo、CAS Validate フロー
- 🔍 Key使用量クォータ照会([neko-api-key-tool](https://github.com/Calcium-Ion/neko-api-key-tool)と併用)
diff --git a/README.md b/README.md
index 8f23d5dcd380..70c9b367f136 100644
--- a/README.md
+++ b/README.md
@@ -202,6 +202,7 @@ docker run --name new-api -d --restart always \
- 🤖 LinuxDO authorization login
- 📱 Telegram authorization login
- 🔑 OIDC unified authentication
+- Enterprise SSO (JWT Direct): direct JWT validation, ticket exchange, UserInfo, and CAS validate flows
- 🔍 Key quota query usage (with [neko-api-key-tool](https://github.com/Calcium-Ion/neko-api-key-tool))
### 🚀 Advanced Features
diff --git a/README.zh_CN.md b/README.zh_CN.md
index 92e5baa1d212..4009336b1363 100644
--- a/README.zh_CN.md
+++ b/README.zh_CN.md
@@ -202,6 +202,7 @@ docker run --name new-api -d --restart always \
- 🤖 LinuxDO 授权登录
- 📱 Telegram 授权登录
- 🔑 OIDC 统一认证
+- 企业 SSO(JWT Direct):支持直验 JWT、票据换 JWT、UserInfo 和 CAS Validate 接入
- 🔍 Key 查询使用额度(配合 [neko-api-key-tool](https://github.com/Calcium-Ion/neko-api-key-tool))
### 🚀 高级功能
diff --git a/README.zh_TW.md b/README.zh_TW.md
index 63664f0c863f..de831c2364c3 100644
--- a/README.zh_TW.md
+++ b/README.zh_TW.md
@@ -202,6 +202,7 @@ docker run --name new-api -d --restart always \
- 🤖 LinuxDO 授權登錄
- 📱 Telegram 授權登錄
- 🔑 OIDC 統一認證
+- 企業 SSO(JWT Direct):支援直接驗證 JWT、票據換 JWT、UserInfo 與 CAS Validate 接入
- 🔍 Key 查詢使用額度(配合 [neko-api-key-tool](https://github.com/Calcium-Ion/neko-api-key-tool))
### 🚀 高級功能
diff --git a/common/json.go b/common/json.go
index 54f8baa34229..00fec9c233af 100644
--- a/common/json.go
+++ b/common/json.go
@@ -6,6 +6,8 @@ import (
"io"
)
+type RawMessage = json.RawMessage
+
func Unmarshal(data []byte, v any) error {
return json.Unmarshal(data, v)
}
diff --git a/controller/bind_session_test.go b/controller/bind_session_test.go
new file mode 100644
index 000000000000..129db8886af2
--- /dev/null
+++ b/controller/bind_session_test.go
@@ -0,0 +1,183 @@
+package controller
+
+import (
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/hex"
+ "sync/atomic"
+ "net/http"
+ "net/http/httptest"
+ "net/url"
+ "sort"
+ "strconv"
+ "testing"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-contrib/sessions/cookie"
+ "github.com/gin-gonic/gin"
+)
+
+func newBindSessionTestRouter(t *testing.T) *gin.Engine {
+ t.Helper()
+
+ router := gin.New()
+ store := cookie.NewStore([]byte("bind-session-test-secret"))
+ router.Use(sessions.Sessions("session", store))
+ router.GET("/api/oauth/email/bind", EmailBind)
+ router.GET("/api/oauth/wechat/bind", WeChatBind)
+ router.GET("/api/oauth/telegram/bind", TelegramBind)
+ return router
+}
+
+func performBindSessionTestRequest(t *testing.T, serverURL string, path string) oauthJWTAPIResponse {
+ t.Helper()
+
+ response, err := http.Get(serverURL + path)
+ if err != nil {
+ t.Fatalf("failed to perform request: %v", err)
+ }
+ defer response.Body.Close()
+
+ var payload oauthJWTAPIResponse
+ if err := common.DecodeJson(response.Body, &payload); err != nil {
+ t.Fatalf("failed to decode response: %v", err)
+ }
+ return payload
+}
+
+func TestEmailBindRequiresLoggedInSession(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ router := newBindSessionTestRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ email := "bind-email@example.com"
+ code := "123456"
+ common.RegisterVerificationCodeWithKey(email, code, common.EmailVerificationPurpose)
+ t.Cleanup(func() {
+ common.DeleteKey(email, common.EmailVerificationPurpose)
+ })
+
+ response := performBindSessionTestRequest(
+ t,
+ server.URL,
+ "/api/oauth/email/bind?email="+url.QueryEscape(email)+"&code="+url.QueryEscape(code),
+ )
+
+ if response.Success {
+ t.Fatalf("expected email bind without session to fail")
+ }
+ if response.Message != "未登录" {
+ t.Fatalf("unexpected error message: %s", response.Message)
+ }
+}
+
+func TestWeChatBindRequiresLoggedInSession(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ router := newBindSessionTestRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ oldEnabled := common.WeChatAuthEnabled
+ oldAddress := common.WeChatServerAddress
+ oldToken := common.WeChatServerToken
+ defer func() {
+ common.WeChatAuthEnabled = oldEnabled
+ common.WeChatServerAddress = oldAddress
+ common.WeChatServerToken = oldToken
+ }()
+
+ var upstreamRequests int32
+ wechatServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ atomic.AddInt32(&upstreamRequests, 1)
+ _, _ = w.Write([]byte(`{"success":true,"message":"","data":"wechat-user-1"}`))
+ }))
+ defer wechatServer.Close()
+
+ common.WeChatAuthEnabled = true
+ common.WeChatServerAddress = wechatServer.URL
+ common.WeChatServerToken = "test-wechat-token"
+
+ response := performBindSessionTestRequest(
+ t,
+ server.URL,
+ "/api/oauth/wechat/bind?code=wechat-code",
+ )
+
+ if response.Success {
+ t.Fatalf("expected wechat bind without session to fail")
+ }
+ if response.Message != "未登录" {
+ t.Fatalf("unexpected error message: %s", response.Message)
+ }
+ if atomic.LoadInt32(&upstreamRequests) != 0 {
+ t.Fatalf("expected unauthenticated wechat bind to avoid upstream requests, got %d", atomic.LoadInt32(&upstreamRequests))
+ }
+}
+
+func TestTelegramBindRequiresLoggedInSession(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ router := newBindSessionTestRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ oldEnabled := common.TelegramOAuthEnabled
+ oldToken := common.TelegramBotToken
+ defer func() {
+ common.TelegramOAuthEnabled = oldEnabled
+ common.TelegramBotToken = oldToken
+ }()
+
+ common.TelegramOAuthEnabled = true
+ common.TelegramBotToken = "test-telegram-bot-token"
+
+ params := buildTelegramAuthParams(common.TelegramBotToken, "123456789")
+ response := performBindSessionTestRequest(
+ t,
+ server.URL,
+ "/api/oauth/telegram/bind?"+params.Encode(),
+ )
+
+ if response.Success {
+ t.Fatalf("expected telegram bind without session to fail")
+ }
+ if response.Message != "未登录" {
+ t.Fatalf("unexpected error message: %s", response.Message)
+ }
+}
+
+func buildTelegramAuthParams(botToken string, telegramID string) url.Values {
+ params := url.Values{}
+ params.Set("id", telegramID)
+ params.Set("first_name", "Bind")
+ params.Set("auth_date", strconv.FormatInt(time.Now().Unix(), 10))
+ params.Set("hash", telegramAuthHash(params, botToken))
+ return params
+}
+
+func telegramAuthHash(params url.Values, token string) string {
+ items := make([]string, 0, len(params))
+ for key, values := range params {
+ if key == "hash" || len(values) == 0 {
+ continue
+ }
+ items = append(items, key+"="+values[0])
+ }
+ sort.Strings(items)
+
+ payload := ""
+ for index, item := range items {
+ if index > 0 {
+ payload += "\n"
+ }
+ payload += item
+ }
+
+ sha256hash := sha256.New()
+ _, _ = sha256hash.Write([]byte(token))
+ hmacHash := hmac.New(sha256.New, sha256hash.Sum(nil))
+ _, _ = hmacHash.Write([]byte(payload))
+ return hex.EncodeToString(hmacHash.Sum(nil))
+}
diff --git a/controller/custom_oauth.go b/controller/custom_oauth.go
index c21ec7910bce..8e1e0c1de4cb 100644
--- a/controller/custom_oauth.go
+++ b/controller/custom_oauth.go
@@ -18,24 +18,50 @@ import (
// CustomOAuthProviderResponse is the response structure for custom OAuth providers
// It excludes sensitive fields like client_secret
type CustomOAuthProviderResponse struct {
- Id int `json:"id"`
- Name string `json:"name"`
- Slug string `json:"slug"`
- Icon string `json:"icon"`
- Enabled bool `json:"enabled"`
- ClientId string `json:"client_id"`
- AuthorizationEndpoint string `json:"authorization_endpoint"`
- TokenEndpoint string `json:"token_endpoint"`
- UserInfoEndpoint string `json:"user_info_endpoint"`
- Scopes string `json:"scopes"`
- UserIdField string `json:"user_id_field"`
- UsernameField string `json:"username_field"`
- DisplayNameField string `json:"display_name_field"`
- EmailField string `json:"email_field"`
- WellKnown string `json:"well_known"`
- AuthStyle int `json:"auth_style"`
- AccessPolicy string `json:"access_policy"`
- AccessDeniedMessage string `json:"access_denied_message"`
+ Id int `json:"id"`
+ Name string `json:"name"`
+ Slug string `json:"slug"`
+ Icon string `json:"icon"`
+ Kind string `json:"kind"`
+ Enabled bool `json:"enabled"`
+ ClientId string `json:"client_id"`
+ AuthorizationEndpoint string `json:"authorization_endpoint"`
+ TokenEndpoint string `json:"token_endpoint"`
+ UserInfoEndpoint string `json:"user_info_endpoint"`
+ Scopes string `json:"scopes"`
+ Issuer string `json:"issuer"`
+ Audience string `json:"audience"`
+ JwksURL string `json:"jwks_url"`
+ PublicKey string `json:"public_key"`
+ JWTSource string `json:"jwt_source"`
+ JWTHeader string `json:"jwt_header"`
+ JWTIdentityMode string `json:"jwt_identity_mode"`
+ JWTAcquireMode string `json:"jwt_acquire_mode"`
+ AuthorizationServiceField string `json:"authorization_service_field"`
+ TicketExchangeURL string `json:"ticket_exchange_url"`
+ TicketExchangeMethod string `json:"ticket_exchange_method"`
+ TicketExchangePayloadMode string `json:"ticket_exchange_payload_mode"`
+ TicketExchangeTicketField string `json:"ticket_exchange_ticket_field"`
+ TicketExchangeTokenField string `json:"ticket_exchange_token_field"`
+ TicketExchangeServiceField string `json:"ticket_exchange_service_field"`
+ UserIdField string `json:"user_id_field"`
+ UsernameField string `json:"username_field"`
+ DisplayNameField string `json:"display_name_field"`
+ EmailField string `json:"email_field"`
+ GroupField string `json:"group_field"`
+ RoleField string `json:"role_field"`
+ GroupMapping string `json:"group_mapping"`
+ RoleMapping string `json:"role_mapping"`
+ AutoRegister bool `json:"auto_register"`
+ AutoMergeByEmail bool `json:"auto_merge_by_email"`
+ SyncGroupOnLogin bool `json:"sync_group_on_login"`
+ SyncRoleOnLogin bool `json:"sync_role_on_login"`
+ GroupMappingMode string `json:"group_mapping_mode"`
+ RoleMappingMode string `json:"role_mapping_mode"`
+ WellKnown string `json:"well_known"`
+ AuthStyle int `json:"auth_style"`
+ AccessPolicy string `json:"access_policy"`
+ AccessDeniedMessage string `json:"access_denied_message"`
}
type UserOAuthBindingResponse struct {
@@ -47,25 +73,71 @@ type UserOAuthBindingResponse struct {
}
func toCustomOAuthProviderResponse(p *model.CustomOAuthProvider) *CustomOAuthProviderResponse {
+ jwtSource := p.JWTSource
+ if strings.TrimSpace(jwtSource) == "" {
+ jwtSource = model.CustomJWTSourceQuery
+ }
+ authorizationServiceField := p.AuthorizationServiceField
+ if strings.TrimSpace(authorizationServiceField) == "" {
+ authorizationServiceField = "service"
+ }
+ ticketExchangeMethod := p.TicketExchangeMethod
+ if strings.TrimSpace(ticketExchangeMethod) == "" {
+ ticketExchangeMethod = model.CustomTicketExchangeMethodGET
+ }
+ ticketExchangePayloadMode := p.TicketExchangePayloadMode
+ if strings.TrimSpace(ticketExchangePayloadMode) == "" {
+ ticketExchangePayloadMode = model.CustomTicketExchangePayloadModeQuery
+ }
+ ticketExchangeTicketField := p.TicketExchangeTicketField
+ if strings.TrimSpace(ticketExchangeTicketField) == "" {
+ ticketExchangeTicketField = "ticket"
+ }
return &CustomOAuthProviderResponse{
- Id: p.Id,
- Name: p.Name,
- Slug: p.Slug,
- Icon: p.Icon,
- Enabled: p.Enabled,
- ClientId: p.ClientId,
- AuthorizationEndpoint: p.AuthorizationEndpoint,
- TokenEndpoint: p.TokenEndpoint,
- UserInfoEndpoint: p.UserInfoEndpoint,
- Scopes: p.Scopes,
- UserIdField: p.UserIdField,
- UsernameField: p.UsernameField,
- DisplayNameField: p.DisplayNameField,
- EmailField: p.EmailField,
- WellKnown: p.WellKnown,
- AuthStyle: p.AuthStyle,
- AccessPolicy: p.AccessPolicy,
- AccessDeniedMessage: p.AccessDeniedMessage,
+ Id: p.Id,
+ Name: p.Name,
+ Slug: p.Slug,
+ Icon: p.Icon,
+ Kind: p.GetKind(),
+ Enabled: p.Enabled,
+ ClientId: p.ClientId,
+ AuthorizationEndpoint: p.AuthorizationEndpoint,
+ TokenEndpoint: p.TokenEndpoint,
+ UserInfoEndpoint: p.UserInfoEndpoint,
+ Scopes: p.Scopes,
+ Issuer: p.Issuer,
+ Audience: p.Audience,
+ JwksURL: p.JwksURL,
+ PublicKey: p.PublicKey,
+ JWTSource: jwtSource,
+ JWTHeader: p.JWTHeader,
+ JWTIdentityMode: p.GetJWTIdentityMode(),
+ JWTAcquireMode: p.GetJWTAcquireMode(),
+ AuthorizationServiceField: authorizationServiceField,
+ TicketExchangeURL: p.TicketExchangeURL,
+ TicketExchangeMethod: ticketExchangeMethod,
+ TicketExchangePayloadMode: ticketExchangePayloadMode,
+ TicketExchangeTicketField: ticketExchangeTicketField,
+ TicketExchangeTokenField: p.TicketExchangeTokenField,
+ TicketExchangeServiceField: p.TicketExchangeServiceField,
+ UserIdField: p.UserIdField,
+ UsernameField: p.UsernameField,
+ DisplayNameField: p.DisplayNameField,
+ EmailField: p.EmailField,
+ GroupField: p.GroupField,
+ RoleField: p.RoleField,
+ GroupMapping: p.GroupMapping,
+ RoleMapping: p.RoleMapping,
+ AutoRegister: p.AutoRegister,
+ AutoMergeByEmail: p.AutoMergeByEmail,
+ SyncGroupOnLogin: p.SyncGroupOnLogin,
+ SyncRoleOnLogin: p.SyncRoleOnLogin,
+ GroupMappingMode: p.GroupMappingMode,
+ RoleMappingMode: p.RoleMappingMode,
+ WellKnown: p.WellKnown,
+ AuthStyle: p.AuthStyle,
+ AccessPolicy: p.AccessPolicy,
+ AccessDeniedMessage: p.AccessDeniedMessage,
}
}
@@ -113,24 +185,52 @@ func GetCustomOAuthProvider(c *gin.Context) {
// CreateCustomOAuthProviderRequest is the request structure for creating a custom OAuth provider
type CreateCustomOAuthProviderRequest struct {
- Name string `json:"name" binding:"required"`
- Slug string `json:"slug" binding:"required"`
- Icon string `json:"icon"`
- Enabled bool `json:"enabled"`
- ClientId string `json:"client_id" binding:"required"`
- ClientSecret string `json:"client_secret" binding:"required"`
- AuthorizationEndpoint string `json:"authorization_endpoint" binding:"required"`
- TokenEndpoint string `json:"token_endpoint" binding:"required"`
- UserInfoEndpoint string `json:"user_info_endpoint" binding:"required"`
- Scopes string `json:"scopes"`
- UserIdField string `json:"user_id_field"`
- UsernameField string `json:"username_field"`
- DisplayNameField string `json:"display_name_field"`
- EmailField string `json:"email_field"`
- WellKnown string `json:"well_known"`
- AuthStyle int `json:"auth_style"`
- AccessPolicy string `json:"access_policy"`
- AccessDeniedMessage string `json:"access_denied_message"`
+ Name string `json:"name" binding:"required"`
+ Slug string `json:"slug" binding:"required"`
+ Icon string `json:"icon"`
+ Kind string `json:"kind"`
+ Enabled bool `json:"enabled"`
+ ClientId string `json:"client_id"`
+ ClientSecret string `json:"client_secret"`
+ AuthorizationEndpoint string `json:"authorization_endpoint"`
+ TokenEndpoint string `json:"token_endpoint"`
+ UserInfoEndpoint string `json:"user_info_endpoint"`
+ Scopes string `json:"scopes"`
+ Issuer string `json:"issuer"`
+ Audience string `json:"audience"`
+ JwksURL string `json:"jwks_url"`
+ PublicKey string `json:"public_key"`
+ JWTSource string `json:"jwt_source"`
+ JWTHeader string `json:"jwt_header"`
+ JWTIdentityMode string `json:"jwt_identity_mode"`
+ JWTAcquireMode string `json:"jwt_acquire_mode"`
+ AuthorizationServiceField string `json:"authorization_service_field"`
+ TicketExchangeURL string `json:"ticket_exchange_url"`
+ TicketExchangeMethod string `json:"ticket_exchange_method"`
+ TicketExchangePayloadMode string `json:"ticket_exchange_payload_mode"`
+ TicketExchangeTicketField string `json:"ticket_exchange_ticket_field"`
+ TicketExchangeTokenField string `json:"ticket_exchange_token_field"`
+ TicketExchangeServiceField string `json:"ticket_exchange_service_field"`
+ TicketExchangeExtraParams string `json:"ticket_exchange_extra_params"`
+ TicketExchangeHeaders string `json:"ticket_exchange_headers"`
+ UserIdField string `json:"user_id_field"`
+ UsernameField string `json:"username_field"`
+ DisplayNameField string `json:"display_name_field"`
+ EmailField string `json:"email_field"`
+ GroupField string `json:"group_field"`
+ RoleField string `json:"role_field"`
+ GroupMapping string `json:"group_mapping"`
+ RoleMapping string `json:"role_mapping"`
+ AutoRegister bool `json:"auto_register"`
+ AutoMergeByEmail bool `json:"auto_merge_by_email"`
+ SyncGroupOnLogin bool `json:"sync_group_on_login"`
+ SyncRoleOnLogin bool `json:"sync_role_on_login"`
+ GroupMappingMode string `json:"group_mapping_mode"`
+ RoleMappingMode string `json:"role_mapping_mode"`
+ WellKnown string `json:"well_known"`
+ AuthStyle int `json:"auth_style"`
+ AccessPolicy string `json:"access_policy"`
+ AccessDeniedMessage string `json:"access_denied_message"`
}
type FetchCustomOAuthDiscoveryRequest struct {
@@ -231,24 +331,52 @@ func CreateCustomOAuthProvider(c *gin.Context) {
}
provider := &model.CustomOAuthProvider{
- Name: req.Name,
- Slug: req.Slug,
- Icon: req.Icon,
- Enabled: req.Enabled,
- ClientId: req.ClientId,
- ClientSecret: req.ClientSecret,
- AuthorizationEndpoint: req.AuthorizationEndpoint,
- TokenEndpoint: req.TokenEndpoint,
- UserInfoEndpoint: req.UserInfoEndpoint,
- Scopes: req.Scopes,
- UserIdField: req.UserIdField,
- UsernameField: req.UsernameField,
- DisplayNameField: req.DisplayNameField,
- EmailField: req.EmailField,
- WellKnown: req.WellKnown,
- AuthStyle: req.AuthStyle,
- AccessPolicy: req.AccessPolicy,
- AccessDeniedMessage: req.AccessDeniedMessage,
+ Name: req.Name,
+ Slug: req.Slug,
+ Icon: req.Icon,
+ Kind: req.Kind,
+ Enabled: req.Enabled,
+ ClientId: req.ClientId,
+ ClientSecret: req.ClientSecret,
+ AuthorizationEndpoint: req.AuthorizationEndpoint,
+ TokenEndpoint: req.TokenEndpoint,
+ UserInfoEndpoint: req.UserInfoEndpoint,
+ Scopes: req.Scopes,
+ Issuer: req.Issuer,
+ Audience: req.Audience,
+ JwksURL: req.JwksURL,
+ PublicKey: req.PublicKey,
+ JWTSource: req.JWTSource,
+ JWTHeader: req.JWTHeader,
+ JWTIdentityMode: req.JWTIdentityMode,
+ JWTAcquireMode: req.JWTAcquireMode,
+ AuthorizationServiceField: req.AuthorizationServiceField,
+ TicketExchangeURL: req.TicketExchangeURL,
+ TicketExchangeMethod: req.TicketExchangeMethod,
+ TicketExchangePayloadMode: req.TicketExchangePayloadMode,
+ TicketExchangeTicketField: req.TicketExchangeTicketField,
+ TicketExchangeTokenField: req.TicketExchangeTokenField,
+ TicketExchangeServiceField: req.TicketExchangeServiceField,
+ TicketExchangeExtraParams: req.TicketExchangeExtraParams,
+ TicketExchangeHeaders: req.TicketExchangeHeaders,
+ UserIdField: req.UserIdField,
+ UsernameField: req.UsernameField,
+ DisplayNameField: req.DisplayNameField,
+ EmailField: req.EmailField,
+ GroupField: req.GroupField,
+ RoleField: req.RoleField,
+ GroupMapping: req.GroupMapping,
+ RoleMapping: req.RoleMapping,
+ AutoRegister: req.AutoRegister,
+ AutoMergeByEmail: req.AutoMergeByEmail,
+ SyncGroupOnLogin: req.SyncGroupOnLogin,
+ SyncRoleOnLogin: req.SyncRoleOnLogin,
+ GroupMappingMode: req.GroupMappingMode,
+ RoleMappingMode: req.RoleMappingMode,
+ WellKnown: req.WellKnown,
+ AuthStyle: req.AuthStyle,
+ AccessPolicy: req.AccessPolicy,
+ AccessDeniedMessage: req.AccessDeniedMessage,
}
if err := model.CreateCustomOAuthProvider(provider); err != nil {
@@ -258,6 +386,7 @@ func CreateCustomOAuthProvider(c *gin.Context) {
// Register the provider in the OAuth registry
oauth.RegisterOrUpdateCustomProvider(provider)
+ invalidateCustomOAuthStatusCache()
c.JSON(http.StatusOK, gin.H{
"success": true,
@@ -268,24 +397,52 @@ func CreateCustomOAuthProvider(c *gin.Context) {
// UpdateCustomOAuthProviderRequest is the request structure for updating a custom OAuth provider
type UpdateCustomOAuthProviderRequest struct {
- Name string `json:"name"`
- Slug string `json:"slug"`
- Icon *string `json:"icon"` // Optional: if nil, keep existing
- Enabled *bool `json:"enabled"` // Optional: if nil, keep existing
- ClientId string `json:"client_id"`
- ClientSecret string `json:"client_secret"` // Optional: if empty, keep existing
- AuthorizationEndpoint string `json:"authorization_endpoint"`
- TokenEndpoint string `json:"token_endpoint"`
- UserInfoEndpoint string `json:"user_info_endpoint"`
- Scopes string `json:"scopes"`
- UserIdField string `json:"user_id_field"`
- UsernameField string `json:"username_field"`
- DisplayNameField string `json:"display_name_field"`
- 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
- AccessPolicy *string `json:"access_policy"` // Optional: if nil, keep existing
- AccessDeniedMessage *string `json:"access_denied_message"` // Optional: if nil, keep existing
+ Name *string `json:"name"`
+ Slug *string `json:"slug"`
+ Icon *string `json:"icon"` // Optional: if nil, keep existing
+ Enabled *bool `json:"enabled"` // Optional: if nil, keep existing
+ Kind *string `json:"kind"`
+ ClientId *string `json:"client_id"`
+ ClientSecret *string `json:"client_secret"`
+ AuthorizationEndpoint *string `json:"authorization_endpoint"`
+ TokenEndpoint *string `json:"token_endpoint"`
+ UserInfoEndpoint *string `json:"user_info_endpoint"`
+ Scopes *string `json:"scopes"`
+ Issuer *string `json:"issuer"`
+ Audience *string `json:"audience"`
+ JwksURL *string `json:"jwks_url"`
+ PublicKey *string `json:"public_key"`
+ JWTSource *string `json:"jwt_source"`
+ JWTHeader *string `json:"jwt_header"`
+ JWTIdentityMode *string `json:"jwt_identity_mode"`
+ JWTAcquireMode *string `json:"jwt_acquire_mode"`
+ AuthorizationServiceField *string `json:"authorization_service_field"`
+ TicketExchangeURL *string `json:"ticket_exchange_url"`
+ TicketExchangeMethod *string `json:"ticket_exchange_method"`
+ TicketExchangePayloadMode *string `json:"ticket_exchange_payload_mode"`
+ TicketExchangeTicketField *string `json:"ticket_exchange_ticket_field"`
+ TicketExchangeTokenField *string `json:"ticket_exchange_token_field"`
+ TicketExchangeServiceField *string `json:"ticket_exchange_service_field"`
+ TicketExchangeExtraParams *string `json:"ticket_exchange_extra_params"`
+ TicketExchangeHeaders *string `json:"ticket_exchange_headers"`
+ UserIdField *string `json:"user_id_field"`
+ UsernameField *string `json:"username_field"`
+ DisplayNameField *string `json:"display_name_field"`
+ EmailField *string `json:"email_field"`
+ GroupField *string `json:"group_field"`
+ RoleField *string `json:"role_field"`
+ GroupMapping *string `json:"group_mapping"`
+ RoleMapping *string `json:"role_mapping"`
+ AutoRegister *bool `json:"auto_register"`
+ AutoMergeByEmail *bool `json:"auto_merge_by_email"`
+ SyncGroupOnLogin *bool `json:"sync_group_on_login"`
+ SyncRoleOnLogin *bool `json:"sync_role_on_login"`
+ GroupMappingMode *string `json:"group_mapping_mode"`
+ RoleMappingMode *string `json:"role_mapping_mode"`
+ WellKnown *string `json:"well_known"` // 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
@@ -313,24 +470,24 @@ func UpdateCustomOAuthProvider(c *gin.Context) {
oldSlug := provider.Slug
// Check if new slug is taken by another provider
- if req.Slug != "" && req.Slug != provider.Slug {
- if model.IsSlugTaken(req.Slug, id) {
+ if req.Slug != nil && *req.Slug != provider.Slug {
+ if model.IsSlugTaken(*req.Slug, id) {
common.ApiErrorMsg(c, "该 Slug 已被使用")
return
}
// Check if slug conflicts with built-in providers
- if oauth.IsProviderRegistered(req.Slug) && !oauth.IsCustomProvider(req.Slug) {
+ if oauth.IsProviderRegistered(*req.Slug) && !oauth.IsCustomProvider(*req.Slug) {
common.ApiErrorMsg(c, "该 Slug 与内置 OAuth 提供商冲突")
return
}
}
// Update fields
- if req.Name != "" {
- provider.Name = req.Name
+ if req.Name != nil {
+ provider.Name = *req.Name
}
- if req.Slug != "" {
- provider.Slug = req.Slug
+ if req.Slug != nil {
+ provider.Slug = *req.Slug
}
if req.Icon != nil {
provider.Icon = *req.Icon
@@ -338,35 +495,119 @@ func UpdateCustomOAuthProvider(c *gin.Context) {
if req.Enabled != nil {
provider.Enabled = *req.Enabled
}
- if req.ClientId != "" {
- provider.ClientId = req.ClientId
+ if req.Kind != nil {
+ provider.Kind = *req.Kind
+ }
+ if req.ClientId != nil {
+ provider.ClientId = *req.ClientId
+ }
+ if req.ClientSecret != nil {
+ provider.ClientSecret = *req.ClientSecret
+ }
+ if req.AuthorizationEndpoint != nil {
+ provider.AuthorizationEndpoint = *req.AuthorizationEndpoint
+ }
+ if req.TokenEndpoint != nil {
+ provider.TokenEndpoint = *req.TokenEndpoint
+ }
+ if req.UserInfoEndpoint != nil {
+ provider.UserInfoEndpoint = *req.UserInfoEndpoint
+ }
+ if req.Scopes != nil {
+ provider.Scopes = *req.Scopes
+ }
+ if req.Issuer != nil {
+ provider.Issuer = *req.Issuer
+ }
+ if req.Audience != nil {
+ provider.Audience = *req.Audience
+ }
+ if req.JwksURL != nil {
+ provider.JwksURL = *req.JwksURL
+ }
+ if req.PublicKey != nil {
+ provider.PublicKey = *req.PublicKey
+ }
+ if req.JWTSource != nil {
+ provider.JWTSource = *req.JWTSource
+ }
+ if req.JWTHeader != nil {
+ provider.JWTHeader = *req.JWTHeader
+ }
+ if req.JWTIdentityMode != nil {
+ provider.JWTIdentityMode = *req.JWTIdentityMode
+ }
+ if req.JWTAcquireMode != nil {
+ provider.JWTAcquireMode = *req.JWTAcquireMode
+ }
+ if req.AuthorizationServiceField != nil {
+ provider.AuthorizationServiceField = *req.AuthorizationServiceField
+ }
+ if req.TicketExchangeURL != nil {
+ provider.TicketExchangeURL = *req.TicketExchangeURL
+ }
+ if req.TicketExchangeMethod != nil {
+ provider.TicketExchangeMethod = *req.TicketExchangeMethod
+ }
+ if req.TicketExchangePayloadMode != nil {
+ provider.TicketExchangePayloadMode = *req.TicketExchangePayloadMode
+ }
+ if req.TicketExchangeTicketField != nil {
+ provider.TicketExchangeTicketField = *req.TicketExchangeTicketField
+ }
+ if req.TicketExchangeTokenField != nil {
+ provider.TicketExchangeTokenField = *req.TicketExchangeTokenField
+ }
+ if req.TicketExchangeServiceField != nil {
+ provider.TicketExchangeServiceField = *req.TicketExchangeServiceField
+ }
+ if req.TicketExchangeExtraParams != nil {
+ provider.TicketExchangeExtraParams = *req.TicketExchangeExtraParams
+ }
+ if req.TicketExchangeHeaders != nil {
+ provider.TicketExchangeHeaders = *req.TicketExchangeHeaders
+ }
+ if req.UserIdField != nil {
+ provider.UserIdField = *req.UserIdField
+ }
+ if req.UsernameField != nil {
+ provider.UsernameField = *req.UsernameField
+ }
+ if req.DisplayNameField != nil {
+ provider.DisplayNameField = *req.DisplayNameField
+ }
+ if req.EmailField != nil {
+ provider.EmailField = *req.EmailField
+ }
+ if req.GroupField != nil {
+ provider.GroupField = *req.GroupField
}
- if req.ClientSecret != "" {
- provider.ClientSecret = req.ClientSecret
+ if req.RoleField != nil {
+ provider.RoleField = *req.RoleField
}
- if req.AuthorizationEndpoint != "" {
- provider.AuthorizationEndpoint = req.AuthorizationEndpoint
+ if req.GroupMapping != nil {
+ provider.GroupMapping = *req.GroupMapping
}
- if req.TokenEndpoint != "" {
- provider.TokenEndpoint = req.TokenEndpoint
+ if req.RoleMapping != nil {
+ provider.RoleMapping = *req.RoleMapping
}
- if req.UserInfoEndpoint != "" {
- provider.UserInfoEndpoint = req.UserInfoEndpoint
+ if req.AutoRegister != nil {
+ provider.AutoRegister = *req.AutoRegister
}
- if req.Scopes != "" {
- provider.Scopes = req.Scopes
+ if req.AutoMergeByEmail != nil {
+ provider.AutoMergeByEmail = *req.AutoMergeByEmail
}
- if req.UserIdField != "" {
- provider.UserIdField = req.UserIdField
+ if req.SyncGroupOnLogin != nil {
+ provider.SyncGroupOnLogin = *req.SyncGroupOnLogin
}
- if req.UsernameField != "" {
- provider.UsernameField = req.UsernameField
+ if req.SyncRoleOnLogin != nil {
+ provider.SyncRoleOnLogin = *req.SyncRoleOnLogin
}
- if req.DisplayNameField != "" {
- provider.DisplayNameField = req.DisplayNameField
+ if req.GroupMappingMode != nil {
+ provider.GroupMappingMode = *req.GroupMappingMode
}
- if req.EmailField != "" {
- provider.EmailField = req.EmailField
+ if req.RoleMappingMode != nil {
+ provider.RoleMappingMode = *req.RoleMappingMode
}
if req.WellKnown != nil {
provider.WellKnown = *req.WellKnown
@@ -391,6 +632,7 @@ func UpdateCustomOAuthProvider(c *gin.Context) {
oauth.UnregisterCustomProvider(oldSlug)
}
oauth.RegisterOrUpdateCustomProvider(provider)
+ invalidateCustomOAuthStatusCache()
c.JSON(http.StatusOK, gin.H{
"success": true,
@@ -434,6 +676,7 @@ func DeleteCustomOAuthProvider(c *gin.Context) {
// Unregister the provider from the OAuth registry
oauth.UnregisterCustomProvider(provider.Slug)
+ invalidateCustomOAuthStatusCache()
c.JSON(http.StatusOK, gin.H{
"success": true,
diff --git a/controller/custom_oauth_jwt.go b/controller/custom_oauth_jwt.go
new file mode 100644
index 000000000000..d5cb72ef7dd3
--- /dev/null
+++ b/controller/custom_oauth_jwt.go
@@ -0,0 +1,389 @@
+package controller
+
+import (
+ "fmt"
+ "net/url"
+ "strings"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/i18n"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/QuantumNous/new-api/oauth"
+ "github.com/QuantumNous/new-api/setting/system_setting"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-gonic/gin"
+)
+
+type customOAuthJWTLoginRequest struct {
+ State string `json:"state" form:"state"`
+ Token string `json:"token" form:"token"`
+ IDToken string `json:"id_token" form:"id_token"`
+ JWT string `json:"jwt" form:"jwt"`
+ Ticket string `json:"ticket" form:"ticket"`
+}
+
+type customOAuthJWTLoginResult struct {
+ Action string
+ User *model.User
+ BindAfterStatusCheck bool
+ ProviderUserID string
+ AutoRegisterTriggered bool
+ EmailMergeTriggered bool
+ GroupResult string
+ RoleResult int
+}
+
+type customOAuthJWTAuditInfo struct {
+ ProviderSlug string
+ ProviderKind string
+ ExternalID string
+ TargetUserID int
+ Action string
+ AutoRegisterTriggered bool
+ EmailMergeTriggered bool
+ GroupResult string
+ RoleResult string
+ FailureReason string
+}
+
+func HandleCustomOAuthJWTLogin(c *gin.Context) {
+ providerConfig, provider := loadCustomJWTDirectProvider(c)
+ if provider == nil {
+ return
+ }
+ audit := newCustomOAuthJWTAuditInfo(providerConfig)
+
+ var req customOAuthJWTLoginRequest
+ if err := c.ShouldBind(&req); err != nil {
+ audit.FailureReason = "invalid_request"
+ recordCustomOAuthJWTAudit(audit)
+ common.ApiErrorI18n(c, i18n.MsgInvalidParams)
+ return
+ }
+
+ session := sessions.Default(c)
+ state := strings.TrimSpace(req.State)
+ sessionState, ok := session.Get("oauth_state").(string)
+ if state == "" || !ok || strings.TrimSpace(sessionState) == "" || state != sessionState {
+ audit.FailureReason = "invalid_state"
+ recordCustomOAuthJWTAudit(audit)
+ common.ApiErrorI18n(c, i18n.MsgOAuthStateInvalid)
+ return
+ }
+
+ result, audit, err := completeCustomOAuthJWTLogin(
+ c,
+ providerConfig,
+ provider,
+ session,
+ sessionState,
+ selectJWTLoginCredential(providerConfig, req),
+ req.Ticket,
+ audit,
+ )
+ if err != nil {
+ if audit != nil && audit.FailureReason == "" {
+ audit.FailureReason = oauthAuditFailureReason(err)
+ }
+ recordCustomOAuthJWTAudit(audit)
+ handleCustomOAuthJWTLoginError(c, err)
+ return
+ }
+
+ if result.Action == "bind" {
+ recordCustomOAuthJWTAudit(audit)
+ common.ApiSuccessI18n(c, i18n.MsgOAuthBindSuccess, gin.H{
+ "action": "bind",
+ })
+ return
+ }
+
+ if result.User.Status != common.UserStatusEnabled {
+ audit.FailureReason = "user_disabled"
+ recordCustomOAuthJWTAudit(audit)
+ common.ApiErrorI18n(c, i18n.MsgOAuthUserBanned)
+ return
+ }
+ if result.BindAfterStatusCheck {
+ if err := bindOAuthIdentityToUser(result.User, provider, result.ProviderUserID); err != nil {
+ audit.FailureReason = oauthAuditFailureReason(err)
+ recordCustomOAuthJWTAudit(audit)
+ handleCustomOAuthJWTLoginError(c, err)
+ return
+ }
+ }
+
+ if !setupLoginWithResult(result.User, c) {
+ audit.FailureReason = "session_save_failed"
+ recordCustomOAuthJWTAudit(audit)
+ return
+ }
+ recordCustomOAuthJWTAudit(audit)
+}
+
+func loadCustomJWTDirectProvider(c *gin.Context) (*model.CustomOAuthProvider, *oauth.JWTDirectProvider) {
+ providerName := c.Param("provider")
+ providerConfig, err := model.GetCustomOAuthProviderBySlug(providerName)
+ if err != nil || providerConfig == nil || !providerConfig.IsJWTDirect() {
+ common.ApiErrorI18n(c, i18n.MsgOAuthUnknownProvider)
+ return nil, nil
+ }
+ if !providerConfig.Enabled {
+ common.ApiErrorI18n(c, i18n.MsgOAuthNotEnabled, providerParams(providerConfig.Name))
+ return nil, nil
+ }
+ return providerConfig, oauth.NewJWTDirectProvider(providerConfig)
+}
+
+func completeCustomOAuthJWTLogin(
+ c *gin.Context,
+ providerConfig *model.CustomOAuthProvider,
+ provider *oauth.JWTDirectProvider,
+ session sessions.Session,
+ state string,
+ rawToken string,
+ ticket string,
+ audit *customOAuthJWTAuditInfo,
+) (*customOAuthJWTLoginResult, *customOAuthJWTAuditInfo, error) {
+ rawToken = strings.TrimSpace(rawToken)
+ ticket = strings.TrimSpace(ticket)
+ if providerConfig.RequiresTicketAcquire() {
+ if ticket == "" {
+ if audit != nil {
+ audit.FailureReason = "missing_exchange_ticket"
+ }
+ return nil, audit, oauth.NewOAuthError(i18n.MsgOAuthTicketMissing, nil)
+ }
+ } else if rawToken == "" {
+ if audit != nil {
+ audit.FailureReason = "missing_jwt_token"
+ }
+ return nil, audit, oauth.NewOAuthError(i18n.MsgOAuthJWTMissing, nil)
+ }
+
+ callbackURL := ""
+ if providerConfig.RequiresTicketAcquire() {
+ validatedCallbackURL, callbackErr := buildCustomOAuthJWTCallbackURL(providerConfig.Slug, state)
+ if callbackErr != nil {
+ if audit != nil {
+ audit.FailureReason = "invalid_callback_url"
+ }
+ return nil, audit, oauth.NewOAuthError(i18n.MsgOAuthTokenFailed, map[string]any{"Provider": providerConfig.Name})
+ }
+ callbackURL = validatedCallbackURL
+ }
+
+ identity, err := provider.ResolveIdentityFromInput(
+ c.Request.Context(),
+ rawToken,
+ ticket,
+ callbackURL,
+ state,
+ )
+ if err != nil {
+ return nil, audit, err
+ }
+ if audit != nil {
+ audit.ExternalID = redactOAuthAuditID(identity.User.ProviderUserID)
+ audit.GroupResult = safeOAuthAuditValue(identity.Group)
+ audit.RoleResult = oauthRoleLabel(identity.Role)
+ }
+
+ if session.Get("username") != nil {
+ if audit != nil {
+ if sessionUserID, ok := session.Get("id").(int); ok {
+ audit.TargetUserID = sessionUserID
+ }
+ }
+ currentUser, currentUserErr := getSessionUser(c)
+ if currentUserErr != nil {
+ if audit != nil {
+ if strings.TrimSpace(currentUserErr.Error()) == "该用户已被禁用" {
+ audit.FailureReason = "user_disabled"
+ } else {
+ audit.FailureReason = oauthAuditFailureReason(currentUserErr)
+ }
+ }
+ if strings.TrimSpace(currentUserErr.Error()) == "该用户已被禁用" {
+ return nil, audit, oauth.NewOAuthError(i18n.MsgOAuthUserBanned, nil)
+ }
+ return nil, audit, currentUserErr
+ }
+ if currentUser.Status != common.UserStatusEnabled {
+ if audit != nil {
+ audit.FailureReason = "user_disabled"
+ }
+ return nil, audit, oauth.NewOAuthError(i18n.MsgOAuthUserBanned, nil)
+ }
+ if err := bindOAuthIdentityToCurrentUser(c, provider, identity.User); err != nil {
+ return nil, audit, err
+ }
+ if audit != nil {
+ audit.Action = "bind"
+ }
+ return &customOAuthJWTLoginResult{Action: "bind"}, audit, nil
+ }
+
+ resolvedUser, err := findOrCreateOAuthUserWithOptions(c, provider, identity.User, session, oauthFindOrCreateOptions{
+ AllowAutoRegister: providerConfig.AutoRegister,
+ AllowAutoMergeByEmail: providerConfig.AutoMergeByEmail,
+ InitialRole: identity.Role,
+ InitialGroup: identity.Group,
+ })
+ if err != nil {
+ return nil, audit, err
+ }
+
+ if resolvedUser.User.Status == common.UserStatusEnabled {
+ if err := syncOAuthUserLoginAttributes(
+ resolvedUser.User,
+ providerConfig.Name,
+ identity.Group,
+ providerConfig.SyncGroupOnLogin,
+ identity.Role,
+ providerConfig.SyncRoleOnLogin,
+ ); err != nil {
+ return nil, audit, err
+ }
+ }
+
+ result := &customOAuthJWTLoginResult{
+ Action: "login",
+ User: resolvedUser.User,
+ BindAfterStatusCheck: resolvedUser.BindAfterStatusCheck,
+ ProviderUserID: identity.User.ProviderUserID,
+ AutoRegisterTriggered: resolvedUser.AutoRegisterTriggered,
+ EmailMergeTriggered: resolvedUser.EmailMergeTriggered,
+ GroupResult: identity.Group,
+ RoleResult: identity.Role,
+ }
+ if audit != nil {
+ audit.Action = result.Action
+ audit.TargetUserID = result.User.Id
+ audit.AutoRegisterTriggered = result.AutoRegisterTriggered
+ audit.EmailMergeTriggered = result.EmailMergeTriggered
+ }
+ return result, audit, nil
+}
+
+func buildCustomOAuthJWTCallbackURL(providerSlug string, state string) (string, error) {
+ baseURL := strings.TrimSpace(system_setting.ServerAddress)
+ if baseURL == "" {
+ return "", fmt.Errorf("server address is empty")
+ }
+ callbackURL, err := url.Parse(baseURL)
+ if err != nil {
+ return "", fmt.Errorf("invalid server address: %w", err)
+ }
+ if callbackURL == nil || strings.TrimSpace(callbackURL.Host) == "" {
+ return "", fmt.Errorf("server address host is empty")
+ }
+ if callbackURL.Scheme != "http" && callbackURL.Scheme != "https" {
+ return "", fmt.Errorf("server address scheme must be http or https")
+ }
+
+ callbackURL.RawQuery = ""
+ callbackURL.Fragment = ""
+ callbackURL.Path = strings.TrimRight(callbackURL.Path, "/") + "/oauth/" + providerSlug
+ if strings.TrimSpace(state) != "" {
+ query := callbackURL.Query()
+ query.Set("state", state)
+ callbackURL.RawQuery = query.Encode()
+ }
+ return callbackURL.String(), nil
+}
+
+func handleCustomOAuthJWTLoginError(c *gin.Context, err error) {
+ if boundErr, ok := err.(*OAuthAlreadyBoundError); ok {
+ common.ApiErrorI18n(c, i18n.MsgOAuthAlreadyBound, providerParams(boundErr.Provider))
+ return
+ }
+ switch err.(type) {
+ case *oauth.OAuthError, *oauth.AccessDeniedError, *oauth.TrustLevelError:
+ handleOAuthError(c, err)
+ default:
+ handleOAuthUserError(c, err)
+ }
+}
+
+func firstNonEmpty(values ...string) string {
+ for _, value := range values {
+ trimmed := strings.TrimSpace(value)
+ if trimmed != "" {
+ return trimmed
+ }
+ }
+ return ""
+}
+
+func selectJWTLoginCredential(providerConfig *model.CustomOAuthProvider, req customOAuthJWTLoginRequest) string {
+ if providerConfig != nil && providerConfig.GetJWTIdentityMode() == model.CustomJWTIdentityModeUserInfo {
+ return firstNonEmpty(req.Token, req.IDToken, req.JWT)
+ }
+ return firstNonEmpty(req.IDToken, req.JWT, req.Token)
+}
+
+func oauthAuditFailureReason(err error) string {
+ if err == nil {
+ return ""
+ }
+
+ switch e := err.(type) {
+ case *oauth.OAuthError:
+ return "oauth_error:" + safeOAuthAuditValue(e.MsgKey)
+ case *oauth.AccessDeniedError:
+ return "access_denied"
+ case *oauth.TrustLevelError:
+ return "trust_level_denied"
+ case *OAuthAlreadyBoundError:
+ return "oauth_already_bound"
+ case *OAuthUserDeletedError:
+ return "oauth_user_deleted"
+ case *OAuthRegistrationDisabledError, *OAuthAutoRegisterDisabledError:
+ return "registration_disabled"
+ default:
+ return "internal_error"
+ }
+}
+
+func newCustomOAuthJWTAuditInfo(providerConfig *model.CustomOAuthProvider) *customOAuthJWTAuditInfo {
+ if providerConfig == nil {
+ return &customOAuthJWTAuditInfo{}
+ }
+ return &customOAuthJWTAuditInfo{
+ ProviderSlug: providerConfig.Slug,
+ ProviderKind: providerConfig.GetKind(),
+ }
+}
+
+func recordCustomOAuthJWTAudit(audit *customOAuthJWTAuditInfo) {
+ if audit == nil {
+ return
+ }
+
+ content := fmt.Sprintf(
+ "企业认证审计 provider_slug=%s provider_kind=%s action=%s external_id=%s target_user_id=%d auto_register=%t email_merge=%t group_result=%s role_result=%s failure_reason=%s",
+ safeOAuthAuditValue(audit.ProviderSlug),
+ safeOAuthAuditValue(audit.ProviderKind),
+ safeOAuthAuditValue(audit.Action),
+ safeOAuthAuditValue(audit.ExternalID),
+ audit.TargetUserID,
+ audit.AutoRegisterTriggered,
+ audit.EmailMergeTriggered,
+ safeOAuthAuditValue(audit.GroupResult),
+ safeOAuthAuditValue(audit.RoleResult),
+ safeOAuthAuditValue(audit.FailureReason),
+ )
+ common.SysLog("[EnterpriseAuth] " + content)
+ if audit.TargetUserID > 0 {
+ model.RecordLog(audit.TargetUserID, model.LogTypeSystem, content)
+ }
+}
+
+func redactOAuthAuditID(value string) string {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ return ""
+ }
+ return "hmac_sha256:" + common.GenerateHMAC(value)
+}
diff --git a/controller/custom_oauth_jwt_test.go b/controller/custom_oauth_jwt_test.go
new file mode 100644
index 000000000000..0769fa820f67
--- /dev/null
+++ b/controller/custom_oauth_jwt_test.go
@@ -0,0 +1,1414 @@
+package controller
+
+import (
+ "bytes"
+ "crypto/rand"
+ "crypto/rsa"
+ "crypto/x509"
+ "encoding/pem"
+ "fmt"
+ "net/http"
+ "net/http/cookiejar"
+ "net/http/httptest"
+ "strconv"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/QuantumNous/new-api/oauth"
+ "github.com/QuantumNous/new-api/setting/system_setting"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-contrib/sessions/cookie"
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "github.com/golang-jwt/jwt/v5"
+ "gorm.io/gorm"
+)
+
+type oauthJWTAPIResponse struct {
+ Success bool `json:"success"`
+ Message string `json:"message"`
+ Data common.RawMessage `json:"data"`
+}
+
+type oauthJWTLoginResponse struct {
+ ID int `json:"id"`
+ Username string `json:"username"`
+ Role int `json:"role"`
+ Group string `json:"group"`
+}
+
+type oauthJWTBindResponse struct {
+ Action string `json:"action"`
+}
+
+func setupCustomOAuthJWTControllerTestDB(t *testing.T) {
+ t.Helper()
+
+ prevDB := model.DB
+ prevLogDB := model.LOG_DB
+ prevUsingSQLite := common.UsingSQLite
+ prevUsingMySQL := common.UsingMySQL
+ prevUsingPostgreSQL := common.UsingPostgreSQL
+ prevRedisEnabled := common.RedisEnabled
+ prevRegisterEnabled := common.RegisterEnabled
+ prevQuotaForNewUser := common.QuotaForNewUser
+ prevQuotaForInvitee := common.QuotaForInvitee
+ prevQuotaForInviter := common.QuotaForInviter
+
+ gin.SetMode(gin.TestMode)
+ common.UsingSQLite = true
+ common.UsingMySQL = false
+ common.UsingPostgreSQL = false
+ common.RedisEnabled = false
+ common.RegisterEnabled = true
+ common.QuotaForNewUser = 0
+ common.QuotaForInvitee = 0
+ common.QuotaForInviter = 0
+
+ dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_"))
+ db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("failed to open sqlite db: %v", err)
+ }
+ model.DB = db
+ model.LOG_DB = db
+ if err := db.AutoMigrate(&model.User{}, &model.Log{}, &model.CustomOAuthProvider{}, &model.UserOAuthBinding{}); err != nil {
+ t.Fatalf("failed to migrate test tables: %v", err)
+ }
+
+ t.Cleanup(func() {
+ sqlDB, err := db.DB()
+ if err == nil {
+ _ = sqlDB.Close()
+ }
+ model.DB = prevDB
+ model.LOG_DB = prevLogDB
+ common.UsingSQLite = prevUsingSQLite
+ common.UsingMySQL = prevUsingMySQL
+ common.UsingPostgreSQL = prevUsingPostgreSQL
+ common.RedisEnabled = prevRedisEnabled
+ common.RegisterEnabled = prevRegisterEnabled
+ common.QuotaForNewUser = prevQuotaForNewUser
+ common.QuotaForInvitee = prevQuotaForInvitee
+ common.QuotaForInviter = prevQuotaForInviter
+ })
+}
+
+func newCustomOAuthJWTRouter(t *testing.T) *gin.Engine {
+ t.Helper()
+ router := gin.New()
+ store := cookie.NewStore([]byte("test-session-secret"))
+ router.Use(sessions.Sessions("session", store))
+ router.GET("/api/oauth/state", GenerateOAuthCode)
+ router.POST("/api/auth/external/:provider/jwt/login", HandleCustomOAuthJWTLogin)
+ router.GET("/test/login-as/:id", func(c *gin.Context) {
+ var user model.User
+ if err := model.DB.First(&user, c.Param("id")).Error; err != nil {
+ c.JSON(http.StatusNotFound, gin.H{"success": false, "message": err.Error()})
+ return
+ }
+ session := sessions.Default(c)
+ session.Set("id", user.Id)
+ session.Set("username", user.Username)
+ session.Set("role", user.Role)
+ session.Set("status", user.Status)
+ session.Set("group", user.Group)
+ if err := session.Save(); err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"success": true})
+ })
+ return router
+}
+
+type jwtDirectProviderTestOptions struct {
+ AutoRegister bool
+ AutoMergeByEmail bool
+ SyncGroupOnLogin bool
+ SyncRoleOnLogin bool
+ JWTIdentityMode string
+ UserInfoEndpoint string
+ JWTHeader string
+ JWTAcquireMode string
+ TicketExchangeURL string
+ TicketExchangeMethod string
+ TicketExchangePayloadMode string
+ TicketExchangeTicketField string
+ TicketExchangeTokenField string
+ TicketExchangeServiceField string
+ TicketExchangeExtraParams string
+ TicketExchangeHeaders string
+}
+
+func createJWTDirectProviderForTest(t *testing.T, privateKey *rsa.PrivateKey, options jwtDirectProviderTestOptions) *model.CustomOAuthProvider {
+ t.Helper()
+ provider := &model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ Enabled: true,
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ ClientId: "new-api-client",
+ Scopes: "openid profile email",
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeControllerRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ JWTIdentityMode: options.JWTIdentityMode,
+ UserInfoEndpoint: options.UserInfoEndpoint,
+ JWTHeader: options.JWTHeader,
+ UserIdField: "sub",
+ UsernameField: "preferred_username",
+ DisplayNameField: "name",
+ EmailField: "email",
+ GroupField: "groups",
+ GroupMapping: `{"engineering":"vip"}`,
+ RoleField: "roles",
+ RoleMapping: `{"platform-admin":"admin"}`,
+ AutoRegister: options.AutoRegister,
+ AutoMergeByEmail: options.AutoMergeByEmail,
+ SyncGroupOnLogin: options.SyncGroupOnLogin,
+ SyncRoleOnLogin: options.SyncRoleOnLogin,
+ JWTAcquireMode: options.JWTAcquireMode,
+ TicketExchangeURL: options.TicketExchangeURL,
+ TicketExchangeMethod: options.TicketExchangeMethod,
+ TicketExchangePayloadMode: options.TicketExchangePayloadMode,
+ TicketExchangeTicketField: options.TicketExchangeTicketField,
+ TicketExchangeTokenField: options.TicketExchangeTokenField,
+ TicketExchangeServiceField: options.TicketExchangeServiceField,
+ TicketExchangeExtraParams: options.TicketExchangeExtraParams,
+ TicketExchangeHeaders: options.TicketExchangeHeaders,
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+ return provider
+}
+
+func createUserForBindTest(t *testing.T, username string) *model.User {
+ t.Helper()
+ password, err := common.Password2Hash("12345678")
+ if err != nil {
+ t.Fatalf("failed to hash password: %v", err)
+ }
+ user := &model.User{
+ Username: username,
+ Password: password,
+ DisplayName: username,
+ Role: common.RoleCommonUser,
+ Status: common.UserStatusEnabled,
+ Group: "default",
+ AffCode: username + "-aff",
+ }
+ if err := model.DB.Create(user).Error; err != nil {
+ t.Fatalf("failed to create user: %v", err)
+ }
+ return user
+}
+
+func createUserWithEmailForTest(t *testing.T, username string, email string) *model.User {
+ t.Helper()
+ password, err := common.Password2Hash("12345678")
+ if err != nil {
+ t.Fatalf("failed to hash password: %v", err)
+ }
+ user := &model.User{
+ Username: username,
+ Password: password,
+ DisplayName: username,
+ Email: email,
+ Role: common.RoleCommonUser,
+ Status: common.UserStatusEnabled,
+ Group: "default",
+ AffCode: username + "-aff",
+ }
+ if err := model.DB.Create(user).Error; err != nil {
+ t.Fatalf("failed to create user: %v", err)
+ }
+ return user
+}
+
+func getLatestSystemLogForUser(t *testing.T, userID int) *model.Log {
+ t.Helper()
+ var log model.Log
+ if err := model.LOG_DB.Where("user_id = ? AND type = ?", userID, model.LogTypeSystem).Order("created_at desc").First(&log).Error; err != nil {
+ t.Fatalf("failed to load latest system log for user %d: %v", userID, err)
+ }
+ return &log
+}
+
+func TestHandleCustomOAuthJWTLoginCreatesUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ })
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-user-1",
+ "preferred_username": "alice",
+ "name": "Alice",
+ "email": "alice@example.com",
+ "groups": []string{"engineering"},
+ "roles": []string{"platform-admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if !response.Success {
+ t.Fatalf("expected success response, got message: %s", response.Message)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.Role != common.RoleAdminUser {
+ t.Fatalf("expected admin role, got %d", loginData.Role)
+ }
+ if loginData.Group != "vip" {
+ t.Fatalf("expected mapped group vip, got %s", loginData.Group)
+ }
+
+ var user model.User
+ if err := model.DB.Where("username = ?", "alice").First(&user).Error; err != nil {
+ t.Fatalf("expected created user alice, got error: %v", err)
+ }
+ if user.Role != common.RoleAdminUser || user.Group != "vip" {
+ t.Fatalf("unexpected persisted user role/group: role=%d group=%s", user.Role, user.Group)
+ }
+ if !model.IsProviderUserIdTaken(provider.Id, "ext-user-1") {
+ t.Fatal("expected oauth binding to be created")
+ }
+ log := getLatestSystemLogForUser(t, user.Id)
+ if !strings.Contains(log.Content, "provider_slug=acme-sso") ||
+ !strings.Contains(log.Content, "provider_kind=jwt_direct") ||
+ !strings.Contains(log.Content, "action=login") ||
+ !strings.Contains(log.Content, "external_id="+redactOAuthAuditID("ext-user-1")) ||
+ !strings.Contains(log.Content, "auto_register=true") ||
+ !strings.Contains(log.Content, "email_merge=false") ||
+ !strings.Contains(log.Content, "group_result=vip") ||
+ !strings.Contains(log.Content, "role_result=admin") {
+ t.Fatalf("unexpected enterprise auth audit log: %s", log.Content)
+ }
+ if strings.Contains(log.Content, "external_id=ext-user-1") {
+ t.Fatalf("expected enterprise auth audit log to redact external id, got %s", log.Content)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginRejectsWhenAutoRegisterDisabled(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: false,
+ })
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-user-2",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if response.Success {
+ t.Fatal("expected auto-register disabled login to fail")
+ }
+
+ var count int64
+ if err := model.DB.Model(&model.User{}).Count(&count).Error; err != nil {
+ t.Fatalf("failed to count users: %v", err)
+ }
+ if count != 0 {
+ t.Fatalf("expected no users to be created, got %d", count)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginBindsExistingSessionUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ })
+ user := createUserForBindTest(t, "bind-user")
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ loginReq, err := http.NewRequest(http.MethodGet, server.URL+"/test/login-as/"+strconv.Itoa(user.Id), nil)
+ if err != nil {
+ t.Fatalf("failed to build login-as request: %v", err)
+ }
+ loginResp, err := client.Do(loginReq)
+ if err != nil {
+ t.Fatalf("failed to establish session: %v", err)
+ }
+ _ = loginResp.Body.Close()
+
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-bind-1",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if !response.Success {
+ t.Fatalf("expected bind response success, got message: %s", response.Message)
+ }
+
+ var bindData oauthJWTBindResponse
+ if err := common.Unmarshal(response.Data, &bindData); err != nil {
+ t.Fatalf("failed to decode bind response: %v", err)
+ }
+ if bindData.Action != "bind" {
+ t.Fatalf("expected bind action, got %s", bindData.Action)
+ }
+ if !model.IsProviderUserIdTaken(provider.Id, "ext-bind-1") {
+ t.Fatal("expected binding to be created for current user")
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginDoesNotSyncAttributesDuringBind(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ SyncGroupOnLogin: true,
+ SyncRoleOnLogin: true,
+ })
+ user := createUserForBindTest(t, "bind-no-sync-user")
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ loginReq, err := http.NewRequest(http.MethodGet, server.URL+"/test/login-as/"+strconv.Itoa(user.Id), nil)
+ if err != nil {
+ t.Fatalf("failed to build login-as request: %v", err)
+ }
+ loginResp, err := client.Do(loginReq)
+ if err != nil {
+ t.Fatalf("failed to establish session: %v", err)
+ }
+ _ = loginResp.Body.Close()
+
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-bind-sync-ignored",
+ "groups": []string{"engineering"},
+ "roles": []string{"platform-admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if !response.Success {
+ t.Fatalf("expected bind response success, got message: %s", response.Message)
+ }
+
+ reloadedUser, err := model.GetUserById(user.Id, false)
+ if err != nil {
+ t.Fatalf("failed to reload bound user: %v", err)
+ }
+ if reloadedUser.Role != common.RoleCommonUser || reloadedUser.Group != "default" {
+ t.Fatalf("expected bind flow to keep local attributes unchanged, got role=%d group=%s", reloadedUser.Role, reloadedUser.Group)
+ }
+ if !model.IsProviderUserIdTaken(provider.Id, "ext-bind-sync-ignored") {
+ t.Fatal("expected binding to be created during bind flow")
+ }
+ log := getLatestSystemLogForUser(t, user.Id)
+ if !strings.Contains(log.Content, "action=bind") ||
+ !strings.Contains(log.Content, "external_id="+redactOAuthAuditID("ext-bind-sync-ignored")) {
+ t.Fatalf("expected bind audit log, got %s", log.Content)
+ }
+ if strings.Contains(log.Content, "external_id=ext-bind-sync-ignored") {
+ t.Fatalf("expected bind audit log to redact external id, got %s", log.Content)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginMergesByEmailWhenEnabled(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: false,
+ AutoMergeByEmail: true,
+ })
+ existingUser := createUserWithEmailForTest(t, "merged-user", "alice@example.com")
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-user-merge",
+ "preferred_username": "alice",
+ "email": "alice@example.com",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if !response.Success {
+ t.Fatalf("expected merge login to succeed, got message: %s", response.Message)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.ID != existingUser.Id {
+ t.Fatalf("expected merged existing user id %d, got %d", existingUser.Id, loginData.ID)
+ }
+
+ var count int64
+ if err := model.DB.Model(&model.User{}).Count(&count).Error; err != nil {
+ t.Fatalf("failed to count users: %v", err)
+ }
+ if count != 1 {
+ t.Fatalf("expected no new user to be created, got %d users", count)
+ }
+ if !model.IsProviderUserIdTaken(provider.Id, "ext-user-merge") {
+ t.Fatal("expected oauth binding to be created for merged user")
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginSyncsExistingBoundUserOnLogin(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ SyncGroupOnLogin: true,
+ SyncRoleOnLogin: true,
+ })
+ user := createUserForBindTest(t, "existing-bound-user")
+ if err := model.CreateUserOAuthBinding(&model.UserOAuthBinding{
+ UserId: user.Id,
+ ProviderId: provider.Id,
+ ProviderUserId: "ext-sync-existing",
+ }); err != nil {
+ t.Fatalf("failed to seed oauth binding: %v", err)
+ }
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-sync-existing",
+ "preferred_username": "existing-bound-user",
+ "groups": []string{"engineering"},
+ "roles": []string{"platform-admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if !response.Success {
+ t.Fatalf("expected existing bound login to succeed, got message: %s", response.Message)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.Role != common.RoleAdminUser || loginData.Group != "vip" {
+ t.Fatalf("expected synced login response, got role=%d group=%s", loginData.Role, loginData.Group)
+ }
+
+ reloadedUser, err := model.GetUserById(user.Id, false)
+ if err != nil {
+ t.Fatalf("failed to reload synced user: %v", err)
+ }
+ if reloadedUser.Role != common.RoleAdminUser || reloadedUser.Group != "vip" {
+ t.Fatalf("expected synced persisted user, got role=%d group=%s", reloadedUser.Role, reloadedUser.Group)
+ }
+ if !strings.Contains(reloadedUser.GetSetting().SidebarModules, "\"admin\"") {
+ t.Fatalf("expected admin sidebar section to be added after role promotion, got %s", reloadedUser.GetSetting().SidebarModules)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginSyncsMergedUserOnLogin(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: false,
+ AutoMergeByEmail: true,
+ SyncGroupOnLogin: true,
+ SyncRoleOnLogin: true,
+ })
+ existingUser := createUserWithEmailForTest(t, "merged-sync-user", "merged-sync@example.com")
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-merge-sync",
+ "preferred_username": "merged-sync-user",
+ "email": "merged-sync@example.com",
+ "groups": []string{"engineering"},
+ "roles": []string{"platform-admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if !response.Success {
+ t.Fatalf("expected merged sync login to succeed, got message: %s", response.Message)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.ID != existingUser.Id || loginData.Role != common.RoleAdminUser || loginData.Group != "vip" {
+ t.Fatalf("expected merged sync response for user %d, got id=%d role=%d group=%s", existingUser.Id, loginData.ID, loginData.Role, loginData.Group)
+ }
+
+ reloadedUser, err := model.GetUserById(existingUser.Id, false)
+ if err != nil {
+ t.Fatalf("failed to reload merged synced user: %v", err)
+ }
+ if reloadedUser.Role != common.RoleAdminUser || reloadedUser.Group != "vip" {
+ t.Fatalf("expected merged synced persisted user, got role=%d group=%s", reloadedUser.Role, reloadedUser.Group)
+ }
+ if !model.IsProviderUserIdTaken(provider.Id, "ext-merge-sync") {
+ t.Fatal("expected merged user oauth binding to be created")
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginRejectsEmailMergeConflict(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: false,
+ AutoMergeByEmail: true,
+ })
+ createUserWithEmailForTest(t, "merge-user-1", "alice@example.com")
+ createUserWithEmailForTest(t, "merge-user-2", "alice@example.com")
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-user-merge-conflict",
+ "email": "alice@example.com",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if response.Success {
+ t.Fatal("expected ambiguous email merge to fail")
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginRejectsBindWhenAlreadyBound(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ })
+ boundUser := createUserForBindTest(t, "bound-user")
+ otherUser := createUserForBindTest(t, "other-user")
+ if err := model.CreateUserOAuthBinding(&model.UserOAuthBinding{
+ UserId: boundUser.Id,
+ ProviderId: provider.Id,
+ ProviderUserId: "ext-bound-user",
+ }); err != nil {
+ t.Fatalf("failed to seed oauth binding: %v", err)
+ }
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ loginReq, err := http.NewRequest(http.MethodGet, server.URL+"/test/login-as/"+strconv.Itoa(otherUser.Id), nil)
+ if err != nil {
+ t.Fatalf("failed to build login-as request: %v", err)
+ }
+ loginResp, err := client.Do(loginReq)
+ if err != nil {
+ t.Fatalf("failed to establish session: %v", err)
+ }
+ _ = loginResp.Body.Close()
+
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-bound-user",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if response.Success {
+ t.Fatal("expected bind with existing external id to fail")
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginRejectsBindForDisabledSessionUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ })
+ disabledUser := createUserForBindTest(t, "disabled-bind-user")
+ if err := model.DB.Model(disabledUser).Update("status", common.UserStatusDisabled).Error; err != nil {
+ t.Fatalf("failed to disable user: %v", err)
+ }
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ loginReq, err := http.NewRequest(http.MethodGet, server.URL+"/test/login-as/"+strconv.Itoa(disabledUser.Id), nil)
+ if err != nil {
+ t.Fatalf("failed to build login-as request: %v", err)
+ }
+ loginResp, err := client.Do(loginReq)
+ if err != nil {
+ t.Fatalf("failed to establish session: %v", err)
+ }
+ _ = loginResp.Body.Close()
+
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-disabled-bind-user",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if response.Success {
+ t.Fatal("expected disabled session user bind to fail")
+ }
+ if model.IsProviderUserIdTaken(provider.Id, "ext-disabled-bind-user") {
+ t.Fatal("expected disabled session user not to receive oauth binding")
+ }
+ log := getLatestSystemLogForUser(t, disabledUser.Id)
+ if !strings.Contains(log.Content, "failure_reason=user_disabled") {
+ t.Fatalf("expected disabled bind audit log, got %s", log.Content)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginDoesNotBindDisabledMergedUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ provider := createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: false,
+ AutoMergeByEmail: true,
+ })
+ disabledUser := createUserWithEmailForTest(t, "disabled-merge-user", "disabled@example.com")
+ disabledUser.Status = common.UserStatusDisabled
+ if err := model.DB.Model(disabledUser).Update("status", common.UserStatusDisabled).Error; err != nil {
+ t.Fatalf("failed to disable user: %v", err)
+ }
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-disabled-merge",
+ "email": "disabled@example.com",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ response := postJWTLoginForTest(t, client, server.URL, state, token)
+ if response.Success {
+ t.Fatal("expected disabled merged user login to fail")
+ }
+ if model.IsProviderUserIdTaken(provider.Id, "ext-disabled-merge") {
+ t.Fatal("expected disabled merged user not to receive oauth binding")
+ }
+ log := getLatestSystemLogForUser(t, disabledUser.Id)
+ if !strings.Contains(log.Content, "email_merge=true") || !strings.Contains(log.Content, "failure_reason=user_disabled") {
+ t.Fatalf("expected disabled merged audit log, got %s", log.Content)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginWithTicketExchangeAndUserInfoModeCreatesUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "data": map[string]any{
+ "access_token": "opaque-access-token",
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal exchange response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer exchangeServer.Close()
+
+ userInfoServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if got := r.Header.Get("x-access-token"); got != "opaque-access-token" {
+ t.Fatalf("expected exchanged token in x-access-token header, got %q", got)
+ }
+ payload, err := common.Marshal(map[string]any{
+ "info": map[string]any{
+ "userCode": "1410833903245320192",
+ "loginid": "liangmingsen",
+ "userName": "梁明森",
+ "mailbox": "liangmingsen@qdama.cn",
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal userinfo payload: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer userInfoServer.Close()
+
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ JWTIdentityMode: model.CustomJWTIdentityModeUserInfo,
+ UserInfoEndpoint: userInfoServer.URL,
+ JWTHeader: "x-access-token",
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ TicketExchangeMethod: http.MethodGet,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ TicketExchangeTicketField: "ticket",
+ TicketExchangeTokenField: "data.access_token",
+ TicketExchangeServiceField: "service",
+ })
+
+ // Override field mappings for qdama-like userinfo payload.
+ if err := model.DB.Model(&model.CustomOAuthProvider{}).
+ Where("slug = ?", "acme-sso").
+ Updates(map[string]any{
+ "user_id_field": "info.userCode",
+ "username_field": "info.loginid",
+ "display_name_field": "info.userName",
+ "email_field": "info.mailbox",
+ "group_field": "",
+ "role_field": "",
+ "group_mapping": "",
+ "role_mapping": "",
+ }).Error; err != nil {
+ t.Fatalf("failed to update provider mappings: %v", err)
+ }
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ previousServerAddress := system_setting.ServerAddress
+ system_setting.ServerAddress = server.URL
+ t.Cleanup(func() {
+ system_setting.ServerAddress = previousServerAddress
+ })
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ response := postJWTTicketLoginForTest(t, client, server.URL, state, "ST-123")
+ if !response.Success {
+ t.Fatalf("expected success response, got message: %s", response.Message)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.Username != "liangmingsen" {
+ t.Fatalf("unexpected login username: %s", loginData.Username)
+ }
+
+ var user model.User
+ if err := model.DB.Where("username = ?", "liangmingsen").First(&user).Error; err != nil {
+ t.Fatalf("expected created user liangmingsen, got error: %v", err)
+ }
+ if user.Email != "liangmingsen@qdama.cn" {
+ t.Fatalf("unexpected persisted email: %s", user.Email)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginWithTicketExchangeCreatesUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+
+ var callbackURLSeen string
+ var stateSeen string
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-ticket-login",
+ "preferred_username": "ticket-user",
+ "name": "Ticket User",
+ "email": "ticket-user@example.com",
+ "groups": []string{"engineering"},
+ "roles": []string{"platform-admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if err := r.ParseForm(); err != nil {
+ t.Fatalf("failed to parse exchange request: %v", err)
+ }
+ if got := r.Form.Get("st"); got != "ST-123" {
+ t.Fatalf("expected ticket field st=ST-123, got %q", got)
+ }
+ callbackURLSeen = r.Form.Get("service")
+ stateSeen = r.Header.Get("X-State")
+ payload, err := common.Marshal(map[string]any{
+ "data": map[string]any{
+ "token": token,
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal exchange response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer exchangeServer.Close()
+
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ TicketExchangeMethod: http.MethodPost,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeForm,
+ TicketExchangeTicketField: "st",
+ TicketExchangeTokenField: "data.token",
+ TicketExchangeServiceField: "service",
+ TicketExchangeHeaders: `{"X-State":"{state}"}`,
+ })
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ previousServerAddress := system_setting.ServerAddress
+ system_setting.ServerAddress = server.URL
+ t.Cleanup(func() {
+ system_setting.ServerAddress = previousServerAddress
+ })
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ response := postJWTTicketLoginForTest(t, client, server.URL, state, "ST-123")
+ if !response.Success {
+ t.Fatalf("expected success response, got message: %s", response.Message)
+ }
+
+ if callbackURLSeen != server.URL+"/oauth/acme-sso?state="+state {
+ t.Fatalf("expected callback url %q, got %q", server.URL+"/oauth/acme-sso?state="+state, callbackURLSeen)
+ }
+ if stateSeen != state {
+ t.Fatalf("expected state header %q, got %q", state, stateSeen)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.Username != "ticket-user" || loginData.Role != common.RoleAdminUser || loginData.Group != "vip" {
+ t.Fatalf("unexpected login response: %+v", loginData)
+ }
+}
+
+func TestSelectJWTLoginCredential(t *testing.T) {
+ t.Run("claims mode prefers id token", func(t *testing.T) {
+ provider := &model.CustomOAuthProvider{
+ JWTIdentityMode: model.CustomJWTIdentityModeClaims,
+ }
+
+ token := selectJWTLoginCredential(provider, customOAuthJWTLoginRequest{
+ Token: "access-token",
+ IDToken: "id-token",
+ JWT: "fallback-jwt",
+ })
+
+ if token != "id-token" {
+ t.Fatalf("expected id token for claims mode, got %q", token)
+ }
+ })
+
+ t.Run("userinfo mode prefers access token", func(t *testing.T) {
+ provider := &model.CustomOAuthProvider{
+ JWTIdentityMode: model.CustomJWTIdentityModeUserInfo,
+ }
+
+ token := selectJWTLoginCredential(provider, customOAuthJWTLoginRequest{
+ Token: "access-token",
+ IDToken: "id-token",
+ JWT: "fallback-jwt",
+ })
+
+ if token != "access-token" {
+ t.Fatalf("expected access token for userinfo mode, got %q", token)
+ }
+ })
+}
+
+func TestOAuthAuditFailureReason(t *testing.T) {
+ err := oauth.NewOAuthErrorWithRaw("oauth_test_failed", nil, "token=secret user=alice@example.com")
+ reason := oauthAuditFailureReason(err)
+
+ if reason != "oauth_error:oauth_test_failed" {
+ t.Fatalf("unexpected audit failure reason: %q", reason)
+ }
+ if strings.Contains(reason, "secret") || strings.Contains(reason, "alice@example.com") {
+ t.Fatalf("audit failure reason leaked raw upstream details: %q", reason)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginWithTicketValidateCreatesUser(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+
+ validationServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if got := r.URL.Query().Get("ticket"); got != "ST-CAS-123" {
+ t.Fatalf("expected ticket query param ST-CAS-123, got %q", got)
+ }
+ if got := r.URL.Query().Get("service"); !strings.Contains(got, "/oauth/acme-sso?state=") {
+ t.Fatalf("expected service callback url to contain oauth callback, got %q", got)
+ }
+ w.Header().Set("Content-Type", "application/xml")
+ _, _ = w.Write([]byte(`
+
+
+ cas-user-1
+
+ cas-user
+ CAS User
+ cas-user@example.com
+ engineering
+ platform-admin
+
+
+`))
+ }))
+ defer validationServer.Close()
+
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketValidate,
+ TicketExchangeURL: validationServer.URL,
+ TicketExchangeMethod: http.MethodGet,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ TicketExchangeTicketField: "ticket",
+ TicketExchangeServiceField: "service",
+ })
+
+ if err := model.DB.Model(&model.CustomOAuthProvider{}).
+ Where("slug = ?", "acme-sso").
+ Updates(map[string]any{
+ "user_id_field": "authenticationSuccess.user",
+ "username_field": "authenticationSuccess.attributes.loginid",
+ "display_name_field": "authenticationSuccess.attributes.userName",
+ "email_field": "authenticationSuccess.attributes.mailbox",
+ "group_field": "authenticationSuccess.attributes.group",
+ "group_mapping": `{"engineering":"vip"}`,
+ "role_field": "authenticationSuccess.attributes.role",
+ "role_mapping": `{"platform-admin":"admin"}`,
+ }).Error; err != nil {
+ t.Fatalf("failed to update provider mappings for ticket validate: %v", err)
+ }
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ previousServerAddress := system_setting.ServerAddress
+ system_setting.ServerAddress = server.URL
+ t.Cleanup(func() {
+ system_setting.ServerAddress = previousServerAddress
+ })
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ response := postJWTTicketLoginForTest(t, client, server.URL, state, "ST-CAS-123")
+ if !response.Success {
+ t.Fatalf("expected success response, got message: %s", response.Message)
+ }
+
+ var loginData oauthJWTLoginResponse
+ if err := common.Unmarshal(response.Data, &loginData); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ if loginData.Username != "cas-user" || loginData.Role != common.RoleAdminUser || loginData.Group != "vip" {
+ t.Fatalf("unexpected login response: %+v", loginData)
+ }
+
+ var user model.User
+ if err := model.DB.Where("username = ?", "cas-user").First(&user).Error; err != nil {
+ t.Fatalf("expected created user cas-user, got error: %v", err)
+ }
+ if user.Email != "cas-user@example.com" {
+ t.Fatalf("unexpected persisted email: %s", user.Email)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginWithTicketExchangeRequiresValidServerAddress(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+
+ exchangeCallCount := 0
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ exchangeCallCount++
+ w.WriteHeader(http.StatusOK)
+ }))
+ defer exchangeServer.Close()
+
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ TicketExchangeMethod: http.MethodPost,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeForm,
+ TicketExchangeTicketField: "ticket",
+ TicketExchangeTokenField: "data.token",
+ TicketExchangeServiceField: "service",
+ })
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ previousServerAddress := system_setting.ServerAddress
+ system_setting.ServerAddress = ""
+ t.Cleanup(func() {
+ system_setting.ServerAddress = previousServerAddress
+ })
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ response := postJWTTicketLoginForTest(t, client, server.URL, state, "ST-123")
+
+ if response.Success {
+ t.Fatalf("expected ticket login to fail when server address is empty")
+ }
+ if exchangeCallCount != 0 {
+ t.Fatalf("expected ticket exchange not to be called without valid server address, got %d", exchangeCallCount)
+ }
+ if response.Message == "" {
+ t.Fatalf("expected ticket login failure to include message")
+ }
+}
+
+func TestBuildCustomOAuthJWTCallbackURLRequiresValidServerAddress(t *testing.T) {
+ previousServerAddress := system_setting.ServerAddress
+ t.Cleanup(func() {
+ system_setting.ServerAddress = previousServerAddress
+ })
+
+ system_setting.ServerAddress = "://bad"
+ _, err := buildCustomOAuthJWTCallbackURL("acme-sso", "state-1")
+ if err == nil {
+ t.Fatalf("expected invalid server address to fail callback url build")
+ }
+
+ system_setting.ServerAddress = "https://example.com/base/"
+ callbackURL, err := buildCustomOAuthJWTCallbackURL("acme-sso", "state-2")
+ if err != nil {
+ t.Fatalf("expected valid server address to build callback url, got %v", err)
+ }
+ expected := "https://example.com/base/oauth/acme-sso?state=state-2"
+ if callbackURL != expected {
+ t.Fatalf("expected callback url %q, got %q", expected, callbackURL)
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginReturns200ForInvalidState(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate private key: %v", err)
+ }
+ createJWTDirectProviderForTest(t, privateKey, jwtDirectProviderTestOptions{
+ AutoRegister: true,
+ })
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ _ = fetchOAuthStateForTest(t, client, server.URL)
+ token := signJWTForControllerTest(t, privateKey, jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-user-invalid-state",
+ "preferred_username": "invalid-state-user",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ payload, marshalErr := common.Marshal(map[string]any{
+ "state": "mismatched-state",
+ "id_token": token,
+ })
+ if marshalErr != nil {
+ t.Fatalf("failed to marshal login payload: %v", marshalErr)
+ }
+ req, reqErr := http.NewRequest(http.MethodPost, server.URL+"/api/auth/external/acme-sso/jwt/login", bytes.NewReader(payload))
+ if reqErr != nil {
+ t.Fatalf("failed to build login request: %v", reqErr)
+ }
+ req.Header.Set("Content-Type", "application/json")
+ resp, doErr := client.Do(req)
+ if doErr != nil {
+ t.Fatalf("failed to post jwt login: %v", doErr)
+ }
+ defer resp.Body.Close()
+
+ if resp.StatusCode != http.StatusOK {
+ t.Fatalf("expected invalid state response to keep http 200 envelope, got %d", resp.StatusCode)
+ }
+ var response oauthJWTAPIResponse
+ if err := common.DecodeJson(resp.Body, &response); err != nil {
+ t.Fatalf("failed to decode invalid state response: %v", err)
+ }
+ if response.Success {
+ t.Fatalf("expected invalid state response to fail")
+ }
+ if response.Message == "" {
+ t.Fatalf("expected invalid state response to include an error message")
+ }
+}
+
+func TestHandleCustomOAuthJWTLoginReturns200ForUnknownProvider(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ router := newCustomOAuthJWTRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newTestHTTPClient(t)
+ state := fetchOAuthStateForTest(t, client, server.URL)
+ payload, err := common.Marshal(map[string]any{
+ "state": state,
+ "id_token": "fake-token",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal login payload: %v", err)
+ }
+ req, err := http.NewRequest(http.MethodPost, server.URL+"/api/auth/external/missing-provider/jwt/login", bytes.NewReader(payload))
+ if err != nil {
+ t.Fatalf("failed to build login request: %v", err)
+ }
+ req.Header.Set("Content-Type", "application/json")
+ resp, err := client.Do(req)
+ if err != nil {
+ t.Fatalf("failed to post jwt login: %v", err)
+ }
+ defer resp.Body.Close()
+
+ if resp.StatusCode != http.StatusOK {
+ t.Fatalf("expected unknown provider response to keep http 200 envelope, got %d", resp.StatusCode)
+ }
+ var response oauthJWTAPIResponse
+ if err := common.DecodeJson(resp.Body, &response); err != nil {
+ t.Fatalf("failed to decode unknown provider response: %v", err)
+ }
+ if response.Success {
+ t.Fatalf("expected unknown provider response to fail")
+ }
+ if response.Message == "" {
+ t.Fatalf("expected unknown provider response to include an error message")
+ }
+}
+
+func newTestHTTPClient(t *testing.T) *http.Client {
+ t.Helper()
+ jar, err := cookiejar.New(nil)
+ if err != nil {
+ t.Fatalf("failed to create cookie jar: %v", err)
+ }
+ return &http.Client{Jar: jar}
+}
+
+func fetchOAuthStateForTest(t *testing.T, client *http.Client, baseURL string) string {
+ t.Helper()
+ req, err := http.NewRequest(http.MethodGet, baseURL+"/api/oauth/state", nil)
+ if err != nil {
+ t.Fatalf("failed to build state request: %v", err)
+ }
+ resp, err := client.Do(req)
+ if err != nil {
+ t.Fatalf("failed to fetch oauth state: %v", err)
+ }
+ defer resp.Body.Close()
+
+ var response oauthJWTAPIResponse
+ if err := common.DecodeJson(resp.Body, &response); err != nil {
+ t.Fatalf("failed to decode state response: %v", err)
+ }
+ var state string
+ if err := common.Unmarshal(response.Data, &state); err != nil {
+ t.Fatalf("failed to decode state payload: %v", err)
+ }
+ return state
+}
+
+func postJWTLoginForTest(t *testing.T, client *http.Client, baseURL string, state string, token string) oauthJWTAPIResponse {
+ t.Helper()
+ payload, err := common.Marshal(map[string]any{
+ "state": state,
+ "id_token": token,
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal login payload: %v", err)
+ }
+ req, err := http.NewRequest(http.MethodPost, baseURL+"/api/auth/external/acme-sso/jwt/login", bytes.NewReader(payload))
+ if err != nil {
+ t.Fatalf("failed to build login request: %v", err)
+ }
+ req.Header.Set("Content-Type", "application/json")
+ resp, err := client.Do(req)
+ if err != nil {
+ t.Fatalf("failed to post jwt login: %v", err)
+ }
+ defer resp.Body.Close()
+
+ var response oauthJWTAPIResponse
+ if err := common.DecodeJson(resp.Body, &response); err != nil {
+ t.Fatalf("failed to decode login response: %v", err)
+ }
+ return response
+}
+
+func postJWTTicketLoginForTest(t *testing.T, client *http.Client, baseURL string, state string, ticket string) oauthJWTAPIResponse {
+ t.Helper()
+ payload, err := common.Marshal(map[string]any{
+ "state": state,
+ "ticket": ticket,
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal ticket login payload: %v", err)
+ }
+ req, err := http.NewRequest(http.MethodPost, baseURL+"/api/auth/external/acme-sso/jwt/login", bytes.NewReader(payload))
+ if err != nil {
+ t.Fatalf("failed to build ticket login request: %v", err)
+ }
+ req.Header.Set("Content-Type", "application/json")
+ resp, err := client.Do(req)
+ if err != nil {
+ t.Fatalf("failed to post ticket login: %v", err)
+ }
+ defer resp.Body.Close()
+
+ var response oauthJWTAPIResponse
+ if err := common.DecodeJson(resp.Body, &response); err != nil {
+ t.Fatalf("failed to decode ticket login response: %v", err)
+ }
+ return response
+}
+
+func signJWTForControllerTest(t *testing.T, privateKey *rsa.PrivateKey, claims jwt.MapClaims) string {
+ t.Helper()
+ token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims)
+ tokenString, err := token.SignedString(privateKey)
+ if err != nil {
+ t.Fatalf("failed to sign jwt token: %v", err)
+ }
+ return tokenString
+}
+
+func mustEncodeControllerRSAPublicKeyPEM(t *testing.T, publicKey *rsa.PublicKey) string {
+ t.Helper()
+ publicKeyDER, err := x509.MarshalPKIXPublicKey(publicKey)
+ if err != nil {
+ t.Fatalf("failed to marshal public key: %v", err)
+ }
+ return string(pem.EncodeToMemory(&pem.Block{
+ Type: "PUBLIC KEY",
+ Bytes: publicKeyDER,
+ }))
+}
diff --git a/controller/custom_oauth_update_test.go b/controller/custom_oauth_update_test.go
new file mode 100644
index 000000000000..c9e82ecb6561
--- /dev/null
+++ b/controller/custom_oauth_update_test.go
@@ -0,0 +1,338 @@
+package controller
+
+import (
+ "bytes"
+ "net/http"
+ "net/http/httptest"
+ "strconv"
+ "strings"
+ "testing"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/gin-gonic/gin"
+)
+
+func TestUpdateCustomOAuthProviderAllowsClearingJWTOptionalFields(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ provider := &model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ Enabled: true,
+ ClientId: "new-api-client",
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ Scopes: "openid profile email",
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ JwksURL: "https://issuer.example.com/.well-known/jwks.json",
+ UserIdField: "sub",
+ GroupField: "groups",
+ RoleField: "roles",
+ GroupMapping: `{"engineering":"vip"}`,
+ RoleMapping: `{"platform-admin":"admin"}`,
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+
+ payload, err := common.Marshal(map[string]any{
+ "name": provider.Name,
+ "slug": provider.Slug,
+ "kind": model.CustomOAuthProviderKindJWTDirect,
+ "enabled": true,
+ "client_id": provider.ClientId,
+ "authorization_endpoint": "",
+ "issuer": provider.Issuer,
+ "audience": "",
+ "jwks_url": provider.JwksURL,
+ "user_id_field": provider.UserIdField,
+ "group_field": "",
+ "role_field": "",
+ "group_mapping": "",
+ "role_mapping": "",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal update payload: %v", err)
+ }
+
+ recorder := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(recorder)
+ ctx.Params = gin.Params{{Key: "id", Value: strconv.Itoa(provider.Id)}}
+ ctx.Request = httptest.NewRequest(http.MethodPut, "/api/custom-oauth-provider/"+strconv.Itoa(provider.Id), bytes.NewReader(payload))
+ ctx.Request.Header.Set("Content-Type", "application/json")
+
+ UpdateCustomOAuthProvider(ctx)
+
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("expected 200 response, got %d with body %s", recorder.Code, recorder.Body.String())
+ }
+
+ updatedProvider, err := model.GetCustomOAuthProviderById(provider.Id)
+ if err != nil {
+ t.Fatalf("failed to reload provider: %v", err)
+ }
+ if updatedProvider.AuthorizationEndpoint != "" {
+ t.Fatalf("expected authorization_endpoint to be cleared, got %q", updatedProvider.AuthorizationEndpoint)
+ }
+ if updatedProvider.Audience != "" {
+ t.Fatalf("expected audience to be cleared, got %q", updatedProvider.Audience)
+ }
+ if updatedProvider.GroupField != "" {
+ t.Fatalf("expected group_field to be cleared, got %q", updatedProvider.GroupField)
+ }
+ if updatedProvider.RoleField != "" {
+ t.Fatalf("expected role_field to be cleared, got %q", updatedProvider.RoleField)
+ }
+ if updatedProvider.GroupMapping != "" || updatedProvider.RoleMapping != "" {
+ t.Fatalf("expected mappings to be cleared, got group=%q role=%q", updatedProvider.GroupMapping, updatedProvider.RoleMapping)
+ }
+}
+
+func TestUpdateCustomOAuthProviderRejectsUnsupportedJWTSyncRoleTargets(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ provider := &model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ Enabled: true,
+ ClientId: "new-api-client",
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ Scopes: "openid profile email",
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ JwksURL: "https://issuer.example.com/.well-known/jwks.json",
+ UserIdField: "sub",
+ RoleField: "roles",
+ RoleMapping: `{"platform-admin":"admin"}`,
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+
+ payload, err := common.Marshal(map[string]any{
+ "role_mapping": `{"member":"guest"}`,
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal update payload: %v", err)
+ }
+
+ recorder := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(recorder)
+ ctx.Params = gin.Params{{Key: "id", Value: strconv.Itoa(provider.Id)}}
+ ctx.Request = httptest.NewRequest(http.MethodPut, "/api/custom-oauth-provider/"+strconv.Itoa(provider.Id), bytes.NewReader(payload))
+ ctx.Request.Header.Set("Content-Type", "application/json")
+
+ UpdateCustomOAuthProvider(ctx)
+
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("expected 200 response envelope, got %d with body %s", recorder.Code, recorder.Body.String())
+ }
+ if !strings.Contains(recorder.Body.String(), "\"success\":false") {
+ t.Fatalf("expected update to fail for unsupported role target, got body %s", recorder.Body.String())
+ }
+}
+
+func TestUpdateCustomOAuthProviderRejectsInvalidTicketExchangeURL(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ provider := &model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ Enabled: true,
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ Issuer: "https://issuer.example.com",
+ JwksURL: "https://issuer.example.com/.well-known/jwks.json",
+ UserIdField: "sub",
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+
+ payload, err := common.Marshal(map[string]any{
+ "jwt_acquire_mode": model.CustomJWTAcquireModeTicketExchange,
+ "ticket_exchange_url": "not-a-url",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal update payload: %v", err)
+ }
+
+ recorder := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(recorder)
+ ctx.Params = gin.Params{{Key: "id", Value: strconv.Itoa(provider.Id)}}
+ ctx.Request = httptest.NewRequest(http.MethodPut, "/api/custom-oauth-provider/"+strconv.Itoa(provider.Id), bytes.NewReader(payload))
+ ctx.Request.Header.Set("Content-Type", "application/json")
+
+ UpdateCustomOAuthProvider(ctx)
+
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("expected 200 response envelope, got %d with body %s", recorder.Code, recorder.Body.String())
+ }
+ if !strings.Contains(recorder.Body.String(), "\"success\":false") {
+ t.Fatalf("expected invalid ticket_exchange_url to fail, got body %s", recorder.Body.String())
+ }
+}
+
+func TestUpdateCustomOAuthProviderAllowsJWTUserInfoModeWithoutVerificationKey(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ provider := &model.CustomOAuthProvider{
+ Name: "Qdama SSO",
+ Slug: "qdama-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ Enabled: true,
+ AuthorizationEndpoint: "https://cas.qdama.cn/login",
+ Issuer: "https://issuer.example.com",
+ JwksURL: "https://issuer.example.com/.well-known/jwks.json",
+ UserIdField: "sub",
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+
+ payload, err := common.Marshal(map[string]any{
+ "jwt_identity_mode": model.CustomJWTIdentityModeUserInfo,
+ "issuer": "",
+ "jwks_url": "",
+ "public_key": "",
+ "user_info_endpoint": "https://my-api.qdama.cn/v1/api/my/pc/getInfo",
+ "jwt_header": "x-access-token",
+ "user_id_field": "info.userCode",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal update payload: %v", err)
+ }
+
+ recorder := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(recorder)
+ ctx.Params = gin.Params{{Key: "id", Value: strconv.Itoa(provider.Id)}}
+ ctx.Request = httptest.NewRequest(http.MethodPut, "/api/custom-oauth-provider/"+strconv.Itoa(provider.Id), bytes.NewReader(payload))
+ ctx.Request.Header.Set("Content-Type", "application/json")
+
+ UpdateCustomOAuthProvider(ctx)
+
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("expected 200 response, got %d with body %s", recorder.Code, recorder.Body.String())
+ }
+ if strings.Contains(recorder.Body.String(), "\"success\":false") {
+ t.Fatalf("expected userinfo mode update to succeed, got body %s", recorder.Body.String())
+ }
+}
+
+func TestUpdateCustomOAuthProviderRejectsTicketValidateWithUserInfoMode(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ provider := &model.CustomOAuthProvider{
+ Name: "CAS SSO",
+ Slug: "cas-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ Enabled: true,
+ AuthorizationEndpoint: "https://cas.example.com/login",
+ Issuer: "https://issuer.example.com",
+ JwksURL: "https://issuer.example.com/.well-known/jwks.json",
+ UserIdField: "sub",
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+
+ payload, err := common.Marshal(map[string]any{
+ "jwt_acquire_mode": model.CustomJWTAcquireModeTicketValidate,
+ "jwt_identity_mode": model.CustomJWTIdentityModeUserInfo,
+ "ticket_exchange_url": "https://cas.example.com/serviceValidate",
+ "user_info_endpoint": "https://api.example.com/userinfo",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal update payload: %v", err)
+ }
+
+ recorder := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(recorder)
+ ctx.Params = gin.Params{{Key: "id", Value: strconv.Itoa(provider.Id)}}
+ ctx.Request = httptest.NewRequest(http.MethodPut, "/api/custom-oauth-provider/"+strconv.Itoa(provider.Id), bytes.NewReader(payload))
+ ctx.Request.Header.Set("Content-Type", "application/json")
+
+ UpdateCustomOAuthProvider(ctx)
+
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("expected 200 response envelope, got %d with body %s", recorder.Code, recorder.Body.String())
+ }
+ if !strings.Contains(recorder.Body.String(), "\"success\":false") {
+ t.Fatalf("expected ticket_validate + userinfo update to fail, got body %s", recorder.Body.String())
+ }
+ if !strings.Contains(recorder.Body.String(), "ticket_validate mode only support claims") {
+ t.Fatalf("expected ticket_validate userinfo validation message, got body %s", recorder.Body.String())
+ }
+}
+
+func TestUpdateCustomOAuthProviderAllowsClearingClientSecret(t *testing.T) {
+ setupCustomOAuthJWTControllerTestDB(t)
+
+ provider := &model.CustomOAuthProvider{
+ Name: "Acme OAuth",
+ Slug: "acme-oauth",
+ Kind: model.CustomOAuthProviderKindOAuthCode,
+ Enabled: true,
+ ClientId: "client-id",
+ ClientSecret: "secret-to-clear",
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ TokenEndpoint: "https://issuer.example.com/oauth2/token",
+ UserInfoEndpoint: "https://issuer.example.com/oauth2/userinfo",
+ UserIdField: "id",
+ }
+ if err := model.CreateCustomOAuthProvider(provider); err != nil {
+ t.Fatalf("failed to create provider: %v", err)
+ }
+
+ payload, err := common.Marshal(map[string]any{
+ "client_secret": "",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal update payload: %v", err)
+ }
+
+ recorder := httptest.NewRecorder()
+ ctx, _ := gin.CreateTestContext(recorder)
+ ctx.Params = gin.Params{{Key: "id", Value: strconv.Itoa(provider.Id)}}
+ ctx.Request = httptest.NewRequest(http.MethodPut, "/api/custom-oauth-provider/"+strconv.Itoa(provider.Id), bytes.NewReader(payload))
+ ctx.Request.Header.Set("Content-Type", "application/json")
+
+ UpdateCustomOAuthProvider(ctx)
+
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("expected 200 response, got %d with body %s", recorder.Code, recorder.Body.String())
+ }
+
+ updatedProvider, err := model.GetCustomOAuthProviderById(provider.Id)
+ if err != nil {
+ t.Fatalf("failed to reload provider: %v", err)
+ }
+ if updatedProvider.ClientSecret != "" {
+ t.Fatalf("expected client_secret to be cleared, got %q", updatedProvider.ClientSecret)
+ }
+}
+
+func TestCustomOAuthProviderResponseOmitsTicketExchangeSecrets(t *testing.T) {
+ response := toCustomOAuthProviderResponse(&model.CustomOAuthProvider{
+ Name: "CAS SSO",
+ Slug: "cas-sso",
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ TicketExchangeExtraParams: `{"token":"sensitive"}`,
+ TicketExchangeHeaders: `{"Authorization":"Bearer sensitive"}`,
+ })
+
+ payload, err := common.Marshal(response)
+ if err != nil {
+ t.Fatalf("failed to marshal response: %v", err)
+ }
+
+ if strings.Contains(string(payload), "ticket_exchange_extra_params") {
+ t.Fatalf("expected response payload to omit ticket_exchange_extra_params, got %s", string(payload))
+ }
+ if strings.Contains(string(payload), "ticket_exchange_headers") {
+ t.Fatalf("expected response payload to omit ticket_exchange_headers, got %s", string(payload))
+ }
+}
diff --git a/controller/misc.go b/controller/misc.go
index 519caed57b81..117ac8349072 100644
--- a/controller/misc.go
+++ b/controller/misc.go
@@ -5,13 +5,13 @@ import (
"fmt"
"net/http"
"strings"
+ "sync"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/middleware"
"github.com/QuantumNous/new-api/model"
- "github.com/QuantumNous/new-api/oauth"
"github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/console_setting"
"github.com/QuantumNous/new-api/setting/operation_setting"
@@ -20,6 +20,84 @@ import (
"github.com/gin-gonic/gin"
)
+type customOAuthStatusInfo struct {
+ Id int `json:"id"`
+ Name string `json:"name"`
+ Slug string `json:"slug"`
+ Icon string `json:"icon"`
+ Kind string `json:"kind"`
+ ClientId string `json:"client_id"`
+ AuthorizationEndpoint string `json:"authorization_endpoint"`
+ Scopes string `json:"scopes"`
+ JWTSource string `json:"jwt_source"`
+ JWTIdentityMode string `json:"jwt_identity_mode"`
+ JWTAcquireMode string `json:"jwt_acquire_mode"`
+ AuthorizationServiceField string `json:"authorization_service_field"`
+ BrowserLoginSupported bool `json:"browser_login_supported"`
+}
+
+var (
+ customOAuthStatusCacheMu sync.RWMutex
+ customOAuthStatusCacheData []customOAuthStatusInfo
+ customOAuthStatusCacheInit bool
+)
+
+func invalidateCustomOAuthStatusCache() {
+ customOAuthStatusCacheMu.Lock()
+ defer customOAuthStatusCacheMu.Unlock()
+ customOAuthStatusCacheData = nil
+ customOAuthStatusCacheInit = false
+}
+
+func getCustomOAuthStatusPayload() []customOAuthStatusInfo {
+ customOAuthStatusCacheMu.RLock()
+ if customOAuthStatusCacheInit {
+ cached := append([]customOAuthStatusInfo(nil), customOAuthStatusCacheData...)
+ customOAuthStatusCacheMu.RUnlock()
+ return cached
+ }
+ customOAuthStatusCacheMu.RUnlock()
+
+ customProviders, err := model.GetEnabledCustomOAuthProviders()
+ if err != nil {
+ common.SysError("failed to load enabled custom auth providers: " + err.Error())
+ return nil
+ }
+
+ providersInfo := make([]customOAuthStatusInfo, 0, len(customProviders))
+ for _, config := range customProviders {
+ jwtSource := config.JWTSource
+ if strings.TrimSpace(jwtSource) == "" {
+ jwtSource = model.CustomJWTSourceQuery
+ }
+ authorizationServiceField := config.AuthorizationServiceField
+ if strings.TrimSpace(authorizationServiceField) == "" {
+ authorizationServiceField = "service"
+ }
+ providersInfo = append(providersInfo, customOAuthStatusInfo{
+ Id: config.Id,
+ Name: config.Name,
+ Slug: config.Slug,
+ Icon: config.Icon,
+ Kind: config.GetKind(),
+ ClientId: config.ClientId,
+ AuthorizationEndpoint: config.AuthorizationEndpoint,
+ Scopes: config.Scopes,
+ JWTSource: jwtSource,
+ JWTIdentityMode: config.GetJWTIdentityMode(),
+ JWTAcquireMode: config.GetJWTAcquireMode(),
+ AuthorizationServiceField: authorizationServiceField,
+ BrowserLoginSupported: config.SupportsBrowserLogin(),
+ })
+ }
+
+ customOAuthStatusCacheMu.Lock()
+ customOAuthStatusCacheData = append([]customOAuthStatusInfo(nil), providersInfo...)
+ customOAuthStatusCacheInit = true
+ customOAuthStatusCacheMu.Unlock()
+ return providersInfo
+}
+
func TestStatus(c *gin.Context) {
err := model.PingDB()
if err != nil {
@@ -42,6 +120,7 @@ func TestStatus(c *gin.Context) {
func GetStatus(c *gin.Context) {
cs := console_setting.GetConsoleSetting()
+ customProviders := getCustomOAuthStatusPayload()
common.OptionMapRWMutex.RLock()
defer common.OptionMapRWMutex.RUnlock()
@@ -130,32 +209,9 @@ func GetStatus(c *gin.Context) {
data["faq"] = console_setting.GetFAQ()
}
- // Add enabled custom OAuth providers
- customProviders := oauth.GetEnabledCustomProviders()
+ // Add enabled custom auth providers
if len(customProviders) > 0 {
- type CustomOAuthInfo struct {
- Id int `json:"id"`
- Name string `json:"name"`
- Slug string `json:"slug"`
- Icon string `json:"icon"`
- ClientId string `json:"client_id"`
- AuthorizationEndpoint string `json:"authorization_endpoint"`
- Scopes string `json:"scopes"`
- }
- providersInfo := make([]CustomOAuthInfo, 0, len(customProviders))
- for _, p := range customProviders {
- config := p.GetConfig()
- providersInfo = append(providersInfo, CustomOAuthInfo{
- Id: config.Id,
- Name: config.Name,
- Slug: config.Slug,
- Icon: config.Icon,
- ClientId: config.ClientId,
- AuthorizationEndpoint: config.AuthorizationEndpoint,
- Scopes: config.Scopes,
- })
- }
- data["custom_oauth_providers"] = providersInfo
+ data["custom_oauth_providers"] = customProviders
}
c.JSON(http.StatusOK, gin.H{
diff --git a/controller/oauth.go b/controller/oauth.go
index 9951f22b035f..92062baa4dca 100644
--- a/controller/oauth.go
+++ b/controller/oauth.go
@@ -4,6 +4,7 @@ import (
"fmt"
"net/http"
"strconv"
+ "strings"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/i18n"
@@ -104,27 +105,36 @@ func HandleOAuth(c *gin.Context) {
}
// 7. Find or create user
- user, err := findOrCreateOAuthUser(c, provider, oauthUser, session)
- if err != nil {
- switch err.(type) {
- case *OAuthUserDeletedError:
- common.ApiErrorI18n(c, i18n.MsgOAuthUserDeleted)
- case *OAuthRegistrationDisabledError:
- common.ApiErrorI18n(c, i18n.MsgUserRegisterDisabled)
- default:
- common.ApiError(c, err)
+ options := oauthFindOrCreateOptions{
+ AllowAutoRegister: true,
+ AllowAutoMergeByEmail: false,
+ InitialRole: common.RoleCommonUser,
+ }
+ if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok {
+ if config := genericProvider.GetConfig(); config != nil {
+ options.AllowAutoMergeByEmail = config.AutoMergeByEmail
}
+ }
+ resolvedUser, err := findOrCreateOAuthUserWithOptions(c, provider, oauthUser, session, options)
+ if err != nil {
+ handleOAuthUserError(c, err)
return
}
// 8. Check user status
- if user.Status != common.UserStatusEnabled {
+ if resolvedUser.User.Status != common.UserStatusEnabled {
common.ApiErrorI18n(c, i18n.MsgOAuthUserBanned)
return
}
+ if resolvedUser.BindAfterStatusCheck {
+ if err := bindOAuthIdentityToUser(resolvedUser.User, provider, oauthUser.ProviderUserID); err != nil {
+ common.ApiError(c, err)
+ return
+ }
+ }
// 9. Setup login
- setupLogin(user, c)
+ setupLogin(resolvedUser.User, c)
}
// handleOAuthBind handles binding OAuth account to existing user
@@ -149,54 +159,94 @@ func handleOAuthBind(c *gin.Context, provider oauth.Provider) {
return
}
+ handleOAuthBindWithUser(c, provider, oauthUser)
+}
+
+func handleOAuthBindWithUser(c *gin.Context, provider oauth.Provider, oauthUser *oauth.OAuthUser) {
+ err := bindOAuthIdentityToCurrentUser(c, provider, oauthUser)
+ if err != nil {
+ if boundErr, ok := err.(*OAuthAlreadyBoundError); ok {
+ common.ApiErrorI18n(c, i18n.MsgOAuthAlreadyBound, providerParams(boundErr.Provider))
+ return
+ }
+ common.ApiError(c, err)
+ return
+ }
+
+ common.ApiSuccessI18n(c, i18n.MsgOAuthBindSuccess, gin.H{
+ "action": "bind",
+ })
+}
+
+func bindOAuthIdentityToCurrentUser(c *gin.Context, provider oauth.Provider, oauthUser *oauth.OAuthUser) error {
// Check if this OAuth account is already bound (check both new ID and legacy ID)
if provider.IsUserIDTaken(oauthUser.ProviderUserID) {
- common.ApiErrorI18n(c, i18n.MsgOAuthAlreadyBound, providerParams(provider.GetName()))
- return
+ return &OAuthAlreadyBoundError{Provider: provider.GetName()}
}
// Also check legacy ID to prevent duplicate bindings during migration period
if legacyID, ok := oauthUser.Extra["legacy_id"].(string); ok && legacyID != "" {
if provider.IsUserIDTaken(legacyID) {
- common.ApiErrorI18n(c, i18n.MsgOAuthAlreadyBound, providerParams(provider.GetName()))
- return
+ return &OAuthAlreadyBoundError{Provider: provider.GetName()}
}
}
// Get current user from session
session := sessions.Default(c)
id := session.Get("id")
+ if id == nil {
+ return fmt.Errorf("missing current user session")
+ }
user := model.User{Id: id.(int)}
- err = user.FillUserById()
+ err := user.FillUserById()
if err != nil {
- common.ApiError(c, err)
- return
+ return err
}
// Handle binding based on provider type
- if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok {
+ if customBindingProvider, ok := provider.(oauth.CustomBindingProvider); ok {
// Custom provider: use user_oauth_bindings table
- err = model.UpdateUserOAuthBinding(user.Id, genericProvider.GetProviderId(), oauthUser.ProviderUserID)
+ err = ensureUserHasNoCustomProviderBinding(user.Id, customBindingProvider.GetProviderId())
+ if err == nil {
+ err = model.UpdateUserOAuthBinding(user.Id, customBindingProvider.GetProviderId(), oauthUser.ProviderUserID)
+ }
if err != nil {
- common.ApiError(c, err)
- return
+ return err
}
} else {
// Built-in provider: update user record directly
provider.SetProviderUserID(&user, oauthUser.ProviderUserID)
err = user.Update(false)
if err != nil {
- common.ApiError(c, err)
- return
+ return err
}
}
+ return nil
+}
- common.ApiSuccessI18n(c, i18n.MsgOAuthBindSuccess, gin.H{
- "action": "bind",
+// findOrCreateOAuthUser finds existing user or creates new user
+func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *oauth.OAuthUser, session sessions.Session) (*oauthUserResolutionResult, error) {
+ return findOrCreateOAuthUserWithOptions(c, provider, oauthUser, session, oauthFindOrCreateOptions{
+ AllowAutoRegister: true,
+ AllowAutoMergeByEmail: false,
+ InitialRole: common.RoleCommonUser,
})
}
-// 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) {
+type oauthFindOrCreateOptions struct {
+ AllowAutoRegister bool
+ AllowAutoMergeByEmail bool
+ InitialRole int
+ InitialGroup string
+}
+
+type oauthUserResolutionResult struct {
+ User *model.User
+ BindAfterStatusCheck bool
+ AutoRegisterTriggered bool
+ EmailMergeTriggered bool
+}
+
+func findOrCreateOAuthUserWithOptions(c *gin.Context, provider oauth.Provider, oauthUser *oauth.OAuthUser, session sessions.Session, options oauthFindOrCreateOptions) (*oauthUserResolutionResult, error) {
user := &model.User{}
// Check if user already exists with new ID
@@ -209,12 +259,12 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
if user.Id == 0 {
return nil, &OAuthUserDeletedError{}
}
- return user, nil
+ return &oauthUserResolutionResult{User: user}, 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) {
+ if _, ok := provider.(*oauth.GitHubProvider); ok && provider.IsUserIDTaken(legacyID) {
err := provider.FillUserByProviderID(user, legacyID)
if err != nil {
return nil, err
@@ -227,12 +277,34 @@ 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 &oauthUserResolutionResult{User: user}, nil
}
}
}
+ if options.AllowAutoMergeByEmail {
+ mergedUser, err := findOAuthMergeCandidateByEmail(oauthUser.Email)
+ if err != nil {
+ return nil, err
+ }
+ if mergedUser != nil {
+ if customBindingProvider, ok := provider.(oauth.CustomBindingProvider); ok {
+ if err := ensureUserHasNoCustomProviderBinding(mergedUser.Id, customBindingProvider.GetProviderId()); err != nil {
+ return nil, err
+ }
+ }
+ return &oauthUserResolutionResult{
+ User: mergedUser,
+ BindAfterStatusCheck: true,
+ EmailMergeTriggered: true,
+ }, nil
+ }
+ }
+
// User doesn't exist, create new user if registration is enabled
+ if !options.AllowAutoRegister {
+ return nil, &OAuthAutoRegisterDisabledError{}
+ }
if !common.RegisterEnabled {
return nil, &OAuthRegistrationDisabledError{}
}
@@ -259,7 +331,10 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
if oauthUser.Email != "" {
user.Email = oauthUser.Email
}
- user.Role = common.RoleCommonUser
+ user.Role = normalizeOAuthInitialRole(options.InitialRole)
+ if strings.TrimSpace(options.InitialGroup) != "" {
+ user.Group = strings.TrimSpace(options.InitialGroup)
+ }
user.Status = common.UserStatusEnabled
// Handle affiliate code
@@ -270,7 +345,7 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
}
// Use transaction to ensure user creation and OAuth binding are atomic
- if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok {
+ if customBindingProvider, ok := provider.(oauth.CustomBindingProvider); ok {
// Custom provider: create user and binding in a transaction
err := model.DB.Transaction(func(tx *gorm.DB) error {
// Create user
@@ -281,7 +356,7 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
// Create OAuth binding
binding := &model.UserOAuthBinding{
UserId: user.Id,
- ProviderId: genericProvider.GetProviderId(),
+ ProviderId: customBindingProvider.GetProviderId(),
ProviderUserId: oauthUser.ProviderUserID,
}
if err := model.CreateUserOAuthBindingWithTx(tx, binding); err != nil {
@@ -327,7 +402,166 @@ func findOrCreateOAuthUser(c *gin.Context, provider oauth.Provider, oauthUser *o
user.FinalizeOAuthUserCreation(inviterId)
}
- return user, nil
+ return &oauthUserResolutionResult{
+ User: user,
+ AutoRegisterTriggered: true,
+ }, nil
+}
+
+func findOAuthMergeCandidateByEmail(email string) (*model.User, error) {
+ email = strings.TrimSpace(email)
+ if email == "" {
+ return nil, nil
+ }
+ var users []*model.User
+ if err := model.DB.Unscoped().Where("email = ?", email).Limit(2).Find(&users).Error; err != nil {
+ return nil, err
+ }
+ if len(users) == 0 {
+ return nil, nil
+ }
+ if len(users) > 1 {
+ return nil, fmt.Errorf("multiple users matched email %s, auto merge is not allowed", email)
+ }
+ if users[0].DeletedAt.Valid {
+ return nil, &OAuthUserDeletedError{}
+ }
+ return users[0], nil
+}
+
+func ensureUserHasNoCustomProviderBinding(userID, providerID int) error {
+ _, err := model.GetUserOAuthBinding(userID, providerID)
+ if err == nil {
+ return fmt.Errorf("user already has a binding for provider %d", providerID)
+ }
+ if err == gorm.ErrRecordNotFound {
+ return nil
+ }
+ return err
+}
+
+func bindOAuthIdentityToUser(user *model.User, provider oauth.Provider, providerUserID string) error {
+ if customBindingProvider, ok := provider.(oauth.CustomBindingProvider); ok {
+ return model.UpdateUserOAuthBinding(user.Id, customBindingProvider.GetProviderId(), providerUserID)
+ }
+ provider.SetProviderUserID(user, providerUserID)
+ return user.Update(false)
+}
+
+func normalizeOAuthInitialRole(role int) int {
+ switch role {
+ case common.RoleGuestUser, common.RoleCommonUser, common.RoleAdminUser:
+ return role
+ default:
+ return common.RoleCommonUser
+ }
+}
+
+func syncOAuthUserLoginAttributes(user *model.User, providerName string, nextGroup string, syncGroup bool, nextRole int, syncRole bool) error {
+ if user == nil {
+ return fmt.Errorf("user is nil")
+ }
+
+ changes := make([]string, 0, 2)
+ if syncGroup {
+ group := strings.TrimSpace(nextGroup)
+ if group != "" && group != user.Group {
+ changes = append(changes, fmt.Sprintf("group %s -> %s", safeOAuthAuditValue(user.Group), group))
+ user.Group = group
+ }
+ }
+
+ oldRole := user.Role
+ if syncRole && isOAuthSyncRole(nextRole) && nextRole != user.Role {
+ changes = append(changes, fmt.Sprintf("role %s -> %s", oauthRoleLabel(user.Role), oauthRoleLabel(nextRole)))
+ user.Role = nextRole
+ }
+
+ if len(changes) == 0 {
+ return nil
+ }
+
+ if err := ensureOAuthSidebarForRoleChange(user, oldRole, user.Role); err != nil {
+ common.SysLog(fmt.Sprintf("[OAuth] Failed to align sidebar for user %d after role sync: %v", user.Id, err))
+ }
+
+ if err := user.Update(false); err != nil {
+ return err
+ }
+
+ content := fmt.Sprintf("外部登录同步用户属性(%s):%s", providerName, strings.Join(changes, ","))
+ model.RecordLog(user.Id, model.LogTypeSystem, content)
+ common.SysLog(fmt.Sprintf("[OAuth] %s", content))
+ return nil
+}
+
+func ensureOAuthSidebarForRoleChange(user *model.User, oldRole int, newRole int) error {
+ if user == nil || oldRole == newRole || newRole != common.RoleAdminUser {
+ return nil
+ }
+
+ defaultConfig := generateDefaultSidebarConfig(newRole)
+ if defaultConfig == "" {
+ return nil
+ }
+
+ defaultSidebar := make(map[string]any)
+ if err := common.UnmarshalJsonStr(defaultConfig, &defaultSidebar); err != nil {
+ return err
+ }
+
+ adminSection, ok := defaultSidebar["admin"]
+ if !ok {
+ return nil
+ }
+
+ setting := user.GetSetting()
+ sidebar := make(map[string]any)
+ if raw := strings.TrimSpace(setting.SidebarModules); raw != "" {
+ if err := common.UnmarshalJsonStr(raw, &sidebar); err != nil {
+ return err
+ }
+ }
+ if _, exists := sidebar["admin"]; exists {
+ return nil
+ }
+
+ sidebar["admin"] = adminSection
+ setting.SidebarModules = common.MapToJsonStr(sidebar)
+ user.SetSetting(setting)
+ return nil
+}
+
+func isOAuthSyncRole(role int) bool {
+ switch role {
+ case common.RoleCommonUser, common.RoleAdminUser:
+ return true
+ default:
+ return false
+ }
+}
+
+func oauthRoleLabel(role int) string {
+ switch role {
+ case common.RoleAdminUser:
+ return "admin"
+ case common.RoleCommonUser:
+ return "common"
+ case common.RoleGuestUser:
+ return "guest"
+ case common.RoleRootUser:
+ return "root"
+ default:
+ return strconv.Itoa(role)
+ }
+}
+
+func safeOAuthAuditValue(value string) string {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ return "(empty)"
+ }
+ return value
}
// Error types for OAuth
@@ -337,12 +571,37 @@ func (e *OAuthUserDeletedError) Error() string {
return "user has been deleted"
}
+type OAuthAlreadyBoundError struct {
+ Provider string
+}
+
+func (e *OAuthAlreadyBoundError) Error() string {
+ return "oauth account is already bound"
+}
+
type OAuthRegistrationDisabledError struct{}
func (e *OAuthRegistrationDisabledError) Error() string {
return "registration is disabled"
}
+type OAuthAutoRegisterDisabledError struct{}
+
+func (e *OAuthAutoRegisterDisabledError) Error() string {
+ return "provider auto registration is disabled"
+}
+
+func handleOAuthUserError(c *gin.Context, err error) {
+ switch err.(type) {
+ case *OAuthUserDeletedError:
+ common.ApiErrorI18n(c, i18n.MsgOAuthUserDeleted)
+ case *OAuthRegistrationDisabledError, *OAuthAutoRegisterDisabledError:
+ common.ApiErrorI18n(c, i18n.MsgUserRegisterDisabled)
+ default:
+ common.ApiError(c, err)
+ }
+}
+
// handleOAuthError handles OAuth errors and returns translated message
func handleOAuthError(c *gin.Context, err error) {
switch e := err.(type) {
diff --git a/controller/telegram.go b/controller/telegram.go
index f16cdd66c545..5ae105467035 100644
--- a/controller/telegram.go
+++ b/controller/telegram.go
@@ -5,13 +5,11 @@ import (
"crypto/sha256"
"encoding/hex"
"io"
- "net/http"
"sort"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
- "github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
@@ -40,23 +38,14 @@ func TelegramBind(c *gin.Context) {
return
}
- session := sessions.Default(c)
- id := session.Get("id")
- user := model.User{Id: id.(int)}
- if err := user.FillUserById(); err != nil {
+ user, err := getSessionUser(c)
+ if err != nil {
c.JSON(200, gin.H{
"message": err.Error(),
"success": false,
})
return
}
- if user.Id == 0 {
- c.JSON(http.StatusOK, gin.H{
- "success": false,
- "message": "用户已注销",
- })
- return
- }
user.TelegramId = telegramId
if err := user.Update(false); err != nil {
c.JSON(200, gin.H{
diff --git a/controller/user.go b/controller/user.go
index 8229d0d2c2bc..dc76b9078a41 100644
--- a/controller/user.go
+++ b/controller/user.go
@@ -86,6 +86,10 @@ func Login(c *gin.Context) {
// setup session & cookies and then return user info
func setupLogin(user *model.User, c *gin.Context) {
+ _ = setupLoginWithResult(user, c)
+}
+
+func setupLoginWithResult(user *model.User, c *gin.Context) bool {
session := sessions.Default(c)
session.Set("id", user.Id)
session.Set("username", user.Username)
@@ -95,7 +99,7 @@ func setupLogin(user *model.User, c *gin.Context) {
err := session.Save()
if err != nil {
common.ApiErrorI18n(c, i18n.MsgUserSessionSaveFailed)
- return
+ return false
}
c.JSON(http.StatusOK, gin.H{
"message": "",
@@ -109,6 +113,7 @@ func setupLogin(user *model.User, c *gin.Context) {
"group": user.Group,
},
})
+ return true
}
func Logout(c *gin.Context) {
@@ -931,6 +936,11 @@ type emailBindRequest struct {
}
func EmailBind(c *gin.Context) {
+ user, err := getSessionUser(c)
+ if err != nil {
+ common.ApiError(c, err)
+ return
+ }
var req emailBindRequest
if err := common.DecodeJson(c.Request.Body, &req); err != nil {
common.ApiError(c, errors.New("invalid request body"))
@@ -942,16 +952,6 @@ func EmailBind(c *gin.Context) {
common.ApiErrorI18n(c, i18n.MsgUserVerificationCodeError)
return
}
- session := sessions.Default(c)
- id := session.Get("id")
- user := model.User{
- Id: id.(int),
- }
- err := user.FillUserById()
- if err != nil {
- common.ApiError(c, err)
- return
- }
user.Email = email
// no need to check if this email already taken, because we have used verification code to check it
err = user.Update(false)
diff --git a/controller/wechat.go b/controller/wechat.go
index 8889daca77db..cadf4ed69aae 100644
--- a/controller/wechat.go
+++ b/controller/wechat.go
@@ -12,7 +12,6 @@ import (
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
- "github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
@@ -134,6 +133,11 @@ func WeChatBind(c *gin.Context) {
})
return
}
+ user, err := getSessionUser(c)
+ if err != nil {
+ common.ApiError(c, err)
+ return
+ }
var req wechatBindRequest
if err := common.DecodeJson(c.Request.Body, &req); err != nil {
c.JSON(http.StatusOK, gin.H{
@@ -158,16 +162,6 @@ func WeChatBind(c *gin.Context) {
})
return
}
- session := sessions.Default(c)
- id := session.Get("id")
- user := model.User{
- Id: id.(int),
- }
- err = user.FillUserById()
- if err != nil {
- common.ApiError(c, err)
- return
- }
user.WeChatId = wechatId
err = user.Update(false)
if err != nil {
diff --git a/i18n/keys.go b/i18n/keys.go
index 8118dff9c69c..e7964df99c03 100644
--- a/i18n/keys.go
+++ b/i18n/keys.go
@@ -279,6 +279,8 @@ const (
MsgOAuthTokenFailed = "oauth.token_failed"
MsgOAuthUserInfoEmpty = "oauth.user_info_empty"
MsgOAuthTrustLevelLow = "oauth.trust_level_low"
+ MsgOAuthTicketMissing = "oauth.ticket_missing"
+ MsgOAuthJWTMissing = "oauth.jwt_missing"
)
// Model layer error messages (for translation in controller)
diff --git a/i18n/locales/en.yaml b/i18n/locales/en.yaml
index 75a8bc6e775c..aec5fabc3bf4 100644
--- a/i18n/locales/en.yaml
+++ b/i18n/locales/en.yaml
@@ -235,6 +235,8 @@ oauth.connect_failed: "Unable to connect to {{.Provider}} server, please try aga
oauth.token_failed: "Failed to get token from {{.Provider}}, please check settings"
oauth.user_info_empty: "{{.Provider}} returned empty user info, please check settings"
oauth.trust_level_low: "Linux DO trust level does not meet the minimum required by administrator"
+oauth.ticket_missing: "Missing login ticket"
+oauth.jwt_missing: "Missing JWT token"
# Model layer error messages
redeem.failed: "Redemption failed, please try again later"
diff --git a/i18n/locales/zh-CN.yaml b/i18n/locales/zh-CN.yaml
index 1f3b5a7b4bc5..2f36056b5dee 100644
--- a/i18n/locales/zh-CN.yaml
+++ b/i18n/locales/zh-CN.yaml
@@ -236,6 +236,8 @@ oauth.connect_failed: "无法连接至 {{.Provider}} 服务器,请稍后重试
oauth.token_failed: "{{.Provider}} 获取 Token 失败,请检查设置"
oauth.user_info_empty: "{{.Provider}} 获取用户信息为空,请检查设置"
oauth.trust_level_low: "Linux DO 信任等级未达到管理员设置的最低信任等级"
+oauth.ticket_missing: "未提供登录票据"
+oauth.jwt_missing: "未提供 JWT 令牌"
# Model layer error messages
redeem.failed: "兑换失败,请稍后重试"
diff --git a/i18n/locales/zh-TW.yaml b/i18n/locales/zh-TW.yaml
index 1231c0e2480c..fd3e5d15af7a 100644
--- a/i18n/locales/zh-TW.yaml
+++ b/i18n/locales/zh-TW.yaml
@@ -236,6 +236,8 @@ oauth.connect_failed: "無法連接至 {{.Provider}} 伺服器,請稍後重試
oauth.token_failed: "{{.Provider}} 獲取 Token 失敗,請檢查設定"
oauth.user_info_empty: "{{.Provider}} 獲取使用者資訊為空,請檢查設定"
oauth.trust_level_low: "Linux DO 信任等級未達到管理員設定的最低信任等級"
+oauth.ticket_missing: "未提供登入票據"
+oauth.jwt_missing: "未提供 JWT 權杖"
# Model layer error messages
redeem.failed: "兌換失敗,請稍後重試"
diff --git a/middleware/auth.go b/middleware/auth.go
index 342e7f49812f..118e70baaa3f 100644
--- a/middleware/auth.go
+++ b/middleware/auth.go
@@ -30,14 +30,54 @@ func validUserInfo(username string, role int) bool {
return true
}
+func loadSessionUserIdentity(session sessions.Session) (*model.User, error) {
+ idRaw := session.Get("id")
+ if idRaw == nil {
+ return nil, fmt.Errorf("session is missing user id")
+ }
+ id, ok := idRaw.(int)
+ if !ok || id <= 0 {
+ return nil, fmt.Errorf("session user id is invalid")
+ }
+ return model.GetUserIdentityById(id)
+}
+
+func syncSessionUserIdentity(session sessions.Session, user *model.User) error {
+ if user == nil {
+ return fmt.Errorf("user is nil")
+ }
+
+ storedUsername, _ := session.Get("username").(string)
+ storedRole, _ := session.Get("role").(int)
+ storedStatus, _ := session.Get("status").(int)
+ storedGroup, _ := session.Get("group").(string)
+
+ if storedUsername == user.Username &&
+ storedRole == user.Role &&
+ storedStatus == user.Status &&
+ storedGroup == user.Group {
+ return nil
+ }
+
+ session.Set("id", user.Id)
+ session.Set("username", user.Username)
+ session.Set("role", user.Role)
+ session.Set("status", user.Status)
+ session.Set("group", user.Group)
+
+ return session.Save()
+}
+
func authHelper(c *gin.Context, minRole int) {
session := sessions.Default(c)
- username := session.Get("username")
- role := session.Get("role")
- id := session.Get("id")
- status := session.Get("status")
useAccessToken := false
- if username == nil {
+ currentUsername := ""
+ currentRole := 0
+ currentID := 0
+ currentStatus := 0
+ currentGroup := ""
+
+ if session.Get("username") == nil {
// Check access token
accessToken := c.Request.Header.Get("Authorization")
if accessToken == "" {
@@ -59,10 +99,11 @@ func authHelper(c *gin.Context, minRole int) {
return
}
// Token is valid
- username = user.Username
- role = user.Role
- id = user.Id
- status = user.Status
+ currentUsername = user.Username
+ currentRole = user.Role
+ currentID = user.Id
+ currentStatus = user.Status
+ currentGroup = user.Group
useAccessToken = true
} else {
c.JSON(http.StatusOK, gin.H{
@@ -72,6 +113,29 @@ func authHelper(c *gin.Context, minRole int) {
c.Abort()
return
}
+ } else {
+ user, err := loadSessionUserIdentity(session)
+ if err != nil {
+ c.JSON(http.StatusUnauthorized, gin.H{
+ "success": false,
+ "message": "无权进行此操作,会话信息无效",
+ })
+ c.Abort()
+ return
+ }
+ if err := syncSessionUserIdentity(session, user); err != nil {
+ c.JSON(http.StatusUnauthorized, gin.H{
+ "success": false,
+ "message": "无权进行此操作,会话刷新失败",
+ })
+ c.Abort()
+ return
+ }
+ currentUsername = user.Username
+ currentRole = user.Role
+ currentID = user.Id
+ currentStatus = user.Status
+ currentGroup = user.Group
}
// get header New-Api-User
apiUserIdStr := c.Request.Header.Get("New-Api-User")
@@ -93,7 +157,7 @@ func authHelper(c *gin.Context, minRole int) {
return
}
- if id != apiUserId {
+ if currentID != apiUserId {
c.JSON(http.StatusUnauthorized, gin.H{
"success": false,
"message": "无权进行此操作,New-Api-User 与登录用户不匹配",
@@ -101,7 +165,7 @@ func authHelper(c *gin.Context, minRole int) {
c.Abort()
return
}
- if status.(int) == common.UserStatusDisabled {
+ if currentStatus == common.UserStatusDisabled {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "用户已被封禁",
@@ -109,7 +173,7 @@ func authHelper(c *gin.Context, minRole int) {
c.Abort()
return
}
- if role.(int) < minRole {
+ if currentRole < minRole {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "无权进行此操作,权限不足",
@@ -117,7 +181,7 @@ func authHelper(c *gin.Context, minRole int) {
c.Abort()
return
}
- if !validUserInfo(username.(string), role.(int)) {
+ if !validUserInfo(currentUsername, currentRole) {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "无权进行此操作,用户信息无效",
@@ -127,11 +191,11 @@ func authHelper(c *gin.Context, minRole int) {
}
// 防止不同newapi版本冲突,导致数据不通用
c.Header("Auth-Version", "864b7076dbcd0a3c01b5520316720ebf")
- c.Set("username", username)
- c.Set("role", role)
- c.Set("id", id)
- c.Set("group", session.Get("group"))
- c.Set("user_group", session.Get("group"))
+ c.Set("username", currentUsername)
+ c.Set("role", currentRole)
+ c.Set("id", currentID)
+ c.Set("group", currentGroup)
+ c.Set("user_group", currentGroup)
c.Set("use_access_token", useAccessToken)
c.Next()
diff --git a/middleware/auth_test.go b/middleware/auth_test.go
new file mode 100644
index 000000000000..45bd65008c59
--- /dev/null
+++ b/middleware/auth_test.go
@@ -0,0 +1,278 @@
+package middleware
+
+import (
+ "fmt"
+ "net/http"
+ "net/http/cookiejar"
+ "net/http/httptest"
+ "strconv"
+ "strings"
+ "testing"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-contrib/sessions/cookie"
+ "github.com/gin-gonic/gin"
+ "github.com/glebarez/sqlite"
+ "gorm.io/gorm"
+)
+
+type authMiddlewareAPIResponse struct {
+ Success bool `json:"success"`
+ Message string `json:"message"`
+ Data common.RawMessage `json:"data"`
+}
+
+type authMiddlewareInfoResponse struct {
+ ID int `json:"id"`
+ Username string `json:"username"`
+ Role int `json:"role"`
+ Group string `json:"group"`
+}
+
+func setupAuthMiddlewareTestDB(t *testing.T) {
+ t.Helper()
+
+ oldDB := model.DB
+ oldLogDB := model.LOG_DB
+ oldUsingSQLite := common.UsingSQLite
+ oldUsingMySQL := common.UsingMySQL
+ oldUsingPostgreSQL := common.UsingPostgreSQL
+ oldRedisEnabled := common.RedisEnabled
+
+ gin.SetMode(gin.TestMode)
+ common.UsingSQLite = true
+ common.UsingMySQL = false
+ common.UsingPostgreSQL = false
+ common.RedisEnabled = false
+
+ dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_"))
+ db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{})
+ if err != nil {
+ t.Fatalf("failed to open sqlite db: %v", err)
+ }
+ model.DB = db
+ model.LOG_DB = db
+ if err := db.AutoMigrate(&model.User{}); err != nil {
+ t.Fatalf("failed to migrate users table: %v", err)
+ }
+
+ t.Cleanup(func() {
+ sqlDB, err := db.DB()
+ if err == nil {
+ _ = sqlDB.Close()
+ }
+ model.DB = oldDB
+ model.LOG_DB = oldLogDB
+ common.UsingSQLite = oldUsingSQLite
+ common.UsingMySQL = oldUsingMySQL
+ common.UsingPostgreSQL = oldUsingPostgreSQL
+ common.RedisEnabled = oldRedisEnabled
+ })
+}
+
+func newAuthMiddlewareTestRouter(t *testing.T) *gin.Engine {
+ t.Helper()
+
+ router := gin.New()
+ store := cookie.NewStore([]byte("auth-middleware-test-secret"))
+ router.Use(sessions.Sessions("session", store))
+
+ router.GET("/test/login/:id", func(c *gin.Context) {
+ userID, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()})
+ return
+ }
+ user, err := model.GetUserById(userID, false)
+ if err != nil {
+ c.JSON(http.StatusNotFound, gin.H{"success": false, "message": err.Error()})
+ return
+ }
+ session := sessions.Default(c)
+ session.Set("id", user.Id)
+ session.Set("username", user.Username)
+ session.Set("role", user.Role)
+ session.Set("status", user.Status)
+ session.Set("group", user.Group)
+ if err := session.Save(); err != nil {
+ c.JSON(http.StatusInternalServerError, gin.H{"success": false, "message": err.Error()})
+ return
+ }
+ c.JSON(http.StatusOK, gin.H{"success": true})
+ })
+
+ router.GET("/auth/info", UserAuth(), func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{
+ "success": true,
+ "data": gin.H{
+ "id": c.GetInt("id"),
+ "username": c.GetString("username"),
+ "role": c.GetInt("role"),
+ "group": c.GetString("group"),
+ },
+ })
+ })
+
+ router.GET("/admin/ping", AdminAuth(), func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"success": true, "message": ""})
+ })
+
+ return router
+}
+
+func createAuthMiddlewareTestUser(t *testing.T, username string, role int, status int, group string) *model.User {
+ t.Helper()
+
+ password, err := common.Password2Hash("12345678")
+ if err != nil {
+ t.Fatalf("failed to hash password: %v", err)
+ }
+ user := &model.User{
+ Username: username,
+ Password: password,
+ DisplayName: username,
+ Role: role,
+ Status: status,
+ Group: group,
+ AffCode: username + "-aff",
+ }
+ if err := model.DB.Create(user).Error; err != nil {
+ t.Fatalf("failed to create user: %v", err)
+ }
+ return user
+}
+
+func newAuthMiddlewareTestClient(t *testing.T) *http.Client {
+ t.Helper()
+
+ jar, err := cookiejar.New(nil)
+ if err != nil {
+ t.Fatalf("failed to create cookie jar: %v", err)
+ }
+ return &http.Client{Jar: jar}
+}
+
+func establishAuthMiddlewareSession(t *testing.T, client *http.Client, baseURL string, userID int) {
+ t.Helper()
+
+ response, err := client.Get(fmt.Sprintf("%s/test/login/%d", baseURL, userID))
+ if err != nil {
+ t.Fatalf("failed to establish session: %v", err)
+ }
+ defer response.Body.Close()
+
+ var payload authMiddlewareAPIResponse
+ if err := common.DecodeJson(response.Body, &payload); err != nil {
+ t.Fatalf("failed to decode session response: %v", err)
+ }
+ if !payload.Success {
+ t.Fatalf("expected session setup success, got message: %s", payload.Message)
+ }
+}
+
+func performAuthMiddlewareRequest(t *testing.T, client *http.Client, method string, url string, userID int) authMiddlewareAPIResponse {
+ t.Helper()
+
+ request, err := http.NewRequest(method, url, nil)
+ if err != nil {
+ t.Fatalf("failed to create request: %v", err)
+ }
+ request.Header.Set("New-Api-User", strconv.Itoa(userID))
+
+ response, err := client.Do(request)
+ if err != nil {
+ t.Fatalf("failed to execute request: %v", err)
+ }
+ defer response.Body.Close()
+
+ var payload authMiddlewareAPIResponse
+ if err := common.DecodeJson(response.Body, &payload); err != nil {
+ t.Fatalf("failed to decode api response: %v", err)
+ }
+ return payload
+}
+
+func TestUserAuthRejectsDisabledSessionImmediately(t *testing.T) {
+ setupAuthMiddlewareTestDB(t)
+ user := createAuthMiddlewareTestUser(t, "bob", common.RoleCommonUser, common.UserStatusEnabled, "default")
+
+ router := newAuthMiddlewareTestRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newAuthMiddlewareTestClient(t)
+ establishAuthMiddlewareSession(t, client, server.URL, user.Id)
+
+ user.Status = common.UserStatusDisabled
+ if err := user.Update(false); err != nil {
+ t.Fatalf("failed to disable user: %v", err)
+ }
+
+ response := performAuthMiddlewareRequest(t, client, http.MethodGet, server.URL+"/auth/info", user.Id)
+ if response.Success {
+ t.Fatalf("expected disabled session request to fail")
+ }
+ if response.Message != "用户已被封禁" {
+ t.Fatalf("unexpected error message: %s", response.Message)
+ }
+}
+
+func TestAdminAuthRejectsDowngradedSessionImmediately(t *testing.T) {
+ setupAuthMiddlewareTestDB(t)
+ user := createAuthMiddlewareTestUser(t, "alice", common.RoleAdminUser, common.UserStatusEnabled, "vip")
+
+ router := newAuthMiddlewareTestRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newAuthMiddlewareTestClient(t)
+ establishAuthMiddlewareSession(t, client, server.URL, user.Id)
+
+ user.Role = common.RoleCommonUser
+ if err := user.Update(false); err != nil {
+ t.Fatalf("failed to downgrade user: %v", err)
+ }
+
+ response := performAuthMiddlewareRequest(t, client, http.MethodGet, server.URL+"/admin/ping", user.Id)
+ if response.Success {
+ t.Fatalf("expected downgraded admin request to fail")
+ }
+ if response.Message != "无权进行此操作,权限不足" {
+ t.Fatalf("unexpected error message: %s", response.Message)
+ }
+}
+
+func TestUserAuthRefreshesLatestGroupFromDatabase(t *testing.T) {
+ setupAuthMiddlewareTestDB(t)
+ user := createAuthMiddlewareTestUser(t, "charlie", common.RoleCommonUser, common.UserStatusEnabled, "default")
+
+ router := newAuthMiddlewareTestRouter(t)
+ server := httptest.NewServer(router)
+ defer server.Close()
+
+ client := newAuthMiddlewareTestClient(t)
+ establishAuthMiddlewareSession(t, client, server.URL, user.Id)
+
+ user.Group = "vip"
+ if err := user.Update(false); err != nil {
+ t.Fatalf("failed to update user group: %v", err)
+ }
+
+ response := performAuthMiddlewareRequest(t, client, http.MethodGet, server.URL+"/auth/info", user.Id)
+ if !response.Success {
+ t.Fatalf("expected refreshed session request to succeed, got message: %s", response.Message)
+ }
+
+ var info authMiddlewareInfoResponse
+ if err := common.Unmarshal(response.Data, &info); err != nil {
+ t.Fatalf("failed to decode auth info: %v", err)
+ }
+ if info.Group != "vip" {
+ t.Fatalf("expected refreshed group to be vip, got %s", info.Group)
+ }
+ if info.Role != common.RoleCommonUser {
+ t.Fatalf("expected role to remain common user, got %d", info.Role)
+ }
+}
diff --git a/model/custom_oauth_provider.go b/model/custom_oauth_provider.go
index 12b4d11113d0..ff10c21e2ded 100644
--- a/model/custom_oauth_provider.go
+++ b/model/custom_oauth_provider.go
@@ -3,12 +3,52 @@ package model
import (
"errors"
"fmt"
+ "net/url"
"strings"
"time"
"github.com/QuantumNous/new-api/common"
)
+const (
+ CustomOAuthProviderKindOAuthCode = "oauth_code"
+ CustomOAuthProviderKindJWTDirect = "jwt_direct"
+)
+
+const (
+ CustomJWTSourceQuery = "query"
+ CustomJWTSourceFragment = "fragment"
+ CustomJWTSourceBody = "body"
+)
+
+const (
+ CustomOAuthMappingModeExplicitOnly = "explicit_only"
+ CustomOAuthMappingModeMappingFirst = "mapping_first"
+)
+
+const (
+ CustomJWTAcquireModeDirectToken = "direct_token"
+ CustomJWTAcquireModeTicketExchange = "ticket_exchange"
+ CustomJWTAcquireModeTicketValidate = "ticket_validate"
+)
+
+const (
+ CustomJWTIdentityModeClaims = "claims"
+ CustomJWTIdentityModeUserInfo = "userinfo"
+)
+
+const (
+ CustomTicketExchangeMethodGET = "GET"
+ CustomTicketExchangeMethodPOST = "POST"
+)
+
+const (
+ CustomTicketExchangePayloadModeQuery = "query"
+ CustomTicketExchangePayloadModeForm = "form"
+ CustomTicketExchangePayloadModeJSON = "json"
+ CustomTicketExchangePayloadModeMultipart = "multipart"
+)
+
type accessPolicyPayload struct {
Logic string `json:"logic"`
Conditions []accessConditionItem `json:"conditions"`
@@ -38,23 +78,51 @@ var supportedAccessPolicyOps = map[string]struct{}{
// CustomOAuthProvider stores configuration for custom OAuth providers
type CustomOAuthProvider struct {
- Id int `json:"id" gorm:"primaryKey"`
- 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"
- Icon string `json:"icon" gorm:"type:varchar(128);default:''"` // Icon name from @lobehub/icons
- Enabled bool `json:"enabled" gorm:"default:false"` // Whether this provider is enabled
- ClientId string `json:"client_id" gorm:"type:varchar(256)"` // OAuth client ID
- ClientSecret string `json:"-" gorm:"type:varchar(512)"` // OAuth client secret (not returned to frontend)
- AuthorizationEndpoint string `json:"authorization_endpoint" gorm:"type:varchar(512)"` // Authorization URL
- TokenEndpoint string `json:"token_endpoint" gorm:"type:varchar(512)"` // Token exchange URL
- 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
+ Id int `json:"id" gorm:"primaryKey"`
+ 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"
+ Icon string `json:"icon" gorm:"type:varchar(128);default:''"` // Icon name from @lobehub/icons
+ Kind string `json:"kind" gorm:"type:varchar(32);default:'oauth_code'"` // oauth_code / jwt_direct
+ Enabled bool `json:"enabled" gorm:"default:false"` // Whether this provider is enabled
+ ClientId string `json:"client_id" gorm:"type:varchar(256)"` // OAuth client ID
+ ClientSecret string `json:"-" gorm:"type:varchar(512)"` // OAuth client secret (not returned to frontend)
+ AuthorizationEndpoint string `json:"authorization_endpoint" gorm:"type:varchar(512)"` // Authorization URL
+ TokenEndpoint string `json:"token_endpoint" gorm:"type:varchar(512)"` // Token exchange URL
+ 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
+ Issuer string `json:"issuer" gorm:"type:varchar(512)"` // JWT issuer
+ Audience string `json:"audience" gorm:"type:varchar(256)"` // JWT audience
+ JwksURL string `json:"jwks_url" gorm:"type:varchar(512)"` // JWKS endpoint URL
+ PublicKey string `json:"public_key" gorm:"type:text"` // PEM public key
+ JWTSource string `json:"jwt_source" gorm:"type:varchar(32);default:'query'"` // query / fragment / body
+ JWTHeader string `json:"jwt_header" gorm:"type:varchar(128);default:'Authorization'"` // token header for userinfo mode
+ JWTIdentityMode string `json:"jwt_identity_mode" gorm:"type:varchar(32);default:'claims'"` // claims / userinfo
+ JWTAcquireMode string `json:"jwt_acquire_mode" gorm:"type:varchar(32);default:'direct_token'"` // direct_token / ticket_exchange / ticket_validate
+ AuthorizationServiceField string `json:"authorization_service_field" gorm:"type:varchar(64);default:'service'"` // browser login callback param for ticket exchange
+ TicketExchangeURL string `json:"ticket_exchange_url" gorm:"type:varchar(512)"` // ticket processing endpoint URL
+ TicketExchangeMethod string `json:"ticket_exchange_method" gorm:"type:varchar(16);default:'GET'"` // GET / POST
+ TicketExchangePayloadMode string `json:"ticket_exchange_payload_mode" gorm:"type:varchar(16);default:'query'"` // query / form / json / multipart
+ TicketExchangeTicketField string `json:"ticket_exchange_ticket_field" gorm:"type:varchar(64);default:'ticket'"` // ticket field name
+ TicketExchangeTokenField string `json:"ticket_exchange_token_field" gorm:"type:varchar(128)"` // response token field path (exchange mode)
+ TicketExchangeServiceField string `json:"ticket_exchange_service_field" gorm:"type:varchar(64)"` // optional service field name
+ TicketExchangeExtraParams string `json:"ticket_exchange_extra_params" gorm:"type:text"` // JSON object for exchange params
+ TicketExchangeHeaders string `json:"ticket_exchange_headers" gorm:"type:text"` // JSON object for exchange headers
// 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"
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
EmailField string `json:"email_field" gorm:"type:varchar(128);default:'email'"` // Email field path
+ GroupField string `json:"group_field" gorm:"type:varchar(128)"` // Group field path
+ RoleField string `json:"role_field" gorm:"type:varchar(128)"` // Role field path
+ GroupMapping string `json:"group_mapping" gorm:"type:text"` // JSON object for external->internal group mapping
+ RoleMapping string `json:"role_mapping" gorm:"type:text"` // JSON object for external->internal role mapping
+ AutoRegister bool `json:"auto_register" gorm:"default:false"` // Auto create local user on first login
+ AutoMergeByEmail bool `json:"auto_merge_by_email" gorm:"default:false"` // Merge to existing user by email when no binding exists
+ SyncGroupOnLogin bool `json:"sync_group_on_login" gorm:"default:false"` // Sync group for existing users on external login
+ SyncRoleOnLogin bool `json:"sync_role_on_login" gorm:"default:false"` // Sync role for existing users on external login
+ GroupMappingMode string `json:"group_mapping_mode" gorm:"type:varchar(32);default:'explicit_only'"` // explicit_only / mapping_first
+ RoleMappingMode string `json:"role_mapping_mode" gorm:"type:varchar(32);default:'explicit_only'"` // explicit_only / mapping_first
// Advanced options
WellKnown string `json:"well_known" gorm:"type:varchar(512)"` // OIDC discovery endpoint (optional)
@@ -70,6 +138,65 @@ func (CustomOAuthProvider) TableName() string {
return "custom_oauth_providers"
}
+func (p *CustomOAuthProvider) GetKind() string {
+ kind := strings.TrimSpace(p.Kind)
+ if kind == "" {
+ return CustomOAuthProviderKindOAuthCode
+ }
+ return kind
+}
+
+func (p *CustomOAuthProvider) IsJWTDirect() bool {
+ return p.GetKind() == CustomOAuthProviderKindJWTDirect
+}
+
+func (p *CustomOAuthProvider) IsOAuthCode() bool {
+ return p.GetKind() == CustomOAuthProviderKindOAuthCode
+}
+
+func (p *CustomOAuthProvider) GetJWTAcquireMode() string {
+ mode := normalizeCustomJWTAcquireMode(p.JWTAcquireMode)
+ if mode == "" {
+ return CustomJWTAcquireModeDirectToken
+ }
+ return mode
+}
+
+func (p *CustomOAuthProvider) GetJWTIdentityMode() string {
+ mode := normalizeCustomJWTIdentityMode(p.JWTIdentityMode)
+ if mode == "" {
+ return CustomJWTIdentityModeClaims
+ }
+ return mode
+}
+
+func (p *CustomOAuthProvider) SupportsBrowserLogin() bool {
+ if !p.Enabled {
+ return false
+ }
+ if p.IsOAuthCode() {
+ return strings.TrimSpace(p.AuthorizationEndpoint) != "" && strings.TrimSpace(p.ClientId) != ""
+ }
+ if p.IsJWTDirect() {
+ if p.RequiresTicketAcquire() {
+ return strings.TrimSpace(p.AuthorizationEndpoint) != ""
+ }
+ return strings.TrimSpace(p.AuthorizationEndpoint) != "" &&
+ strings.TrimSpace(p.ClientId) != "" &&
+ p.JWTSource != CustomJWTSourceBody
+ }
+ return false
+}
+
+func (p *CustomOAuthProvider) RequiresTicketAcquire() bool {
+ switch p.GetJWTAcquireMode() {
+ case CustomJWTAcquireModeTicketExchange, CustomJWTAcquireModeTicketValidate:
+ return true
+ default:
+ return false
+ }
+}
+
// GetAllCustomOAuthProviders returns all custom OAuth providers
func GetAllCustomOAuthProviders() ([]*CustomOAuthProvider, error) {
var providers []*CustomOAuthProvider
@@ -161,18 +288,56 @@ func validateCustomOAuthProvider(provider *CustomOAuthProvider) error {
}
}
provider.Slug = slug
-
- if provider.ClientId == "" {
- return errors.New("client ID is required")
- }
- if provider.AuthorizationEndpoint == "" {
- return errors.New("authorization endpoint is required")
+ provider.Kind = strings.TrimSpace(provider.Kind)
+ if provider.Kind == "" {
+ provider.Kind = CustomOAuthProviderKindOAuthCode
}
- if provider.TokenEndpoint == "" {
- return errors.New("token endpoint is required")
+ if provider.Kind != CustomOAuthProviderKindOAuthCode && provider.Kind != CustomOAuthProviderKindJWTDirect {
+ return errors.New("provider kind is invalid")
}
- if provider.UserInfoEndpoint == "" {
- return errors.New("user info endpoint is required")
+
+ if provider.IsOAuthCode() {
+ if provider.ClientId == "" {
+ return errors.New("client ID is required")
+ }
+ if provider.AuthorizationEndpoint == "" {
+ return errors.New("authorization endpoint is required")
+ }
+ if provider.TokenEndpoint == "" {
+ return errors.New("token endpoint is required")
+ }
+ if provider.UserInfoEndpoint == "" {
+ return errors.New("user info endpoint is required")
+ }
+ } else {
+ acquireMode := normalizeCustomJWTAcquireMode(provider.JWTAcquireMode)
+ if acquireMode == "" {
+ return errors.New("jwt_acquire_mode is invalid")
+ }
+ provider.JWTAcquireMode = acquireMode
+ identityMode := normalizeCustomJWTIdentityMode(provider.JWTIdentityMode)
+ if identityMode == "" {
+ return errors.New("jwt_identity_mode is invalid")
+ }
+ provider.JWTIdentityMode = identityMode
+ switch provider.JWTIdentityMode {
+ case CustomJWTIdentityModeClaims:
+ if provider.JWTAcquireMode != CustomJWTAcquireModeTicketValidate {
+ if strings.TrimSpace(provider.Issuer) == "" {
+ return errors.New("issuer is required for jwt_direct providers using claims mode")
+ }
+ if strings.TrimSpace(provider.JwksURL) == "" && strings.TrimSpace(provider.PublicKey) == "" {
+ return errors.New("jwks_url or public_key is required for jwt_direct providers using claims mode")
+ }
+ }
+ case CustomJWTIdentityModeUserInfo:
+ if provider.JWTAcquireMode == CustomJWTAcquireModeTicketValidate {
+ return errors.New("jwt_direct providers using ticket_validate mode only support claims identity mode")
+ }
+ if !isValidAbsoluteHTTPURL(provider.UserInfoEndpoint) {
+ return errors.New("user_info_endpoint is required and must be a valid http/https url for jwt_direct providers using userinfo mode")
+ }
+ }
}
// Set defaults for field mappings if empty
@@ -191,6 +356,73 @@ func validateCustomOAuthProvider(provider *CustomOAuthProvider) error {
if provider.Scopes == "" {
provider.Scopes = "openid profile email"
}
+ if provider.JWTSource == "" {
+ provider.JWTSource = CustomJWTSourceQuery
+ }
+ switch provider.JWTSource {
+ case CustomJWTSourceQuery, CustomJWTSourceFragment, CustomJWTSourceBody:
+ default:
+ return errors.New("jwt_source is invalid")
+ }
+ if strings.TrimSpace(provider.JWTHeader) == "" {
+ provider.JWTHeader = "Authorization"
+ }
+ if strings.TrimSpace(provider.AuthorizationServiceField) == "" {
+ provider.AuthorizationServiceField = "service"
+ }
+ provider.TicketExchangeMethod = normalizeTicketExchangeMethod(provider.TicketExchangeMethod)
+ if provider.TicketExchangeMethod == "" {
+ return errors.New("ticket_exchange_method is invalid")
+ }
+ provider.TicketExchangePayloadMode = normalizeTicketExchangePayloadMode(provider.TicketExchangePayloadMode)
+ if provider.TicketExchangePayloadMode == "" {
+ return errors.New("ticket_exchange_payload_mode is invalid")
+ }
+ if strings.TrimSpace(provider.TicketExchangeTicketField) == "" {
+ provider.TicketExchangeTicketField = "ticket"
+ }
+ if provider.RequiresTicketAcquire() {
+ if strings.TrimSpace(provider.TicketExchangeURL) == "" {
+ return errors.New("ticket_exchange_url is required for ticket-based acquire mode")
+ }
+ if !isValidAbsoluteHTTPURL(provider.TicketExchangeURL) {
+ return errors.New("ticket_exchange_url must be a valid http/https url")
+ }
+ if strings.TrimSpace(provider.TicketExchangeExtraParams) != "" {
+ if err := validateJSONStringObject(provider.TicketExchangeExtraParams); err != nil {
+ return fmt.Errorf("ticket_exchange_extra_params is invalid: %w", err)
+ }
+ }
+ if strings.TrimSpace(provider.TicketExchangeHeaders) != "" {
+ if err := validateJSONStringObject(provider.TicketExchangeHeaders); err != nil {
+ return fmt.Errorf("ticket_exchange_headers is invalid: %w", err)
+ }
+ }
+ }
+ groupMappingMode := normalizeCustomOAuthMappingMode(provider.GroupMappingMode)
+ if groupMappingMode == "" {
+ return errors.New("group_mapping_mode is invalid")
+ }
+ provider.GroupMappingMode = groupMappingMode
+
+ roleMappingMode := normalizeCustomOAuthMappingMode(provider.RoleMappingMode)
+ if roleMappingMode == "" {
+ return errors.New("role_mapping_mode is invalid")
+ }
+ provider.RoleMappingMode = roleMappingMode
+ if strings.TrimSpace(provider.GroupMapping) != "" {
+ if err := validateJSONStringObject(provider.GroupMapping); err != nil {
+ return fmt.Errorf("group_mapping is invalid: %w", err)
+ }
+ }
+ if strings.TrimSpace(provider.RoleMapping) != "" {
+ if err := validateJSONStringObject(provider.RoleMapping); err != nil {
+ return fmt.Errorf("role_mapping is invalid: %w", err)
+ }
+ if err := validateRoleMappingTargets(provider.RoleMapping); err != nil {
+ return fmt.Errorf("role_mapping is invalid: %w", err)
+ }
+ }
if strings.TrimSpace(provider.AccessPolicy) != "" {
var policy accessPolicyPayload
if err := common.UnmarshalJsonStr(provider.AccessPolicy, &policy); err != nil {
@@ -204,6 +436,105 @@ func validateCustomOAuthProvider(provider *CustomOAuthProvider) error {
return nil
}
+func isValidAbsoluteHTTPURL(raw string) bool {
+ parsed, err := url.Parse(strings.TrimSpace(raw))
+ if err != nil || parsed == nil {
+ return false
+ }
+ if parsed.Scheme != "http" && parsed.Scheme != "https" {
+ return false
+ }
+ return strings.TrimSpace(parsed.Host) != ""
+}
+
+func validateJSONStringObject(raw string) error {
+ var payload map[string]any
+ if err := common.UnmarshalJsonStr(raw, &payload); err != nil {
+ return errors.New("must be valid JSON object")
+ }
+ if payload == nil {
+ return errors.New("must be a JSON object")
+ }
+ return nil
+}
+
+func normalizeCustomOAuthMappingMode(raw string) string {
+ switch strings.ToLower(strings.TrimSpace(raw)) {
+ case "", CustomOAuthMappingModeExplicitOnly:
+ return CustomOAuthMappingModeExplicitOnly
+ case CustomOAuthMappingModeMappingFirst:
+ return CustomOAuthMappingModeMappingFirst
+ default:
+ return ""
+ }
+}
+
+func normalizeCustomJWTAcquireMode(raw string) string {
+ switch strings.ToLower(strings.TrimSpace(raw)) {
+ case "", CustomJWTAcquireModeDirectToken:
+ return CustomJWTAcquireModeDirectToken
+ case CustomJWTAcquireModeTicketExchange:
+ return CustomJWTAcquireModeTicketExchange
+ case CustomJWTAcquireModeTicketValidate:
+ return CustomJWTAcquireModeTicketValidate
+ default:
+ return ""
+ }
+}
+
+func normalizeCustomJWTIdentityMode(raw string) string {
+ switch strings.ToLower(strings.TrimSpace(raw)) {
+ case "", CustomJWTIdentityModeClaims:
+ return CustomJWTIdentityModeClaims
+ case CustomJWTIdentityModeUserInfo:
+ return CustomJWTIdentityModeUserInfo
+ default:
+ return ""
+ }
+}
+
+func normalizeTicketExchangeMethod(raw string) string {
+ switch strings.ToUpper(strings.TrimSpace(raw)) {
+ case "", CustomTicketExchangeMethodGET:
+ return CustomTicketExchangeMethodGET
+ case CustomTicketExchangeMethodPOST:
+ return CustomTicketExchangeMethodPOST
+ default:
+ return ""
+ }
+}
+
+func normalizeTicketExchangePayloadMode(raw string) string {
+ switch strings.ToLower(strings.TrimSpace(raw)) {
+ case "", CustomTicketExchangePayloadModeQuery:
+ return CustomTicketExchangePayloadModeQuery
+ case CustomTicketExchangePayloadModeForm:
+ return CustomTicketExchangePayloadModeForm
+ case CustomTicketExchangePayloadModeJSON:
+ return CustomTicketExchangePayloadModeJSON
+ case CustomTicketExchangePayloadModeMultipart:
+ return CustomTicketExchangePayloadModeMultipart
+ default:
+ return ""
+ }
+}
+
+func validateRoleMappingTargets(raw string) error {
+ var payload map[string]any
+ if err := common.UnmarshalJsonStr(raw, &payload); err != nil {
+ return errors.New("must be valid JSON object")
+ }
+ for key, value := range payload {
+ target := strings.ToLower(strings.TrimSpace(fmt.Sprint(value)))
+ switch target {
+ case "common", "user", "member", "1", "admin", "administrator", "10":
+ default:
+ return fmt.Errorf("unsupported role target for key %q", key)
+ }
+ }
+ return nil
+}
+
func validateAccessPolicyPayload(policy *accessPolicyPayload) error {
if policy == nil {
return errors.New("policy is nil")
diff --git a/model/user.go b/model/user.go
index 1210b5435d04..4dd879af2a6d 100644
--- a/model/user.go
+++ b/model/user.go
@@ -303,6 +303,16 @@ func GetUserById(id int, selectAll bool) (*User, error) {
return &user, err
}
+func GetUserIdentityById(id int) (*User, error) {
+ if id == 0 {
+ return nil, errors.New("id 为空!")
+ }
+ user := User{}
+ err := DB.Select("id", "username", "role", "status", "group").
+ First(&user, "id = ?", id).Error
+ return &user, err
+}
+
func GetUserIdByAffCode(affCode string) (int, error) {
if affCode == "" {
return 0, errors.New("affCode 为空!")
diff --git a/oauth/jwt_direct.go b/oauth/jwt_direct.go
new file mode 100644
index 000000000000..2c2a7673228e
--- /dev/null
+++ b/oauth/jwt_direct.go
@@ -0,0 +1,987 @@
+package oauth
+
+import (
+ "bytes"
+ "context"
+ "crypto/ecdsa"
+ "crypto/ed25519"
+ "crypto/elliptic"
+ "crypto/rsa"
+ "encoding/base64"
+ "encoding/xml"
+ "errors"
+ "fmt"
+ "io"
+ "math/big"
+ "mime/multipart"
+ "net/http"
+ "net/url"
+ "strings"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/QuantumNous/new-api/setting"
+ "github.com/QuantumNous/new-api/setting/ratio_setting"
+ "github.com/gin-gonic/gin"
+ "github.com/golang-jwt/jwt/v5"
+ "github.com/tidwall/gjson"
+)
+
+type JWTDirectProvider struct {
+ config *model.CustomOAuthProvider
+}
+
+type JWTDirectIdentity struct {
+ User *OAuthUser
+ ClaimsJSON []byte
+ Group string
+ Role int
+}
+
+type jwksDocument struct {
+ Keys []jwkKey `json:"keys"`
+}
+
+type jwkKey struct {
+ Kty string `json:"kty"`
+ Kid string `json:"kid"`
+ Use string `json:"use"`
+ Alg string `json:"alg"`
+ N string `json:"n"`
+ E string `json:"e"`
+ Crv string `json:"crv"`
+ X string `json:"x"`
+ Y string `json:"y"`
+}
+
+type casServiceResponseEnvelope struct {
+ XMLName xml.Name `xml:"serviceResponse"`
+ AuthenticationSuccess *casAuthenticationSuccess `xml:"authenticationSuccess"`
+ AuthenticationFailure *casAuthenticationFailure `xml:"authenticationFailure"`
+}
+
+type casAuthenticationSuccess struct {
+ User string `xml:"user"`
+ Attributes casAttributes `xml:"attributes"`
+ ProxyGrantingTicket string `xml:"proxyGrantingTicket"`
+ Proxies []string `xml:"proxies>proxy"`
+}
+
+type casAuthenticationFailure struct {
+ Code string `xml:"code,attr"`
+ Message string `xml:",chardata"`
+}
+
+type casAttributes map[string]any
+
+func (a *casAttributes) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
+ values := map[string]any{}
+ for {
+ token, err := d.Token()
+ if err != nil {
+ if errors.Is(err, io.EOF) {
+ break
+ }
+ return err
+ }
+
+ switch elem := token.(type) {
+ case xml.StartElement:
+ var raw string
+ if err := d.DecodeElement(&raw, &elem); err != nil {
+ return err
+ }
+ name := strings.TrimSpace(elem.Name.Local)
+ value := strings.TrimSpace(raw)
+ if name == "" {
+ continue
+ }
+ if existing, ok := values[name]; ok {
+ switch typed := existing.(type) {
+ case string:
+ values[name] = []string{typed, value}
+ case []string:
+ values[name] = append(typed, value)
+ }
+ continue
+ }
+ values[name] = value
+ case xml.EndElement:
+ if elem.Name == start.Name {
+ *a = casAttributes(values)
+ return nil
+ }
+ }
+ }
+ *a = casAttributes(values)
+ return nil
+}
+
+func NewJWTDirectProvider(config *model.CustomOAuthProvider) *JWTDirectProvider {
+ return &JWTDirectProvider{config: config}
+}
+
+func (p *JWTDirectProvider) GetName() string {
+ return p.config.Name
+}
+
+func (p *JWTDirectProvider) IsEnabled() bool {
+ return p.config.Enabled
+}
+
+func (p *JWTDirectProvider) ExchangeToken(ctx context.Context, code string, c *gin.Context) (*OAuthToken, error) {
+ return nil, errors.New("jwt_direct provider does not support authorization code exchange")
+}
+
+func (p *JWTDirectProvider) GetUserInfo(ctx context.Context, token *OAuthToken) (*OAuthUser, error) {
+ return nil, errors.New("jwt_direct provider does not support userinfo fetch")
+}
+
+func (p *JWTDirectProvider) IsUserIDTaken(providerUserID string) bool {
+ return model.IsProviderUserIdTaken(p.config.Id, providerUserID)
+}
+
+func (p *JWTDirectProvider) FillUserByProviderID(user *model.User, providerUserID string) error {
+ foundUser, err := model.GetUserByOAuthBinding(p.config.Id, providerUserID)
+ if err != nil {
+ return err
+ }
+ *user = *foundUser
+ return nil
+}
+
+func (p *JWTDirectProvider) SetProviderUserID(user *model.User, providerUserID string) {
+ // JWT direct providers persist bindings in user_oauth_bindings.
+}
+
+func (p *JWTDirectProvider) GetProviderPrefix() string {
+ return p.config.Slug + "_"
+}
+
+func (p *JWTDirectProvider) GetProviderId() int {
+ return p.config.Id
+}
+
+func (p *JWTDirectProvider) ResolveIdentityFromInput(ctx context.Context, rawToken string, ticket string, callbackURL string, state string) (*JWTDirectIdentity, error) {
+ switch p.config.GetJWTAcquireMode() {
+ case model.CustomJWTAcquireModeTicketExchange:
+ exchangedToken, err := p.exchangeTicketForJWT(ctx, ticket, callbackURL, state)
+ if err != nil {
+ return nil, err
+ }
+ rawToken = exchangedToken
+ case model.CustomJWTAcquireModeTicketValidate:
+ claimsJSON, err := p.validateTicketForClaims(ctx, ticket, callbackURL, state)
+ if err != nil {
+ return nil, err
+ }
+ return p.resolveIdentityFromClaimsJSON(claimsJSON)
+ }
+ return p.ResolveIdentity(ctx, rawToken)
+}
+
+func (p *JWTDirectProvider) ResolveIdentity(ctx context.Context, rawToken string) (*JWTDirectIdentity, error) {
+ tokenString := normalizeJWTToken(rawToken)
+ if tokenString == "" {
+ return nil, errors.New("missing jwt token")
+ }
+
+ var claimsJSON []byte
+ var err error
+ if p.config.GetJWTIdentityMode() == model.CustomJWTIdentityModeUserInfo {
+ claimsJSON, err = p.fetchUserInfoClaims(ctx, tokenString)
+ if err != nil {
+ return nil, err
+ }
+ } else {
+ var claims jwt.MapClaims
+ claims, err = p.parseAndValidateClaims(ctx, tokenString)
+ if err != nil {
+ return nil, err
+ }
+ claimsJSON, err = common.Marshal(claims)
+ if err != nil {
+ return nil, fmt.Errorf("marshal jwt claims failed: %w", err)
+ }
+ }
+
+ return p.resolveIdentityFromClaimsJSON(claimsJSON)
+}
+
+func (p *JWTDirectProvider) resolveIdentityFromClaimsJSON(claimsJSON []byte) (*JWTDirectIdentity, error) {
+ if len(bytes.TrimSpace(claimsJSON)) == 0 {
+ return nil, errors.New("identity claims are empty")
+ }
+
+ userID := firstClaimValue(claimsJSON, p.config.UserIdField)
+ if userID == "" {
+ return nil, errors.New("jwt claims missing external user id")
+ }
+
+ username := firstClaimValue(claimsJSON, p.config.UsernameField)
+ displayName := firstClaimValue(claimsJSON, p.config.DisplayNameField)
+ email := firstClaimValue(claimsJSON, p.config.EmailField)
+
+ policyRaw := strings.TrimSpace(p.config.AccessPolicy)
+ if policyRaw != "" {
+ policy, err := parseAccessPolicy(policyRaw)
+ if err != nil {
+ return nil, fmt.Errorf("invalid access policy configuration: %w", err)
+ }
+ allowed, failure := evaluateAccessPolicy(string(claimsJSON), policy)
+ if !allowed {
+ message := renderAccessDeniedMessage(
+ p.config.AccessDeniedMessage,
+ p.config.Name,
+ string(claimsJSON),
+ failure,
+ )
+ return nil, &AccessDeniedError{Message: message}
+ }
+ }
+
+ return &JWTDirectIdentity{
+ User: &OAuthUser{
+ ProviderUserID: userID,
+ Username: username,
+ DisplayName: displayName,
+ Email: email,
+ Extra: map[string]any{
+ "provider": p.config.Slug,
+ },
+ },
+ ClaimsJSON: claimsJSON,
+ Group: resolveMappedGroup(claimsJSON, p.config),
+ Role: resolveMappedRole(claimsJSON, p.config),
+ }, nil
+}
+
+func (p *JWTDirectProvider) exchangeTicketForJWT(ctx context.Context, ticket string, callbackURL string, state string) (string, error) {
+ responseBody, err := p.performTicketAcquireRequest(ctx, ticket, callbackURL, state)
+ if err != nil {
+ return "", err
+ }
+
+ token := extractExchangedToken(
+ responseBody,
+ p.config.TicketExchangeTokenField,
+ p.config.GetJWTIdentityMode() == model.CustomJWTIdentityModeUserInfo,
+ )
+ if token == "" {
+ return "", errors.New("ticket exchange response missing jwt token")
+ }
+ return token, nil
+}
+
+func (p *JWTDirectProvider) validateTicketForClaims(ctx context.Context, ticket string, callbackURL string, state string) ([]byte, error) {
+ responseBody, err := p.performTicketAcquireRequest(ctx, ticket, callbackURL, state)
+ if err != nil {
+ return nil, err
+ }
+ return parseTicketValidationClaims(responseBody)
+}
+
+func (p *JWTDirectProvider) performTicketAcquireRequest(ctx context.Context, ticket string, callbackURL string, state string) ([]byte, error) {
+ ticket = strings.TrimSpace(ticket)
+ if ticket == "" {
+ return nil, errors.New("missing ticket")
+ }
+
+ targetURL := strings.TrimSpace(p.config.TicketExchangeURL)
+ if targetURL == "" {
+ return nil, errors.New("ticket exchange url is not configured")
+ }
+
+ parsedURL, err := url.Parse(targetURL)
+ if err != nil || parsedURL.Host == "" || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") {
+ return nil, errors.New("ticket exchange url is invalid")
+ }
+
+ params := parseStringMapping(p.config.TicketExchangeExtraParams)
+ headers := parseStringMapping(p.config.TicketExchangeHeaders)
+ ticketField := strings.TrimSpace(p.config.TicketExchangeTicketField)
+ if ticketField == "" {
+ ticketField = "ticket"
+ }
+ params[ticketField] = ticket
+ if serviceField := strings.TrimSpace(p.config.TicketExchangeServiceField); serviceField != "" && strings.TrimSpace(callbackURL) != "" {
+ params[serviceField] = callbackURL
+ }
+
+ placeholderValues := map[string]string{
+ "ticket": ticket,
+ "callback_url": callbackURL,
+ "provider_slug": p.config.Slug,
+ "state": state,
+ }
+ for key, value := range params {
+ params[key] = replaceJWTExchangePlaceholders(value, placeholderValues)
+ }
+ for key, value := range headers {
+ headers[key] = replaceJWTExchangePlaceholders(value, placeholderValues)
+ }
+
+ method := normalizeTicketExchangeMethod(p.config.TicketExchangeMethod)
+ payloadMode := normalizeTicketExchangePayloadMode(p.config.TicketExchangePayloadMode)
+
+ var body io.Reader
+ switch method {
+ case model.CustomTicketExchangeMethodGET:
+ appendExchangeQueryParams(parsedURL, params)
+ case model.CustomTicketExchangeMethodPOST:
+ switch payloadMode {
+ case model.CustomTicketExchangePayloadModeQuery:
+ appendExchangeQueryParams(parsedURL, params)
+ case model.CustomTicketExchangePayloadModeForm:
+ values := url.Values{}
+ for key, value := range params {
+ values.Set(key, value)
+ }
+ body = strings.NewReader(values.Encode())
+ headers["Content-Type"] = "application/x-www-form-urlencoded"
+ case model.CustomTicketExchangePayloadModeJSON:
+ payload, marshalErr := common.Marshal(params)
+ if marshalErr != nil {
+ return nil, fmt.Errorf("marshal ticket exchange payload failed: %w", marshalErr)
+ }
+ body = bytes.NewReader(payload)
+ headers["Content-Type"] = "application/json"
+ case model.CustomTicketExchangePayloadModeMultipart:
+ var buffer bytes.Buffer
+ writer := multipart.NewWriter(&buffer)
+ for key, value := range params {
+ if fieldErr := writer.WriteField(key, value); fieldErr != nil {
+ return nil, fmt.Errorf("build multipart exchange payload failed: %w", fieldErr)
+ }
+ }
+ if closeErr := writer.Close(); closeErr != nil {
+ return nil, fmt.Errorf("close multipart exchange payload failed: %w", closeErr)
+ }
+ body = &buffer
+ headers["Content-Type"] = writer.FormDataContentType()
+ default:
+ return nil, errors.New("ticket exchange payload mode is invalid")
+ }
+ default:
+ return nil, errors.New("ticket exchange method is invalid")
+ }
+
+ req, err := http.NewRequestWithContext(ctx, method, parsedURL.String(), body)
+ if err != nil {
+ return nil, err
+ }
+ req.Header.Set("Accept", "application/json, text/plain, */*")
+ for key, value := range headers {
+ if strings.TrimSpace(key) == "" {
+ continue
+ }
+ req.Header.Set(key, value)
+ }
+
+ client := &http.Client{Timeout: 10 * time.Second}
+ resp, err := client.Do(req)
+ if err != nil {
+ return nil, err
+ }
+ defer resp.Body.Close()
+
+ responseBody, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if err != nil {
+ return nil, err
+ }
+ if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
+ return nil, fmt.Errorf("ticket acquire failed: %s %s", resp.Status, strings.TrimSpace(string(responseBody)))
+ }
+ return responseBody, nil
+}
+
+func (p *JWTDirectProvider) fetchUserInfoClaims(ctx context.Context, tokenString string) ([]byte, error) {
+ targetURL := strings.TrimSpace(p.config.UserInfoEndpoint)
+ if targetURL == "" {
+ return nil, errors.New("userinfo endpoint is not configured")
+ }
+
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, targetURL, nil)
+ if err != nil {
+ return nil, err
+ }
+ req.Header.Set("Accept", "application/json, text/plain, */*")
+
+ headerName := strings.TrimSpace(p.config.JWTHeader)
+ if headerName == "" {
+ headerName = "Authorization"
+ }
+ headerValue := tokenString
+ if strings.EqualFold(headerName, "Authorization") {
+ headerValue = "Bearer " + tokenString
+ }
+ req.Header.Set(headerName, headerValue)
+
+ client := &http.Client{Timeout: 10 * time.Second}
+ resp, err := client.Do(req)
+ if err != nil {
+ return nil, err
+ }
+ defer resp.Body.Close()
+
+ body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if err != nil {
+ return nil, err
+ }
+ if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
+ return nil, fmt.Errorf("userinfo request failed: %s %s", resp.Status, strings.TrimSpace(string(body)))
+ }
+ if len(bytes.TrimSpace(body)) == 0 {
+ return nil, errors.New("userinfo response is empty")
+ }
+ return body, nil
+}
+
+func (p *JWTDirectProvider) parseAndValidateClaims(ctx context.Context, tokenString string) (jwt.MapClaims, error) {
+ claims := jwt.MapClaims{}
+ token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (any, error) {
+ if token.Method == nil || token.Method.Alg() == "" || token.Method.Alg() == "none" {
+ return nil, errors.New("unsupported jwt signing algorithm")
+ }
+ return p.resolveVerificationKey(ctx, token)
+ })
+ if err != nil {
+ return nil, fmt.Errorf("jwt verification failed: %w", err)
+ }
+ if token == nil || !token.Valid {
+ return nil, errors.New("jwt token is invalid")
+ }
+
+ if issuer := strings.TrimSpace(p.config.Issuer); issuer != "" {
+ gotIssuer, err := claims.GetIssuer()
+ if err != nil || strings.TrimSpace(gotIssuer) != issuer {
+ return nil, errors.New("jwt issuer mismatch")
+ }
+ }
+ if audience := strings.TrimSpace(p.config.Audience); audience != "" {
+ gotAudience, err := claims.GetAudience()
+ if err != nil || !stringSliceContains(gotAudience, audience) {
+ return nil, errors.New("jwt audience mismatch")
+ }
+ }
+ return claims, nil
+}
+
+func (p *JWTDirectProvider) resolveVerificationKey(ctx context.Context, token *jwt.Token) (any, error) {
+ if strings.TrimSpace(p.config.PublicKey) != "" {
+ return parsePEMPublicKey(p.config.PublicKey)
+ }
+ if strings.TrimSpace(p.config.JwksURL) == "" {
+ return nil, errors.New("jwt verification key is not configured")
+ }
+ return fetchJWKSKey(ctx, p.config.JwksURL, token)
+}
+
+func normalizeJWTToken(raw string) string {
+ value := strings.TrimSpace(raw)
+ if strings.HasPrefix(strings.ToLower(value), "bearer ") {
+ value = strings.TrimSpace(value[7:])
+ }
+ return value
+}
+
+func appendExchangeQueryParams(targetURL *url.URL, params map[string]string) {
+ query := targetURL.Query()
+ for key, value := range params {
+ if strings.TrimSpace(key) == "" {
+ continue
+ }
+ query.Set(key, value)
+ }
+ targetURL.RawQuery = query.Encode()
+}
+
+func replaceJWTExchangePlaceholders(input string, values map[string]string) string {
+ result := input
+ for key, value := range values {
+ result = strings.ReplaceAll(result, "{"+key+"}", value)
+ }
+ return result
+}
+
+func extractExchangedToken(body []byte, tokenField string, allowOpaque bool) string {
+ if len(body) == 0 {
+ return ""
+ }
+
+ candidates := []string{}
+ if strings.TrimSpace(tokenField) != "" {
+ candidates = append(candidates, strings.TrimSpace(tokenField))
+ }
+ candidates = append(candidates,
+ "token",
+ "access_token",
+ "data.token",
+ "data.access_token",
+ "result.token",
+ "result.access_token",
+ "data",
+ )
+
+ for _, candidate := range candidates {
+ result := gjson.GetBytes(body, candidate)
+ if result.Exists() {
+ value := normalizeJWTToken(result.String())
+ if looksLikeJWT(value) || (allowOpaque && value != "") {
+ return value
+ }
+ }
+ }
+
+ trimmed := normalizeJWTToken(string(bytes.TrimSpace(body)))
+ if looksLikeJWT(trimmed) || (allowOpaque && trimmed != "") {
+ return trimmed
+ }
+
+ var payload any
+ if err := common.Unmarshal(body, &payload); err == nil {
+ if str, ok := payload.(string); ok {
+ value := normalizeJWTToken(str)
+ if looksLikeJWT(value) || (allowOpaque && value != "") {
+ return value
+ }
+ }
+ }
+
+ return ""
+}
+
+func parseTicketValidationClaims(body []byte) ([]byte, error) {
+ trimmed := bytes.TrimSpace(body)
+ if len(trimmed) == 0 {
+ return nil, errors.New("ticket validation response is empty")
+ }
+ if trimmed[0] == '<' {
+ return parseTicketValidationXML(trimmed)
+ }
+ return parseTicketValidationJSON(trimmed)
+}
+
+func parseTicketValidationJSON(body []byte) ([]byte, error) {
+ var payload map[string]any
+ if err := common.Unmarshal(body, &payload); err != nil {
+ return nil, fmt.Errorf("parse ticket validation json failed: %w", err)
+ }
+ if payload == nil {
+ return nil, errors.New("ticket validation response is empty")
+ }
+
+ serviceResponse := map[string]any{}
+ if raw, ok := payload["serviceResponse"].(map[string]any); ok {
+ serviceResponse = raw
+ } else {
+ if success, ok := payload["authenticationSuccess"]; ok {
+ serviceResponse["authenticationSuccess"] = success
+ }
+ if failure, ok := payload["authenticationFailure"]; ok {
+ serviceResponse["authenticationFailure"] = failure
+ }
+ }
+
+ if failure, ok := serviceResponse["authenticationFailure"]; ok {
+ return nil, formatTicketValidationFailure(failure)
+ }
+ if len(serviceResponse) == 0 {
+ return body, nil
+ }
+
+ normalized := map[string]any{
+ "serviceResponse": serviceResponse,
+ }
+ if success, ok := serviceResponse["authenticationSuccess"]; ok {
+ normalized["authenticationSuccess"] = success
+ }
+ if failure, ok := serviceResponse["authenticationFailure"]; ok {
+ normalized["authenticationFailure"] = failure
+ }
+ claimsJSON, err := common.Marshal(normalized)
+ if err != nil {
+ return nil, fmt.Errorf("marshal ticket validation claims failed: %w", err)
+ }
+ return claimsJSON, nil
+}
+
+func parseTicketValidationXML(body []byte) ([]byte, error) {
+ var envelope casServiceResponseEnvelope
+ if err := xml.Unmarshal(body, &envelope); err != nil {
+ return nil, fmt.Errorf("parse ticket validation xml failed: %w", err)
+ }
+ if envelope.AuthenticationFailure != nil {
+ return nil, formatTicketValidationFailure(map[string]any{
+ "code": strings.TrimSpace(envelope.AuthenticationFailure.Code),
+ "message": strings.TrimSpace(envelope.AuthenticationFailure.Message),
+ })
+ }
+ if envelope.AuthenticationSuccess == nil {
+ return nil, errors.New("ticket validation response missing authenticationSuccess")
+ }
+
+ success := map[string]any{
+ "user": strings.TrimSpace(envelope.AuthenticationSuccess.User),
+ }
+ if len(envelope.AuthenticationSuccess.Attributes) > 0 {
+ success["attributes"] = map[string]any(envelope.AuthenticationSuccess.Attributes)
+ }
+ if pgt := strings.TrimSpace(envelope.AuthenticationSuccess.ProxyGrantingTicket); pgt != "" {
+ success["proxyGrantingTicket"] = pgt
+ }
+ if len(envelope.AuthenticationSuccess.Proxies) > 0 {
+ success["proxies"] = envelope.AuthenticationSuccess.Proxies
+ }
+
+ normalized := map[string]any{
+ "serviceResponse": map[string]any{
+ "authenticationSuccess": success,
+ },
+ "authenticationSuccess": success,
+ }
+ claimsJSON, err := common.Marshal(normalized)
+ if err != nil {
+ return nil, fmt.Errorf("marshal ticket validation claims failed: %w", err)
+ }
+ return claimsJSON, nil
+}
+
+func formatTicketValidationFailure(raw any) error {
+ switch typed := raw.(type) {
+ case map[string]any:
+ code := strings.TrimSpace(fmt.Sprint(typed["code"]))
+ message := strings.TrimSpace(fmt.Sprint(typed["message"]))
+ if message == "" {
+ message = strings.TrimSpace(fmt.Sprint(typed["description"]))
+ }
+ if code != "" && message != "" {
+ return fmt.Errorf("ticket validation failed: %s: %s", code, message)
+ }
+ if code != "" {
+ return fmt.Errorf("ticket validation failed: %s", code)
+ }
+ if message != "" {
+ return fmt.Errorf("ticket validation failed: %s", message)
+ }
+ case string:
+ if message := strings.TrimSpace(typed); message != "" {
+ return fmt.Errorf("ticket validation failed: %s", message)
+ }
+ }
+ return errors.New("ticket validation failed")
+}
+
+func looksLikeJWT(raw string) bool {
+ parts := strings.Split(strings.TrimSpace(raw), ".")
+ return len(parts) == 3 && parts[0] != "" && parts[1] != "" && parts[2] != ""
+}
+
+func normalizeTicketExchangeMethod(raw string) string {
+ switch strings.ToUpper(strings.TrimSpace(raw)) {
+ case "", model.CustomTicketExchangeMethodGET:
+ return model.CustomTicketExchangeMethodGET
+ case model.CustomTicketExchangeMethodPOST:
+ return model.CustomTicketExchangeMethodPOST
+ default:
+ return ""
+ }
+}
+
+func normalizeTicketExchangePayloadMode(raw string) string {
+ switch strings.ToLower(strings.TrimSpace(raw)) {
+ case "", model.CustomTicketExchangePayloadModeQuery:
+ return model.CustomTicketExchangePayloadModeQuery
+ case model.CustomTicketExchangePayloadModeForm:
+ return model.CustomTicketExchangePayloadModeForm
+ case model.CustomTicketExchangePayloadModeJSON:
+ return model.CustomTicketExchangePayloadModeJSON
+ case model.CustomTicketExchangePayloadModeMultipart:
+ return model.CustomTicketExchangePayloadModeMultipart
+ default:
+ return ""
+ }
+}
+
+func parsePEMPublicKey(raw string) (any, error) {
+ pemData := []byte(strings.TrimSpace(raw))
+ if len(pemData) == 0 {
+ return nil, errors.New("empty public key")
+ }
+ if key, err := jwt.ParseRSAPublicKeyFromPEM(pemData); err == nil {
+ return key, nil
+ }
+ if key, err := jwt.ParseECPublicKeyFromPEM(pemData); err == nil {
+ return key, nil
+ }
+ if key, err := jwt.ParseEdPublicKeyFromPEM(pemData); err == nil {
+ return key, nil
+ }
+ return nil, errors.New("unsupported public key format")
+}
+
+func fetchJWKSKey(ctx context.Context, jwksURL string, token *jwt.Token) (any, error) {
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, jwksURL, nil)
+ if err != nil {
+ return nil, err
+ }
+ req.Header.Set("Accept", "application/json")
+
+ client := &http.Client{Timeout: 10 * time.Second}
+ resp, err := client.Do(req)
+ if err != nil {
+ return nil, err
+ }
+ defer resp.Body.Close()
+ if resp.StatusCode != http.StatusOK {
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
+ return nil, fmt.Errorf("jwks request failed: %s %s", resp.Status, strings.TrimSpace(string(body)))
+ }
+
+ var doc jwksDocument
+ if err := common.DecodeJson(resp.Body, &doc); err != nil {
+ return nil, err
+ }
+ if len(doc.Keys) == 0 {
+ return nil, errors.New("jwks document has no keys")
+ }
+
+ selected, err := selectJWK(doc.Keys, token)
+ if err != nil {
+ return nil, err
+ }
+ return jwkToPublicKey(selected)
+}
+
+func selectJWK(keys []jwkKey, token *jwt.Token) (*jwkKey, error) {
+ kid, _ := token.Header["kid"].(string)
+ alg := ""
+ if token.Method != nil {
+ alg = token.Method.Alg()
+ }
+ if kid != "" {
+ for i := range keys {
+ if keys[i].Kid == kid && (keys[i].Use == "" || keys[i].Use == "sig") {
+ return &keys[i], nil
+ }
+ }
+ return nil, fmt.Errorf("jwks key with kid %q not found", kid)
+ }
+ for i := range keys {
+ if keys[i].Use != "" && keys[i].Use != "sig" {
+ continue
+ }
+ if keys[i].Alg == "" || alg == "" || keys[i].Alg == alg {
+ return &keys[i], nil
+ }
+ }
+ if len(keys) == 1 {
+ return &keys[0], nil
+ }
+ return nil, errors.New("unable to select jwks key")
+}
+
+func jwkToPublicKey(key *jwkKey) (any, error) {
+ switch key.Kty {
+ case "RSA":
+ return jwkToRSAPublicKey(key)
+ case "EC":
+ return jwkToECPublicKey(key)
+ case "OKP":
+ return jwkToEd25519PublicKey(key)
+ default:
+ return nil, fmt.Errorf("unsupported jwk key type: %s", key.Kty)
+ }
+}
+
+func jwkToRSAPublicKey(key *jwkKey) (*rsa.PublicKey, error) {
+ nBytes, err := base64.RawURLEncoding.DecodeString(key.N)
+ if err != nil {
+ return nil, fmt.Errorf("decode rsa modulus failed: %w", err)
+ }
+ eBytes, err := base64.RawURLEncoding.DecodeString(key.E)
+ if err != nil {
+ return nil, fmt.Errorf("decode rsa exponent failed: %w", err)
+ }
+ n := new(big.Int).SetBytes(nBytes)
+ e := new(big.Int).SetBytes(eBytes)
+ return &rsa.PublicKey{N: n, E: int(e.Int64())}, nil
+}
+
+func jwkToECPublicKey(key *jwkKey) (*ecdsa.PublicKey, error) {
+ var curve elliptic.Curve
+ switch key.Crv {
+ case "P-256":
+ curve = elliptic.P256()
+ case "P-384":
+ curve = elliptic.P384()
+ case "P-521":
+ curve = elliptic.P521()
+ default:
+ return nil, fmt.Errorf("unsupported ec curve: %s", key.Crv)
+ }
+ xBytes, err := base64.RawURLEncoding.DecodeString(key.X)
+ if err != nil {
+ return nil, fmt.Errorf("decode ec x failed: %w", err)
+ }
+ yBytes, err := base64.RawURLEncoding.DecodeString(key.Y)
+ if err != nil {
+ return nil, fmt.Errorf("decode ec y failed: %w", err)
+ }
+ return &ecdsa.PublicKey{
+ Curve: curve,
+ X: new(big.Int).SetBytes(xBytes),
+ Y: new(big.Int).SetBytes(yBytes),
+ }, nil
+}
+
+func jwkToEd25519PublicKey(key *jwkKey) (ed25519.PublicKey, error) {
+ xBytes, err := base64.RawURLEncoding.DecodeString(key.X)
+ if err != nil {
+ return nil, fmt.Errorf("decode ed25519 key failed: %w", err)
+ }
+ return ed25519.PublicKey(xBytes), nil
+}
+
+func firstClaimValue(claimsJSON []byte, path string) string {
+ values := extractClaimCandidates(claimsJSON, path)
+ if len(values) == 0 {
+ return ""
+ }
+ return values[0]
+}
+
+func extractClaimCandidates(claimsJSON []byte, path string) []string {
+ path = strings.TrimSpace(path)
+ if path == "" {
+ return nil
+ }
+ result := gjson.GetBytes(claimsJSON, path)
+ if !result.Exists() {
+ return nil
+ }
+ if result.IsArray() {
+ candidates := make([]string, 0, len(result.Array()))
+ for _, item := range result.Array() {
+ value := strings.TrimSpace(item.String())
+ if value != "" {
+ candidates = append(candidates, value)
+ }
+ }
+ return candidates
+ }
+ value := strings.TrimSpace(result.String())
+ if value == "" {
+ return nil
+ }
+ return []string{value}
+}
+
+func resolveMappedGroup(claimsJSON []byte, config *model.CustomOAuthProvider) string {
+ candidates := extractClaimCandidates(claimsJSON, config.GroupField)
+ if len(candidates) == 0 {
+ return ""
+ }
+ mapping := parseStringMapping(config.GroupMapping)
+ for _, candidate := range candidates {
+ if mapped, ok := mapping[candidate]; ok {
+ if isExistingGroup(mapped) {
+ return mapped
+ }
+ continue
+ }
+ if isMappingFirstMode(config.GroupMappingMode) && isExistingGroup(candidate) {
+ return candidate
+ }
+ }
+ return ""
+}
+
+func resolveMappedRole(claimsJSON []byte, config *model.CustomOAuthProvider) int {
+ candidates := extractClaimCandidates(claimsJSON, config.RoleField)
+ if len(candidates) == 0 {
+ return 0
+ }
+ mapping := parseStringMapping(config.RoleMapping)
+ for _, candidate := range candidates {
+ if mapped, ok := mapping[candidate]; ok {
+ if role := parseRoleValue(mapped); role != 0 {
+ return role
+ }
+ continue
+ }
+ if isMappingFirstMode(config.RoleMappingMode) {
+ if role := parseRoleValue(candidate); role != 0 {
+ return role
+ }
+ }
+ }
+ return 0
+}
+
+func isMappingFirstMode(mode string) bool {
+ return strings.EqualFold(strings.TrimSpace(mode), model.CustomOAuthMappingModeMappingFirst)
+}
+
+func parseRoleValue(raw string) int {
+ switch strings.ToLower(strings.TrimSpace(raw)) {
+ case "common", "user", "member", "1":
+ return common.RoleCommonUser
+ case "admin", "administrator", "10":
+ return common.RoleAdminUser
+ default:
+ return 0
+ }
+}
+
+func isSyncableRole(role int) bool {
+ switch role {
+ case common.RoleCommonUser, common.RoleAdminUser:
+ return true
+ default:
+ return false
+ }
+}
+
+func parseStringMapping(raw string) map[string]string {
+ payload := make(map[string]any)
+ if strings.TrimSpace(raw) == "" {
+ return map[string]string{}
+ }
+ if err := common.UnmarshalJsonStr(raw, &payload); err != nil {
+ common.SysError("failed to parse custom auth mapping: " + err.Error())
+ return map[string]string{}
+ }
+ result := make(map[string]string, len(payload))
+ for key, value := range payload {
+ trimmedKey := strings.TrimSpace(key)
+ trimmedValue := strings.TrimSpace(fmt.Sprint(value))
+ if trimmedKey == "" || trimmedValue == "" {
+ continue
+ }
+ result[trimmedKey] = trimmedValue
+ }
+ return result
+}
+
+func isExistingGroup(group string) bool {
+ group = strings.TrimSpace(group)
+ if group == "" {
+ return false
+ }
+ if setting.ContainsAutoGroup(group) {
+ return true
+ }
+ _, ok := ratio_setting.GetGroupRatioCopy()[group]
+ return ok
+}
+
+func stringSliceContains(values []string, target string) bool {
+ for _, value := range values {
+ if strings.TrimSpace(value) == target {
+ return true
+ }
+ }
+ return false
+}
diff --git a/oauth/jwt_direct_test.go b/oauth/jwt_direct_test.go
new file mode 100644
index 000000000000..9bf4afc93930
--- /dev/null
+++ b/oauth/jwt_direct_test.go
@@ -0,0 +1,1159 @@
+package oauth
+
+import (
+ "context"
+ "crypto/rand"
+ "crypto/rsa"
+ "crypto/x509"
+ "encoding/base64"
+ "encoding/pem"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/QuantumNous/new-api/common"
+ "github.com/QuantumNous/new-api/model"
+ "github.com/golang-jwt/jwt/v5"
+)
+
+func TestJWTDirectResolveIdentityWithPEMMapping(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ UsernameField: "preferred_username",
+ DisplayNameField: "name",
+ EmailField: "email",
+ GroupField: "groups",
+ GroupMapping: `{"engineering":"vip"}`,
+ RoleField: "roles",
+ RoleMapping: `{"platform-admin":"admin","root":"root"}`,
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-1",
+ "preferred_username": "alice",
+ "name": "Alice",
+ "email": "alice@example.com",
+ "groups": []string{"engineering"},
+ "roles": []string{"platform-admin", "root"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), token)
+ if err != nil {
+ t.Fatalf("expected identity to resolve, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "external-user-1" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+ if identity.User.Username != "alice" {
+ t.Fatalf("unexpected username: %s", identity.User.Username)
+ }
+ if identity.Group != "vip" {
+ t.Fatalf("expected mapped group vip, got %s", identity.Group)
+ }
+ if identity.Role != common.RoleAdminUser {
+ t.Fatalf("expected admin role, got %d", identity.Role)
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsIssuerMismatch(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://other-issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-2",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ _, err := provider.ResolveIdentity(context.Background(), token)
+ if err == nil {
+ t.Fatal("expected issuer mismatch to fail")
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsAudienceMismatch(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "other-audience",
+ "sub": "external-user-2",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ _, err := provider.ResolveIdentity(context.Background(), token)
+ if err == nil {
+ t.Fatal("expected audience mismatch to fail")
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsExpiredToken(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-expired",
+ "exp": time.Now().Add(-time.Minute).Unix(),
+ })
+
+ _, err := provider.ResolveIdentity(context.Background(), token)
+ if err == nil {
+ t.Fatal("expected expired token to fail")
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsInvalidSignature(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ otherPrivateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ })
+
+ token := mustSignJWT(t, otherPrivateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-invalid-signature",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ _, err := provider.ResolveIdentity(context.Background(), token)
+ if err == nil {
+ t.Fatal("expected invalid signature to fail")
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsMissingExternalID(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ _, err := provider.ResolveIdentity(context.Background(), token)
+ if err == nil {
+ t.Fatal("expected missing external id to fail")
+ }
+}
+
+func TestJWTDirectResolveIdentityDoesNotPromoteRootRoleOrInvalidGroup(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ GroupField: "groups",
+ GroupMapping: `{"engineering":"nonexistent-group"}`,
+ RoleField: "roles",
+ RoleMapping: `{"root-role":"root"}`,
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-4",
+ "groups": []string{"engineering", "totally-unknown"},
+ "roles": []string{"root-role"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), token)
+ if err != nil {
+ t.Fatalf("expected identity resolution to succeed, got error: %v", err)
+ }
+ if identity.Group != "" {
+ t.Fatalf("expected invalid group mapping to be ignored, got %s", identity.Group)
+ }
+ if identity.Role != 0 {
+ t.Fatalf("expected root-like claim not to be promoted, got %d", identity.Role)
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsDirectPassThroughByDefault(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ GroupField: "groups",
+ RoleField: "roles",
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-default-mode",
+ "groups": []string{"default"},
+ "roles": []string{"admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), token)
+ if err != nil {
+ t.Fatalf("expected identity resolution to succeed, got error: %v", err)
+ }
+ if identity.Group != "" {
+ t.Fatalf("expected direct group pass-through to be disabled by default, got %s", identity.Group)
+ }
+ if identity.Role != 0 {
+ t.Fatalf("expected direct role pass-through to be disabled by default, got %d", identity.Role)
+ }
+}
+
+func TestJWTDirectResolveIdentityAllowsPassThroughInMappingFirstMode(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ GroupField: "groups",
+ RoleField: "roles",
+ GroupMappingMode: model.CustomOAuthMappingModeMappingFirst,
+ RoleMappingMode: model.CustomOAuthMappingModeMappingFirst,
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-mapping-first",
+ "groups": []string{"default"},
+ "roles": []string{"admin"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), token)
+ if err != nil {
+ t.Fatalf("expected identity resolution to succeed, got error: %v", err)
+ }
+ if identity.Group != "default" {
+ t.Fatalf("expected mapping_first group pass-through to use default, got %s", identity.Group)
+ }
+ if identity.Role != common.RoleAdminUser {
+ t.Fatalf("expected mapping_first role pass-through to use admin, got %d", identity.Role)
+ }
+}
+
+func TestJWTDirectResolveIdentityRejectsGuestRoleTargets(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ RoleField: "roles",
+ RoleMapping: `{"member":"guest"}`,
+ RoleMappingMode: model.CustomOAuthMappingModeMappingFirst,
+ })
+
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-guest",
+ "roles": []string{"member", "guest"},
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), token)
+ if err != nil {
+ t.Fatalf("expected identity resolution to succeed, got error: %v", err)
+ }
+ if identity.Role != 0 {
+ t.Fatalf("expected guest role targets to be ignored, got %d", identity.Role)
+ }
+}
+
+func TestJWTDirectResolveIdentityWithJWKS(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ jwksServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "keys": []map[string]any{
+ {
+ "kty": "RSA",
+ "kid": "kid-1",
+ "use": "sig",
+ "alg": "RS256",
+ "n": base64.RawURLEncoding.EncodeToString(privateKey.PublicKey.N.Bytes()),
+ "e": base64.RawURLEncoding.EncodeToString(bigEndianExponent(privateKey.PublicKey.E)),
+ },
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal jwks payload: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer jwksServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ JwksURL: jwksServer.URL,
+ UserIdField: "sub",
+ })
+
+ token := mustSignJWT(t, privateKey, "kid-1", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "external-user-3",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), token)
+ if err != nil {
+ t.Fatalf("expected jwks validation to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "external-user-3" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+}
+
+func TestJWTDirectResolveIdentityWithUserInfoMode(t *testing.T) {
+ userInfoServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if got := r.Header.Get("x-access-token"); got != "opaque-token" {
+ t.Fatalf("expected raw token in x-access-token header, got %q", got)
+ }
+ payload, err := common.Marshal(map[string]any{
+ "info": map[string]any{
+ "userCode": "1410833903245320192",
+ "loginid": "liangmingsen",
+ "userName": "梁明森",
+ "mailbox": "liangmingsen@qdama.cn",
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal userinfo payload: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer userInfoServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Qdama SSO",
+ Slug: "qdama-sso",
+ Enabled: true,
+ JWTIdentityMode: model.CustomJWTIdentityModeUserInfo,
+ UserInfoEndpoint: userInfoServer.URL,
+ JWTHeader: "x-access-token",
+ UserIdField: "info.userCode",
+ UsernameField: "info.loginid",
+ DisplayNameField: "info.userName",
+ EmailField: "info.mailbox",
+ })
+
+ identity, err := provider.ResolveIdentity(context.Background(), "opaque-token")
+ if err != nil {
+ t.Fatalf("expected userinfo mode identity resolution to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "1410833903245320192" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+ if identity.User.Username != "liangmingsen" {
+ t.Fatalf("unexpected username: %s", identity.User.Username)
+ }
+ if identity.User.DisplayName != "梁明森" {
+ t.Fatalf("unexpected display name: %s", identity.User.DisplayName)
+ }
+ if identity.User.Email != "liangmingsen@qdama.cn" {
+ t.Fatalf("unexpected email: %s", identity.User.Email)
+ }
+}
+
+func TestJWTDirectPerformTicketAcquireRequestSupportsConfiguredMethodsAndPayloadModes(t *testing.T) {
+ testCases := []struct {
+ name string
+ method string
+ payloadMode string
+ }{
+ {name: "get query", method: model.CustomTicketExchangeMethodGET, payloadMode: model.CustomTicketExchangePayloadModeQuery},
+ {name: "post query", method: model.CustomTicketExchangeMethodPOST, payloadMode: model.CustomTicketExchangePayloadModeQuery},
+ {name: "post form", method: model.CustomTicketExchangeMethodPOST, payloadMode: model.CustomTicketExchangePayloadModeForm},
+ {name: "post json", method: model.CustomTicketExchangeMethodPOST, payloadMode: model.CustomTicketExchangePayloadModeJSON},
+ {name: "post multipart", method: model.CustomTicketExchangeMethodPOST, payloadMode: model.CustomTicketExchangePayloadModeMultipart},
+ }
+
+ for _, testCase := range testCases {
+ t.Run(testCase.name, func(t *testing.T) {
+ callbackURL := "https://new-api.example.com/oauth/acme-sso?state=state-123"
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != testCase.method {
+ t.Fatalf("expected method %s, got %s", testCase.method, r.Method)
+ }
+ if got := r.Header.Get("X-State"); got != "acme-sso:state-123" {
+ t.Fatalf("expected X-State header, got %q", got)
+ }
+ if got := r.Header.Get("X-Ticket"); got != "ticket-123" {
+ t.Fatalf("expected X-Ticket header, got %q", got)
+ }
+
+ params := map[string]string{}
+ switch {
+ case testCase.method == model.CustomTicketExchangeMethodGET || testCase.payloadMode == model.CustomTicketExchangePayloadModeQuery:
+ for key, values := range r.URL.Query() {
+ if len(values) > 0 {
+ params[key] = values[0]
+ }
+ }
+ case testCase.payloadMode == model.CustomTicketExchangePayloadModeForm:
+ if err := r.ParseForm(); err != nil {
+ t.Fatalf("failed to parse form payload: %v", err)
+ }
+ for key, values := range r.PostForm {
+ if len(values) > 0 {
+ params[key] = values[0]
+ }
+ }
+ case testCase.payloadMode == model.CustomTicketExchangePayloadModeJSON:
+ if err := common.DecodeJson(r.Body, ¶ms); err != nil {
+ t.Fatalf("failed to decode json payload: %v", err)
+ }
+ case testCase.payloadMode == model.CustomTicketExchangePayloadModeMultipart:
+ if err := r.ParseMultipartForm(1 << 20); err != nil {
+ t.Fatalf("failed to parse multipart payload: %v", err)
+ }
+ for key, values := range r.MultipartForm.Value {
+ if len(values) > 0 {
+ params[key] = values[0]
+ }
+ }
+ default:
+ t.Fatalf("unexpected payload mode %s", testCase.payloadMode)
+ }
+
+ if got := params["st"]; got != "ticket-123" {
+ t.Fatalf("expected st=ticket-123, got %q", got)
+ }
+ if got := params["svc"]; got != callbackURL {
+ t.Fatalf("expected svc=%q, got %q", callbackURL, got)
+ }
+ if got := params["source"]; got != "acme-sso:state-123" {
+ t.Fatalf("expected source placeholder expansion, got %q", got)
+ }
+ if got := params["raw_callback"]; got != callbackURL {
+ t.Fatalf("expected raw_callback placeholder expansion, got %q", got)
+ }
+
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(`{"ok":true}`))
+ }))
+ defer server.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: server.URL,
+ TicketExchangeMethod: testCase.method,
+ TicketExchangePayloadMode: testCase.payloadMode,
+ TicketExchangeTicketField: "st",
+ TicketExchangeServiceField: "svc",
+ TicketExchangeExtraParams: `{"source":"{provider_slug}:{state}","raw_callback":"{callback_url}"}`,
+ TicketExchangeHeaders: `{"X-State":"{provider_slug}:{state}","X-Ticket":"{ticket}"}`,
+ })
+
+ body, err := provider.performTicketAcquireRequest(
+ context.Background(),
+ "ticket-123",
+ callbackURL,
+ "state-123",
+ )
+ if err != nil {
+ t.Fatalf("expected request to succeed, got error: %v", err)
+ }
+ if strings.TrimSpace(string(body)) != `{"ok":true}` {
+ t.Fatalf("unexpected response body: %s", string(body))
+ }
+ })
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputWithTicketExchange(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ expectedCallbackURL := "https://new-api.example.com/oauth/acme-sso"
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-ticket-1",
+ "preferred_username": "alice",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ t.Fatalf("expected POST exchange method, got %s", r.Method)
+ }
+ if got := r.Header.Get("X-State"); got != "state-123" {
+ t.Fatalf("expected X-State header to be populated, got %q", got)
+ }
+ if err := r.ParseForm(); err != nil {
+ t.Fatalf("failed to parse form payload: %v", err)
+ }
+ if got := r.Form.Get("st"); got != "ticket-123" {
+ t.Fatalf("expected exchanged ticket field st=ticket-123, got %q", got)
+ }
+ if got := r.Form.Get("service"); got != expectedCallbackURL {
+ t.Fatalf("expected service field %q, got %q", expectedCallbackURL, got)
+ }
+ if got := r.Form.Get("source"); got != "acme-sso:state-123" {
+ t.Fatalf("expected placeholder expansion result, got %q", got)
+ }
+ payload, err := common.Marshal(map[string]any{
+ "data": map[string]any{
+ "token": token,
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal exchange response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer exchangeServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ UsernameField: "preferred_username",
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodPOST,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeForm,
+ TicketExchangeTicketField: "st",
+ TicketExchangeServiceField: "service",
+ TicketExchangeExtraParams: `{"source":"{provider_slug}:{state}"}`,
+ TicketExchangeHeaders: `{"X-State":"{state}"}`,
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ticket-123",
+ expectedCallbackURL,
+ "state-123",
+ )
+ if err != nil {
+ t.Fatalf("expected ticket exchange identity resolution to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "ext-ticket-1" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+ if identity.User.Username != "alice" {
+ t.Fatalf("unexpected username: %s", identity.User.Username)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputWithTicketValidateXML(t *testing.T) {
+ expectedCallbackURL := "https://new-api.example.com/oauth/acme-sso"
+ validationServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if got := r.URL.Query().Get("ticket"); got != "ST-XML-123" {
+ t.Fatalf("expected ticket query param ST-XML-123, got %q", got)
+ }
+ if got := r.URL.Query().Get("service"); got != expectedCallbackURL {
+ t.Fatalf("expected service query param %q, got %q", expectedCallbackURL, got)
+ }
+ w.Header().Set("Content-Type", "application/xml")
+ _, _ = w.Write([]byte(`
+
+
+ ext-cas-1
+
+ alice
+ Alice
+ alice@example.com
+ engineering
+ backup
+ platform-admin
+
+
+`))
+ }))
+ defer validationServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "CAS SSO",
+ Slug: "cas-sso",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketValidate,
+ TicketExchangeURL: validationServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodGET,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ TicketExchangeTicketField: "ticket",
+ TicketExchangeServiceField: "service",
+ UserIdField: "authenticationSuccess.user",
+ UsernameField: "authenticationSuccess.attributes.loginid",
+ DisplayNameField: "authenticationSuccess.attributes.userName",
+ EmailField: "authenticationSuccess.attributes.mailbox",
+ GroupField: "authenticationSuccess.attributes.group",
+ GroupMapping: `{"engineering":"vip"}`,
+ RoleField: "authenticationSuccess.attributes.role",
+ RoleMapping: `{"platform-admin":"admin"}`,
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ST-XML-123",
+ expectedCallbackURL,
+ "state-xml",
+ )
+ if err != nil {
+ t.Fatalf("expected ticket validation xml identity resolution to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "ext-cas-1" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+ if identity.User.Username != "alice" || identity.User.DisplayName != "Alice" || identity.User.Email != "alice@example.com" {
+ t.Fatalf("unexpected mapped user: %+v", identity.User)
+ }
+ if identity.Group != "vip" {
+ t.Fatalf("expected mapped group vip, got %s", identity.Group)
+ }
+ if identity.Role != common.RoleAdminUser {
+ t.Fatalf("expected mapped admin role, got %d", identity.Role)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputWithTicketValidateJSON(t *testing.T) {
+ validationServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "serviceResponse": map[string]any{
+ "authenticationSuccess": map[string]any{
+ "user": "ext-cas-json-1",
+ "attributes": map[string]any{
+ "loginid": "bob",
+ "userName": "Bob",
+ "mailbox": "bob@example.com",
+ },
+ },
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal validation response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer validationServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "CAS JSON",
+ Slug: "cas-json",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketValidate,
+ TicketExchangeURL: validationServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodGET,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ UserIdField: "serviceResponse.authenticationSuccess.user",
+ UsernameField: "serviceResponse.authenticationSuccess.attributes.loginid",
+ DisplayNameField: "serviceResponse.authenticationSuccess.attributes.userName",
+ EmailField: "serviceResponse.authenticationSuccess.attributes.mailbox",
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ST-JSON-123",
+ "https://new-api.example.com/oauth/cas-json",
+ "state-json",
+ )
+ if err != nil {
+ t.Fatalf("expected ticket validation json identity resolution to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "ext-cas-json-1" || identity.User.Username != "bob" {
+ t.Fatalf("unexpected mapped user: %+v", identity.User)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputWithTicketValidatePOSTJSON(t *testing.T) {
+ expectedCallbackURL := "https://new-api.example.com/oauth/cas-json-post"
+ validationServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ t.Fatalf("expected POST method, got %s", r.Method)
+ }
+ if got := r.Header.Get("X-Trace"); got != "cas-json-post:state-json-post" {
+ t.Fatalf("expected X-Trace header, got %q", got)
+ }
+
+ var payload map[string]string
+ if err := common.DecodeJson(r.Body, &payload); err != nil {
+ t.Fatalf("failed to decode validation payload: %v", err)
+ }
+ if got := payload["st"]; got != "ST-JSON-POST-123" {
+ t.Fatalf("expected custom ticket field st, got %q", got)
+ }
+ if got := payload["svc"]; got != expectedCallbackURL {
+ t.Fatalf("expected custom service field svc, got %q", got)
+ }
+ if got := payload["source"]; got != "cas-json-post:state-json-post" {
+ t.Fatalf("expected source placeholder expansion, got %q", got)
+ }
+
+ responseBody, err := common.Marshal(map[string]any{
+ "authenticationSuccess": map[string]any{
+ "user": "ext-cas-json-post-1",
+ "attributes": map[string]any{
+ "loginid": "dora",
+ "userName": "Dora",
+ "mailbox": "dora@example.com",
+ },
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal validation response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(responseBody)
+ }))
+ defer validationServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "CAS JSON POST",
+ Slug: "cas-json-post",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketValidate,
+ TicketExchangeURL: validationServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodPOST,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeJSON,
+ TicketExchangeTicketField: "st",
+ TicketExchangeServiceField: "svc",
+ TicketExchangeExtraParams: `{"source":"{provider_slug}:{state}"}`,
+ TicketExchangeHeaders: `{"X-Trace":"{provider_slug}:{state}"}`,
+ UserIdField: "authenticationSuccess.user",
+ UsernameField: "authenticationSuccess.attributes.loginid",
+ DisplayNameField: "authenticationSuccess.attributes.userName",
+ EmailField: "authenticationSuccess.attributes.mailbox",
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ST-JSON-POST-123",
+ expectedCallbackURL,
+ "state-json-post",
+ )
+ if err != nil {
+ t.Fatalf("expected ticket validation post json identity resolution to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "ext-cas-json-post-1" || identity.User.Username != "dora" {
+ t.Fatalf("unexpected mapped user: %+v", identity.User)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputWithTicketValidateDirectJSON(t *testing.T) {
+ validationServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "id": "custom-validate-1",
+ "username": "carol",
+ "display_name": "Carol",
+ "email": "carol@example.com",
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal direct json validation response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer validationServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Custom Validator",
+ Slug: "custom-validator",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketValidate,
+ TicketExchangeURL: validationServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodGET,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ UserIdField: "id",
+ UsernameField: "username",
+ DisplayNameField: "display_name",
+ EmailField: "email",
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ST-DIRECT-123",
+ "https://new-api.example.com/oauth/custom-validator",
+ "state-direct",
+ )
+ if err != nil {
+ t.Fatalf("expected direct json ticket validation to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "custom-validate-1" || identity.User.Username != "carol" {
+ t.Fatalf("unexpected mapped user: %+v", identity.User)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputRejectsTicketValidationFailure(t *testing.T) {
+ validationServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "serviceResponse": map[string]any{
+ "authenticationFailure": map[string]any{
+ "code": "INVALID_TICKET",
+ "message": "ticket expired",
+ },
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal failure response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer validationServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "CAS Failure",
+ Slug: "cas-failure",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketValidate,
+ TicketExchangeURL: validationServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodGET,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ UserIdField: "authenticationSuccess.user",
+ })
+
+ _, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ST-FAIL-123",
+ "https://new-api.example.com/oauth/cas-failure",
+ "state-fail",
+ )
+ if err == nil {
+ t.Fatal("expected ticket validation failure to be rejected")
+ }
+ if !strings.Contains(err.Error(), "ticket validation failed") {
+ t.Fatalf("expected ticket validation failure error, got %v", err)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputRejectsMissingTicket(t *testing.T) {
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: "https://issuer.example.com/api/exchange",
+ })
+
+ _, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "",
+ "https://new-api.example.com/oauth/acme-sso",
+ "state-123",
+ )
+ if err == nil || !strings.Contains(err.Error(), "missing ticket") {
+ t.Fatalf("expected missing ticket error, got %v", err)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputRejectsExchangeFailure(t *testing.T) {
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ http.Error(w, "invalid ticket", http.StatusUnauthorized)
+ }))
+ defer exchangeServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ })
+
+ _, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ticket-123",
+ "https://new-api.example.com/oauth/acme-sso",
+ "state-123",
+ )
+ if err == nil || !strings.Contains(err.Error(), "ticket acquire failed") {
+ t.Fatalf("expected exchange failure error, got %v", err)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputRejectsMissingTokenFromExchangeResponse(t *testing.T) {
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "success": true,
+ "data": map[string]any{
+ "user": "alice",
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal exchange response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer exchangeServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ TicketExchangeTokenField: "data.token",
+ })
+
+ _, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ticket-123",
+ "https://new-api.example.com/oauth/acme-sso",
+ "state-123",
+ )
+ if err == nil || !strings.Contains(err.Error(), "missing jwt token") {
+ t.Fatalf("expected missing jwt token error, got %v", err)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputUsesFallbackTokenField(t *testing.T) {
+ privateKey := mustGenerateRSAPrivateKey(t)
+ token := mustSignJWT(t, privateKey, "", jwt.MapClaims{
+ "iss": "https://issuer.example.com",
+ "aud": "new-api",
+ "sub": "ext-ticket-fallback",
+ "exp": time.Now().Add(time.Hour).Unix(),
+ })
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "data": map[string]any{
+ "access_token": token,
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal exchange response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer exchangeServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Acme SSO",
+ Slug: "acme-sso",
+ Enabled: true,
+ Issuer: "https://issuer.example.com",
+ Audience: "new-api",
+ PublicKey: mustEncodeRSAPublicKeyPEM(t, &privateKey.PublicKey),
+ UserIdField: "sub",
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ticket-123",
+ "https://new-api.example.com/oauth/acme-sso",
+ "state-123",
+ )
+ if err != nil {
+ t.Fatalf("expected fallback token extraction to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "ext-ticket-fallback" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+}
+
+func TestJWTDirectResolveIdentityFromInputWithTicketExchangeAndUserInfoMode(t *testing.T) {
+ exchangeServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ payload, err := common.Marshal(map[string]any{
+ "data": map[string]any{
+ "access_token": "opaque-access-token",
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal exchange response: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer exchangeServer.Close()
+
+ userInfoServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if got := r.Header.Get("x-access-token"); got != "opaque-access-token" {
+ t.Fatalf("expected exchanged token in x-access-token header, got %q", got)
+ }
+ payload, err := common.Marshal(map[string]any{
+ "info": map[string]any{
+ "userCode": "1410833903245320192",
+ "loginid": "liangmingsen",
+ "userName": "梁明森",
+ "mailbox": "liangmingsen@qdama.cn",
+ },
+ })
+ if err != nil {
+ t.Fatalf("failed to marshal userinfo payload: %v", err)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write(payload)
+ }))
+ defer userInfoServer.Close()
+
+ provider := NewJWTDirectProvider(&model.CustomOAuthProvider{
+ Name: "Qdama SSO",
+ Slug: "qdama-sso",
+ Enabled: true,
+ JWTIdentityMode: model.CustomJWTIdentityModeUserInfo,
+ UserInfoEndpoint: userInfoServer.URL,
+ JWTHeader: "x-access-token",
+ UserIdField: "info.userCode",
+ UsernameField: "info.loginid",
+ DisplayNameField: "info.userName",
+ EmailField: "info.mailbox",
+ JWTAcquireMode: model.CustomJWTAcquireModeTicketExchange,
+ TicketExchangeURL: exchangeServer.URL,
+ TicketExchangeMethod: model.CustomTicketExchangeMethodGET,
+ TicketExchangePayloadMode: model.CustomTicketExchangePayloadModeQuery,
+ TicketExchangeTicketField: "ticket",
+ TicketExchangeTokenField: "data.access_token",
+ TicketExchangeServiceField: "service",
+ })
+
+ identity, err := provider.ResolveIdentityFromInput(
+ context.Background(),
+ "",
+ "ST-123",
+ "https://new-api.example.com/oauth/qdama-sso?state=abc",
+ "abc",
+ )
+ if err != nil {
+ t.Fatalf("expected ticket exchange + userinfo mode to succeed, got error: %v", err)
+ }
+ if identity.User.ProviderUserID != "1410833903245320192" {
+ t.Fatalf("unexpected provider user id: %s", identity.User.ProviderUserID)
+ }
+}
+
+func mustGenerateRSAPrivateKey(t *testing.T) *rsa.PrivateKey {
+ t.Helper()
+ privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+ if err != nil {
+ t.Fatalf("failed to generate rsa key: %v", err)
+ }
+ return privateKey
+}
+
+func mustEncodeRSAPublicKeyPEM(t *testing.T, publicKey *rsa.PublicKey) string {
+ t.Helper()
+ publicKeyDER, err := x509.MarshalPKIXPublicKey(publicKey)
+ if err != nil {
+ t.Fatalf("failed to marshal rsa public key: %v", err)
+ }
+ return string(pem.EncodeToMemory(&pem.Block{
+ Type: "PUBLIC KEY",
+ Bytes: publicKeyDER,
+ }))
+}
+
+func mustSignJWT(t *testing.T, privateKey *rsa.PrivateKey, kid string, claims jwt.MapClaims) string {
+ t.Helper()
+ token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims)
+ if kid != "" {
+ token.Header["kid"] = kid
+ }
+ tokenString, err := token.SignedString(privateKey)
+ if err != nil {
+ t.Fatalf("failed to sign jwt: %v", err)
+ }
+ return tokenString
+}
+
+func bigEndianExponent(exponent int) []byte {
+ if exponent == 0 {
+ return []byte{0}
+ }
+ bytes := make([]byte, 0, 4)
+ for exponent > 0 {
+ bytes = append([]byte{byte(exponent & 0xff)}, bytes...)
+ exponent >>= 8
+ }
+ return bytes
+}
diff --git a/oauth/provider.go b/oauth/provider.go
index 785ed25d251b..fd6773c292ff 100644
--- a/oauth/provider.go
+++ b/oauth/provider.go
@@ -34,3 +34,10 @@ type Provider interface {
// GetProviderPrefix returns the prefix for auto-generated usernames (e.g., "github_")
GetProviderPrefix() string
}
+
+// CustomBindingProvider represents provider kinds that store external identities
+// in user_oauth_bindings rather than dedicated columns on the users table.
+type CustomBindingProvider interface {
+ Provider
+ GetProviderId() int
+}
diff --git a/oauth/registry.go b/oauth/registry.go
index 91d19636459c..5b4df6407132 100644
--- a/oauth/registry.go
+++ b/oauth/registry.go
@@ -104,13 +104,19 @@ func LoadCustomProviders() error {
}
// Register each custom provider
+ loadedCount := 0
for _, config := range customProviders {
+ if !config.IsOAuthCode() {
+ common.SysLog("Skip non-oauth_code custom provider in OAuth registry: " + config.Name + " (" + config.Slug + ")")
+ continue
+ }
provider := NewGenericOAuthProvider(config)
RegisterCustom(config.Slug, provider)
+ loadedCount++
common.SysLog("Loaded custom OAuth provider: " + config.Name + " (" + config.Slug + ")")
}
- common.SysLog(fmt.Sprintf("Loaded %d custom OAuth providers", len(customProviders)))
+ common.SysLog(fmt.Sprintf("Loaded %d custom OAuth providers", loadedCount))
return nil
}
@@ -121,9 +127,15 @@ func ReloadCustomProviders() error {
// RegisterOrUpdateCustomProvider registers or updates a single custom provider
func RegisterOrUpdateCustomProvider(config *model.CustomOAuthProvider) {
- provider := NewGenericOAuthProvider(config)
mu.Lock()
defer mu.Unlock()
+ if !config.IsOAuthCode() {
+ common.SysLog("Removing custom OAuth provider from registry because kind is not oauth_code: " + config.Slug)
+ delete(providers, config.Slug)
+ delete(customProviderSlugs, config.Slug)
+ return
+ }
+ provider := NewGenericOAuthProvider(config)
providers[config.Slug] = provider
customProviderSlugs[config.Slug] = true
}
diff --git a/oauth/registry_test.go b/oauth/registry_test.go
new file mode 100644
index 000000000000..c8767200fcd8
--- /dev/null
+++ b/oauth/registry_test.go
@@ -0,0 +1,70 @@
+package oauth
+
+import (
+ "testing"
+
+ "github.com/QuantumNous/new-api/model"
+)
+
+func TestRegisterOrUpdateCustomProviderSkipsJWTDirect(t *testing.T) {
+ slug := "jwt-direct-test-provider"
+ UnregisterCustomProvider(slug)
+ t.Cleanup(func() {
+ UnregisterCustomProvider(slug)
+ })
+
+ RegisterOrUpdateCustomProvider(&model.CustomOAuthProvider{
+ Name: "OAuth Code Test",
+ Slug: slug,
+ Kind: model.CustomOAuthProviderKindOAuthCode,
+ ClientId: "client-id",
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ TokenEndpoint: "https://issuer.example.com/oauth2/token",
+ UserInfoEndpoint: "https://issuer.example.com/oauth2/userinfo",
+ })
+
+ if provider := GetProvider(slug); provider == nil {
+ t.Fatalf("expected oauth_code provider %s to be registered before switching kinds", slug)
+ }
+ if !IsCustomProvider(slug) {
+ t.Fatalf("expected oauth_code provider %s to be marked as custom provider before switching kinds", slug)
+ }
+
+ RegisterOrUpdateCustomProvider(&model.CustomOAuthProvider{
+ Name: "JWT Direct Test",
+ Slug: slug,
+ Kind: model.CustomOAuthProviderKindJWTDirect,
+ })
+
+ if provider := GetProvider(slug); provider != nil {
+ t.Fatalf("expected jwt_direct provider %s to be excluded from oauth registry", slug)
+ }
+ if IsCustomProvider(slug) {
+ t.Fatalf("expected jwt_direct provider %s not to be marked as registry custom provider", slug)
+ }
+}
+
+func TestRegisterOrUpdateCustomProviderKeepsOAuthCode(t *testing.T) {
+ slug := "oauth-code-test-provider"
+ UnregisterCustomProvider(slug)
+ t.Cleanup(func() {
+ UnregisterCustomProvider(slug)
+ })
+
+ RegisterOrUpdateCustomProvider(&model.CustomOAuthProvider{
+ Name: "OAuth Code Test",
+ Slug: slug,
+ Kind: model.CustomOAuthProviderKindOAuthCode,
+ ClientId: "client-id",
+ AuthorizationEndpoint: "https://issuer.example.com/oauth2/authorize",
+ TokenEndpoint: "https://issuer.example.com/oauth2/token",
+ UserInfoEndpoint: "https://issuer.example.com/oauth2/userinfo",
+ })
+
+ if provider := GetProvider(slug); provider == nil {
+ t.Fatalf("expected oauth_code provider %s to remain in oauth registry", slug)
+ }
+ if !IsCustomProvider(slug) {
+ t.Fatalf("expected oauth_code provider %s to be marked as registry custom provider", slug)
+ }
+}
diff --git a/router/api-router.go b/router/api-router.go
index bff158a819ce..d5ba64aa5365 100644
--- a/router/api-router.go
+++ b/router/api-router.go
@@ -44,6 +44,7 @@ func SetApiRouter(router *gin.Engine) {
apiRouter.GET("/oauth/telegram/bind", middleware.CriticalRateLimit(), controller.TelegramBind)
// Standard OAuth providers (GitHub, Discord, OIDC, LinuxDO) - unified route
apiRouter.GET("/oauth/:provider", middleware.CriticalRateLimit(), controller.HandleOAuth)
+ apiRouter.POST("/auth/external/:provider/jwt/login", middleware.CriticalRateLimit(), controller.HandleCustomOAuthJWTLogin)
apiRouter.GET("/ratio_config", middleware.CriticalRateLimit(), controller.GetRatioConfig)
apiRouter.POST("/stripe/webhook", controller.StripeWebhook)
diff --git a/web/src/components/auth/LoginForm.jsx b/web/src/components/auth/LoginForm.jsx
index 7e8c0ce017f1..1f4797432b39 100644
--- a/web/src/components/auth/LoginForm.jsx
+++ b/web/src/components/auth/LoginForm.jsx
@@ -131,8 +131,13 @@ const LoginForm = () => {
return {};
}
}, [statusState?.status]);
+ const customOAuthProviders = Array.isArray(status.custom_oauth_providers)
+ ? status.custom_oauth_providers
+ : [];
const hasCustomOAuthProviders =
- (status.custom_oauth_providers || []).length > 0;
+ customOAuthProviders.some(
+ (provider) => provider.browser_login_supported !== false,
+ );
const hasOAuthLoginOptions = Boolean(
status.github_oauth ||
status.discord_oauth ||
@@ -603,22 +608,24 @@ const LoginForm = () => {
)}
- {status.custom_oauth_providers &&
- status.custom_oauth_providers.map((provider) => (
-
- ))}
+ {customOAuthProviders.length > 0 &&
+ customOAuthProviders
+ .filter((provider) => provider.browser_login_supported !== false)
+ .map((provider) => (
+
+ ))}
{status.telegram_oauth && (
diff --git a/web/src/components/auth/OAuth2Callback.jsx b/web/src/components/auth/OAuth2Callback.jsx
index 55a85c6b14a5..68b793f0de50 100644
--- a/web/src/components/auth/OAuth2Callback.jsx
+++ b/web/src/components/auth/OAuth2Callback.jsx
@@ -35,17 +35,67 @@ const OAuth2Callback = (props) => {
const [searchParams] = useSearchParams();
const [, userDispatch] = useContext(UserContext);
const navigate = useNavigate();
-
+
// 防止 React 18 Strict Mode 下重复执行
const hasExecuted = useRef(false);
// 最大重试次数
const MAX_RETRIES = 3;
+ const getHashParams = () => {
+ const hash = window.location.hash.startsWith('#')
+ ? window.location.hash.slice(1)
+ : window.location.hash;
+ return new URLSearchParams(hash);
+ };
+
+ const getStoredCustomProvider = () => {
+ try {
+ const statusStr = localStorage.getItem('status');
+ if (!statusStr) return null;
+ const status = JSON.parse(statusStr);
+ const customProviders = Array.isArray(status.custom_oauth_providers)
+ ? status.custom_oauth_providers
+ : [];
+ return customProviders.find(
+ (provider) => provider.slug === props.type,
+ );
+ } catch (error) {
+ return null;
+ }
+ };
+
+ const pickFirstParamValue = (query, hash, keys) => {
+ for (const key of keys) {
+ const queryValue = query.get(key);
+ if (queryValue) return queryValue;
+ const hashValue = hash.get(key);
+ if (hashValue) return hashValue;
+ }
+ return '';
+ };
+
+ const handleCallbackSuccess = (data) => {
+ if (data?.action === 'bind') {
+ showSuccess(t('绑定成功!'));
+ navigate('/console/personal');
+ return;
+ }
+
+ userDispatch({ type: 'login', payload: data });
+ setUserData(data);
+ updateAPI();
+ showSuccess(t('登录成功!'));
+ navigate('/console/token');
+ };
+
const sendCode = async (code, state, retry = 0) => {
try {
const { data: resData } = await API.get(
`/api/oauth/${props.type}?code=${code}&state=${state}`,
+ {
+ skipErrorHandler: true,
+ },
);
const { success, message, data } = resData;
@@ -56,17 +106,7 @@ const OAuth2Callback = (props) => {
return;
}
- if (data?.action === 'bind') {
- showSuccess(t('绑定成功!'));
- navigate('/console/personal');
- } else {
- userDispatch({ type: 'login', payload: data });
- localStorage.setItem('user', JSON.stringify(data));
- setUserData(data);
- updateAPI();
- showSuccess(t('登录成功!'));
- navigate('/console/token');
- }
+ handleCallbackSuccess(data);
} catch (error) {
// 网络错误等可重试
if (retry < MAX_RETRIES) {
@@ -81,6 +121,35 @@ const OAuth2Callback = (props) => {
}
};
+ const submitJWTLogin = async (payload, retry = 0) => {
+ try {
+ const { data: resData } = await API.post(
+ `/api/auth/external/${props.type}/jwt/login`,
+ payload,
+ {
+ skipErrorHandler: true,
+ },
+ );
+
+ const { success, message, data } = resData;
+ if (!success) {
+ showError(message || t('授权失败'));
+ navigate('/console/personal');
+ return;
+ }
+
+ handleCallbackSuccess(data);
+ } catch (error) {
+ if (retry < MAX_RETRIES) {
+ await new Promise((resolve) => setTimeout(resolve, (retry + 1) * 2000));
+ return submitJWTLogin(payload, retry + 1);
+ }
+
+ showError(error.message || t('授权失败'));
+ navigate('/console/personal');
+ }
+ };
+
useEffect(() => {
// 防止 React 18 Strict Mode 下重复执行
if (hasExecuted.current) {
@@ -88,6 +157,54 @@ const OAuth2Callback = (props) => {
}
hasExecuted.current = true;
+ const hashParams = getHashParams();
+ const customProvider = getStoredCustomProvider();
+ const providerKind = customProvider?.kind || 'oauth_code';
+ const errorDescription =
+ pickFirstParamValue(searchParams, hashParams, ['error_description']) ||
+ pickFirstParamValue(searchParams, hashParams, ['error']);
+
+ if (errorDescription) {
+ showError(errorDescription);
+ navigate('/console/personal');
+ return;
+ }
+
+ if (providerKind === 'jwt_direct') {
+ const jwtAcquireMode = customProvider?.jwt_acquire_mode || 'direct_token';
+ const state = pickFirstParamValue(searchParams, hashParams, ['state']);
+
+ if (['ticket_exchange', 'ticket_validate'].includes(jwtAcquireMode)) {
+ const ticket = pickFirstParamValue(searchParams, hashParams, [
+ 'ticket',
+ 'st',
+ ]);
+ if (!ticket) {
+ showError(t('未获取到登录票据'));
+ navigate('/console/personal');
+ return;
+ }
+
+ submitJWTLogin({ state, ticket });
+ return;
+ }
+
+ const jwtToken = pickFirstParamValue(searchParams, hashParams, [
+ 'id_token',
+ 'token',
+ 'jwt',
+ 'access_token',
+ ]);
+ if (!jwtToken) {
+ showError(t('未获取到 JWT 令牌'));
+ navigate('/console/personal');
+ return;
+ }
+
+ submitJWTLogin({ state, id_token: jwtToken });
+ return;
+ }
+
const code = searchParams.get('code');
const state = searchParams.get('state');
diff --git a/web/src/components/auth/RegisterForm.jsx b/web/src/components/auth/RegisterForm.jsx
index 0a755b194431..462ad713f0e9 100644
--- a/web/src/components/auth/RegisterForm.jsx
+++ b/web/src/components/auth/RegisterForm.jsx
@@ -129,8 +129,13 @@ const RegisterForm = () => {
return {};
}
}, [statusState?.status]);
+ const customOAuthProviders = Array.isArray(status.custom_oauth_providers)
+ ? status.custom_oauth_providers
+ : [];
const hasCustomOAuthProviders =
- (status.custom_oauth_providers || []).length > 0;
+ customOAuthProviders.some(
+ (provider) => provider.browser_login_supported !== false,
+ );
const hasOAuthRegisterOptions = Boolean(
status.github_oauth ||
status.discord_oauth ||
@@ -494,22 +499,24 @@ const RegisterForm = () => {
)}
- {status.custom_oauth_providers &&
- status.custom_oauth_providers.map((provider) => (
-
- ))}
+ {customOAuthProviders.length > 0 &&
+ customOAuthProviders
+ .filter((provider) => provider.browser_login_supported !== false)
+ .map((provider) => (
+
+ ))}
{status.telegram_oauth && (
diff --git a/web/src/components/settings/CustomOAuthSetting.jsx b/web/src/components/settings/CustomOAuthSetting.jsx
index 0912160bee5d..7ff1d53bcfff 100644
--- a/web/src/components/settings/CustomOAuthSetting.jsx
+++ b/web/src/components/settings/CustomOAuthSetting.jsx
@@ -40,7 +40,12 @@ import {
IconDelete,
IconRefresh,
} from '@douyinfe/semi-icons';
-import { API, showError, showSuccess, getOAuthProviderIcon } from '../../helpers';
+import {
+ API,
+ showError,
+ showSuccess,
+ getOAuthProviderIcon,
+} from '../../helpers';
import { useTranslation } from 'react-i18next';
const { Text } = Typography;
@@ -156,15 +161,65 @@ const PRESET_RESET_VALUES = {
access_denied_message: '',
};
+const CUSTOM_OAUTH_KIND_OPTIONS = [
+ { value: 'oauth_code', label: 'OAuth 2.0 / OIDC 授权码模式' },
+ { value: 'jwt_direct', label: 'JWT 直连登录' },
+];
+
+const JWT_SOURCE_OPTIONS = [
+ { value: 'query', label: '查询参数' },
+ { value: 'fragment', label: 'URL 片段' },
+ { value: 'body', label: '请求体(仅 API)' },
+];
+
+const JWT_ACQUIRE_MODE_OPTIONS = [
+ { value: 'direct_token', label: '直接回调 JWT' },
+ { value: 'ticket_exchange', label: '票据换取 JWT' },
+ { value: 'ticket_validate', label: '票据校验(CAS serviceValidate)' },
+];
+
+const JWT_IDENTITY_MODE_OPTIONS = [
+ { value: 'claims', label: '本地验签并解析 JWT Claims' },
+ { value: 'userinfo', label: '通过用户信息端点解析身份' },
+];
+
+const TICKET_EXCHANGE_METHOD_OPTIONS = [
+ { value: 'GET', label: 'GET' },
+ { value: 'POST', label: 'POST' },
+];
+
+const TICKET_EXCHANGE_PAYLOAD_MODE_OPTIONS = [
+ { value: 'query', label: '查询字符串' },
+ { value: 'form', label: '表单 URL 编码' },
+ { value: 'json', label: 'JSON 请求体' },
+ { value: 'multipart', label: 'Multipart 表单' },
+];
+
+const JWT_MAPPING_MODE_OPTIONS = [
+ { value: 'explicit_only', label: '仅显式映射' },
+ { value: 'mapping_first', label: '映射优先,其次透传' },
+];
+
const DISCOVERY_FIELD_LABELS = {
- authorization_endpoint: 'Authorization Endpoint',
- token_endpoint: 'Token Endpoint',
- user_info_endpoint: 'User Info Endpoint',
- scopes: 'Scopes',
- user_id_field: 'User ID Field',
- username_field: 'Username Field',
- display_name_field: 'Display Name Field',
- email_field: 'Email Field',
+ authorization_endpoint: '授权端点',
+ token_endpoint: '令牌端点',
+ user_info_endpoint: '用户信息端点',
+ scopes: '作用域',
+ user_id_field: '用户 ID 字段',
+ username_field: '用户名字段',
+ display_name_field: '显示名称字段',
+ email_field: '邮箱字段',
+};
+
+const REQUIRED_FIELD_LABELS = {
+ name: '显示名称',
+ slug: 'Slug',
+ client_id: '客户端 ID',
+ client_secret: '客户端密钥',
+ authorization_endpoint: '授权端点',
+ token_endpoint: '令牌端点',
+ user_info_endpoint: '用户信息端点',
+ issuer: '发行者',
};
const ACCESS_POLICY_TEMPLATES = {
@@ -185,8 +240,17 @@ const ACCESS_POLICY_TEMPLATES = {
};
const ACCESS_DENIED_TEMPLATES = {
- level_hint: '需要等级 {{required}},你当前等级 {{current}}(字段:{{field}})',
- org_hint: '仅限指定组织或角色访问。组织={{current.org}},角色={{current.roles}}',
+ level_hint:
+ '需要等级 {{required}},你当前等级 {{current}}(字段:{{field}})',
+ org_hint:
+ '仅限指定组织或角色访问。组织={{current.org}},角色={{current.roles}}',
+};
+
+const TICKET_VALIDATE_SUGGESTED_FIELDS = {
+ user_id_field: 'authenticationSuccess.user',
+ username_field: 'authenticationSuccess.attributes.loginid',
+ display_name_field: 'authenticationSuccess.attributes.userName',
+ email_field: 'authenticationSuccess.attributes.mailbox',
};
const CustomOAuthSetting = ({ serverAddress }) => {
@@ -201,7 +265,57 @@ const CustomOAuthSetting = ({ serverAddress }) => {
const [discoveryLoading, setDiscoveryLoading] = useState(false);
const [discoveryInfo, setDiscoveryInfo] = useState(null);
const [advancedActiveKeys, setAdvancedActiveKeys] = useState([]);
+ const [clientSecretDirty, setClientSecretDirty] = useState(false);
+ const [clearClientSecret, setClearClientSecret] = useState(false);
const formApiRef = React.useRef(null);
+ const customOAuthKindOptions = CUSTOM_OAUTH_KIND_OPTIONS.map((option) => ({
+ ...option,
+ label: t(option.label),
+ }));
+ const jwtSourceOptions = JWT_SOURCE_OPTIONS.map((option) => ({
+ ...option,
+ label: t(option.label),
+ }));
+ const jwtAcquireModeOptions = JWT_ACQUIRE_MODE_OPTIONS.map((option) => ({
+ ...option,
+ label: t(option.label),
+ }));
+ const jwtIdentityModeOptions = JWT_IDENTITY_MODE_OPTIONS.map((option) => ({
+ ...option,
+ label: t(option.label),
+ }));
+ const ticketExchangeMethodOptions = TICKET_EXCHANGE_METHOD_OPTIONS.map(
+ (option) => ({
+ ...option,
+ label: t(option.label),
+ }),
+ );
+ const ticketExchangePayloadModeOptions =
+ TICKET_EXCHANGE_PAYLOAD_MODE_OPTIONS.map((option) => ({
+ ...option,
+ label: t(option.label),
+ }));
+ const jwtMappingModeOptions = JWT_MAPPING_MODE_OPTIONS.map((option) => ({
+ ...option,
+ label: t(option.label),
+ }));
+ const discoveryFieldLabels = Object.fromEntries(
+ Object.entries(DISCOVERY_FIELD_LABELS).map(([field, label]) => [
+ field,
+ t(label),
+ ]),
+ );
+ const currentProviderKind = formValues.kind || 'oauth_code';
+ const isJWTDirect = currentProviderKind === 'jwt_direct';
+ const currentJWTIdentityMode = formValues.jwt_identity_mode || 'claims';
+ const currentJWTAcquireMode = formValues.jwt_acquire_mode || 'direct_token';
+ const isJWTTicketExchange =
+ isJWTDirect && currentJWTAcquireMode === 'ticket_exchange';
+ const isJWTTicketValidateMode =
+ isJWTDirect && currentJWTAcquireMode === 'ticket_validate';
+ const isJWTTicketBasedMode = isJWTTicketExchange || isJWTTicketValidateMode;
+ const isJWTUserInfoMode =
+ isJWTDirect && currentJWTIdentityMode === 'userinfo';
const mergeFormValues = (newValues) => {
setFormValues((prev) => ({ ...prev, ...newValues }));
@@ -216,10 +330,27 @@ const CustomOAuthSetting = ({ serverAddress }) => {
return values && typeof values === 'object' ? values : formValues;
};
+ const applyTicketValidateSuggestions = (values = {}) => {
+ const nextValues = {};
+ Object.entries(TICKET_VALIDATE_SUGGESTED_FIELDS).forEach(
+ ([field, suggestedValue]) => {
+ const currentValue = (values[field] ?? formValues[field] ?? '').trim();
+ if (
+ !currentValue ||
+ ['sub', 'preferred_username', 'name', 'email'].includes(currentValue)
+ ) {
+ nextValues[field] = suggestedValue;
+ }
+ },
+ );
+ return nextValues;
+ };
+
const normalizeBaseUrl = (url) => (url || '').trim().replace(/\/+$/, '');
const inferBaseUrlFromProvider = (provider) => {
- const endpoint = provider?.authorization_endpoint || provider?.token_endpoint;
+ const endpoint =
+ provider?.authorization_endpoint || provider?.token_endpoint;
if (!endpoint) return '';
try {
const url = new URL(endpoint);
@@ -237,6 +368,8 @@ const CustomOAuthSetting = ({ serverAddress }) => {
setModalVisible(false);
resetDiscoveryState();
setAdvancedActiveKeys([]);
+ setClientSecretDirty(false);
+ setClearClientSecret(false);
};
const fetchProviders = async () => {
@@ -261,13 +394,32 @@ const CustomOAuthSetting = ({ serverAddress }) => {
const handleAdd = () => {
setEditingProvider(null);
setFormValues({
+ kind: 'oauth_code',
enabled: false,
icon: '',
scopes: 'openid profile email',
+ jwt_source: 'query',
+ jwt_identity_mode: 'claims',
+ jwt_acquire_mode: 'direct_token',
+ jwt_header: 'Authorization',
+ authorization_service_field: 'service',
+ ticket_exchange_method: 'GET',
+ ticket_exchange_payload_mode: 'query',
+ ticket_exchange_ticket_field: 'ticket',
+ ticket_exchange_token_field: '',
+ ticket_exchange_service_field: '',
+ ticket_exchange_extra_params: '',
+ ticket_exchange_headers: '',
user_id_field: 'sub',
username_field: 'preferred_username',
display_name_field: 'name',
email_field: 'email',
+ auto_register: false,
+ auto_merge_by_email: false,
+ sync_group_on_login: false,
+ sync_role_on_login: false,
+ group_mapping_mode: 'explicit_only',
+ role_mapping_mode: 'explicit_only',
auth_style: 0,
access_policy: '',
access_denied_message: '',
@@ -276,16 +428,43 @@ const CustomOAuthSetting = ({ serverAddress }) => {
setBaseUrl('');
resetDiscoveryState();
setAdvancedActiveKeys([]);
+ setClientSecretDirty(false);
+ setClearClientSecret(false);
setModalVisible(true);
};
const handleEdit = (provider) => {
setEditingProvider(provider);
- setFormValues({ ...provider });
+ setFormValues({
+ ...provider,
+ kind: provider.kind || 'oauth_code',
+ jwt_source: provider.jwt_source || 'query',
+ jwt_header: provider.jwt_header || 'Authorization',
+ jwt_identity_mode: provider.jwt_identity_mode || 'claims',
+ jwt_acquire_mode: provider.jwt_acquire_mode || 'direct_token',
+ authorization_service_field:
+ provider.authorization_service_field || 'service',
+ ticket_exchange_method: provider.ticket_exchange_method || 'GET',
+ ticket_exchange_payload_mode:
+ provider.ticket_exchange_payload_mode || 'query',
+ ticket_exchange_ticket_field:
+ provider.ticket_exchange_ticket_field || 'ticket',
+ ticket_exchange_token_field: provider.ticket_exchange_token_field || '',
+ ticket_exchange_service_field:
+ provider.ticket_exchange_service_field || '',
+ ticket_exchange_extra_params: provider.ticket_exchange_extra_params || '',
+ ticket_exchange_headers: provider.ticket_exchange_headers || '',
+ sync_group_on_login: !!provider.sync_group_on_login,
+ sync_role_on_login: !!provider.sync_role_on_login,
+ group_mapping_mode: provider.group_mapping_mode || 'explicit_only',
+ role_mapping_mode: provider.role_mapping_mode || 'explicit_only',
+ });
setSelectedPreset(OAUTH_PRESETS[provider.slug] ? provider.slug : '');
setBaseUrl(inferBaseUrlFromProvider(provider));
resetDiscoveryState();
setAdvancedActiveKeys([]);
+ setClientSecretDirty(false);
+ setClearClientSecret(false);
setModalVisible(true);
};
@@ -305,38 +484,93 @@ const CustomOAuthSetting = ({ serverAddress }) => {
const handleSubmit = async () => {
const currentValues = getLatestFormValues();
+ const providerKind = currentValues.kind || 'oauth_code';
+
+ const requiredFields = ['name', 'slug'];
+ if (providerKind === 'oauth_code') {
+ requiredFields.push(
+ 'client_id',
+ 'authorization_endpoint',
+ 'token_endpoint',
+ 'user_info_endpoint',
+ );
- // Validate required fields
- const requiredFields = [
- 'name',
- 'slug',
- 'client_id',
- 'authorization_endpoint',
- 'token_endpoint',
- 'user_info_endpoint',
- ];
-
- if (!editingProvider) {
- requiredFields.push('client_secret');
+ if (!editingProvider) {
+ requiredFields.push('client_secret');
+ }
+ } else {
+ const acquireMode = currentValues.jwt_acquire_mode || 'direct_token';
+ const identityMode = currentValues.jwt_identity_mode || 'claims';
+ if (acquireMode === 'ticket_validate' && identityMode !== 'claims') {
+ showError(
+ t('Ticket Validation 模式仅支持 claims 身份解析方式'),
+ );
+ return;
+ }
+ if (identityMode === 'userinfo') {
+ requiredFields.push('user_info_endpoint');
+ } else if (acquireMode !== 'ticket_validate') {
+ requiredFields.push('issuer');
+ if (!currentValues.jwks_url && !currentValues.public_key) {
+ showError(t('JWT 直连至少需要配置 JWKS URL 或公钥'));
+ return;
+ }
+ }
+ if (
+ ['ticket_exchange', 'ticket_validate'].includes(acquireMode) &&
+ !currentValues.ticket_exchange_url
+ ) {
+ showError(t('票据处理模式必须填写 Ticket Processing URL'));
+ return;
+ }
+ if (
+ ['ticket_exchange', 'ticket_validate'].includes(acquireMode) &&
+ currentValues.ticket_exchange_url &&
+ !currentValues.ticket_exchange_url.startsWith('http://') &&
+ !currentValues.ticket_exchange_url.startsWith('https://')
+ ) {
+ showError(t('票据处理模式必须填写有效的 Ticket Processing URL'));
+ return;
+ }
}
for (const field of requiredFields) {
if (!currentValues[field]) {
- showError(t(`请填写 ${field}`));
+ const fieldLabel = REQUIRED_FIELD_LABELS[field] || field;
+ showError(
+ t('请填写 {{fieldLabel}}', {
+ fieldLabel: t(fieldLabel),
+ }),
+ );
return;
}
}
- // Validate endpoint URLs must be full URLs
- const endpointFields = ['authorization_endpoint', 'token_endpoint', 'user_info_endpoint'];
+ const endpointFields =
+ providerKind === 'oauth_code'
+ ? ['authorization_endpoint', 'token_endpoint', 'user_info_endpoint']
+ : [
+ 'authorization_endpoint',
+ ...(currentValues.jwt_identity_mode === 'userinfo'
+ ? ['user_info_endpoint']
+ : currentValues.jwt_acquire_mode === 'ticket_validate'
+ ? []
+ : ['issuer', 'jwks_url']),
+ ];
for (const field of endpointFields) {
const value = currentValues[field];
- if (value && !value.startsWith('http://') && !value.startsWith('https://')) {
+ if (
+ value &&
+ !value.startsWith('http://') &&
+ !value.startsWith('https://')
+ ) {
// Check if user selected a preset but forgot to fill issuer URL
- if (selectedPreset && !baseUrl) {
- showError(t('请先填写 Issuer URL,以自动生成完整的端点 URL'));
+ if (providerKind === 'oauth_code' && selectedPreset && !baseUrl) {
+ showError(t('请先填写发行者 URL,以自动生成完整的端点 URL'));
} else {
- showError(t('端点 URL 必须是完整地址(以 http:// 或 https:// 开头)'));
+ showError(
+ t('端点 URL 必须是完整地址(以 http:// 或 https:// 开头)'),
+ );
}
return;
}
@@ -346,12 +580,46 @@ const CustomOAuthSetting = ({ serverAddress }) => {
const payload = { ...currentValues, enabled: !!currentValues.enabled };
delete payload.preset;
delete payload.base_url;
+ if (editingProvider) {
+ if (clearClientSecret) {
+ payload.client_secret = '';
+ } else if (!clientSecretDirty || payload.client_secret === '') {
+ delete payload.client_secret;
+ }
+ }
+ if (editingProvider) {
+ const hiddenJWTSecretFields = [
+ 'ticket_exchange_extra_params',
+ 'ticket_exchange_headers',
+ ];
+ hiddenJWTSecretFields.forEach((field) => {
+ if (
+ editingProvider[field] === undefined &&
+ payload[field] === ''
+ ) {
+ delete payload[field];
+ }
+ });
+ }
+ if (providerKind !== 'jwt_direct') {
+ delete payload.jwt_identity_mode;
+ delete payload.jwt_acquire_mode;
+ delete payload.authorization_service_field;
+ delete payload.ticket_exchange_url;
+ delete payload.ticket_exchange_method;
+ delete payload.ticket_exchange_payload_mode;
+ delete payload.ticket_exchange_ticket_field;
+ delete payload.ticket_exchange_token_field;
+ delete payload.ticket_exchange_service_field;
+ delete payload.ticket_exchange_extra_params;
+ delete payload.ticket_exchange_headers;
+ }
let res;
if (editingProvider) {
res = await API.put(
`/api/custom-oauth-provider/${editingProvider.id}`,
- payload
+ payload,
);
} else {
res = await API.post('/api/custom-oauth-provider/', payload);
@@ -372,6 +640,11 @@ const CustomOAuthSetting = ({ serverAddress }) => {
}
};
+ const issuerRules =
+ isJWTUserInfoMode || isJWTTicketValidateMode
+ ? []
+ : [{ required: true, message: t('请输入发行者') }];
+
const handleFetchFromDiscovery = async () => {
const cleanBaseUrl = normalizeBaseUrl(baseUrl);
const configuredWellKnown = (formValues.well_known || '').trim();
@@ -380,7 +653,7 @@ const CustomOAuthSetting = ({ serverAddress }) => {
(cleanBaseUrl ? `${cleanBaseUrl}/.well-known/openid-configuration` : '');
if (!wellKnownUrl) {
- showError(t('请先填写 Discovery URL 或 Issuer URL'));
+ showError(t('请先填写 Discovery URL 或发行者 URL'));
return;
}
@@ -542,7 +815,17 @@ const CustomOAuthSetting = ({ serverAddress }) => {
key: 'name',
},
{
- title: 'Slug',
+ title: t('类型'),
+ dataIndex: 'kind',
+ key: 'kind',
+ render: (kind) => (
+
+ {kind === 'jwt_direct' ? t('JWT 直连') : t('OAuth 授权码')}
+
+ ),
+ },
+ {
+ title: t('标识符 (Slug)'),
dataIndex: 'slug',
key: 'slug',
render: (slug) =>
{slug},
@@ -558,7 +841,7 @@ const CustomOAuthSetting = ({ serverAddress }) => {
),
},
{
- title: t('Client ID'),
+ title: t('客户端 ID'),
dataIndex: 'client_id',
key: 'client_id',
render: (id) => {
@@ -573,7 +856,7 @@ const CustomOAuthSetting = ({ serverAddress }) => {
}
- size="small"
+ size='small'
onClick={() => handleEdit(record)}
>
{t('编辑')}
@@ -582,7 +865,7 @@ const CustomOAuthSetting = ({ serverAddress }) => {
title={t('确定要删除此 OAuth 提供商吗?')}
onConfirm={() => handleDelete(record.id)}
>
- } size="small" type="danger">
+ } size='small' type='danger'>
{t('删除')}
@@ -592,22 +875,27 @@ const CustomOAuthSetting = ({ serverAddress }) => {
];
const discoveryAutoFilledLabels = (discoveryInfo?.autoFilledFields || [])
- .map((field) => DISCOVERY_FIELD_LABELS[field] || field)
+ .map((field) => discoveryFieldLabels[field] || field)
.join(', ');
return (
{t(
- '配置自定义 OAuth 提供商,支持 GitHub Enterprise、GitLab、Gitea、Nextcloud、Keycloak、ORY 等兼容 OAuth 2.0 协议的身份提供商'
+ '配置自定义外部身份提供商,支持 OAuth Code Flow 和 JWT Direct 两种接入模式',
)}
- {t('回调 URL 格式')}: {serverAddress || t('网站地址')}/oauth/
+ {t('浏览器回调 URL')}: {serverAddress || t('网站地址')}/oauth/
{'{slug}'}
+
+ {t('说明')}:{' '}
+ {t(
+ 'JWT Direct 支持 direct_token、ticket_exchange、ticket_validate 三种获取模式,并支持 claims 或 userinfo 两类身份解析方式',
+ )}
>
}
style={{ marginBottom: 20 }}
@@ -615,24 +903,24 @@ const CustomOAuthSetting = ({ serverAddress }) => {
}
- theme="solid"
+ theme='solid'
onClick={handleAdd}
style={{ marginBottom: 16 }}
>
- {t('添加 OAuth 提供商')}
+ {t('添加身份提供商')}
{
mergeFormValues({ enabled: !!checked })}
+ onChange={(checked) =>
+ mergeFormValues({ enabled: !!checked })
+ }
+ id='components-settings-customoauthsetting-switch-1'
/>
{formValues.enabled ? t('已启用') : t('已禁用')}
@@ -674,12 +965,33 @@ const CustomOAuthSetting = ({ serverAddress }) => {
getFormApi={(api) => (formApiRef.current = api)}
>
- {t('Configuration')}
+ {t('配置')}
-
- {t('先填写配置,再自动填充 OAuth 端点,能显著减少手工输入')}
+
+ {isJWTTicketExchange
+ ? t(
+ '浏览器回调页先接收 ticket,后端再向票据交换接口换取 JWT,并继续复用现有验签、映射、建号和绑定链路',
+ )
+ : isJWTTicketValidateMode
+ ? t(
+ '浏览器回调页先接收 ticket,后端再向票据校验接口取回身份声明,直接复用现有字段映射、建号和绑定链路',
+ )
+ : isJWTUserInfoMode
+ ? t(
+ 'JWT Direct 使用前端回调页接收 token,再由后端调用用户信息接口验证 token 并提取身份',
+ )
+ : isJWTDirect
+ ? t(
+ 'JWT Direct 使用前端回调页接收 JWT,再由后端完成验签、建号、绑定与登录',
+ )
+ : t(
+ '先填写配置,再自动填充 OAuth 端点,能显著减少手工输入',
+ )}
- {discoveryInfo && (
+ {!isJWTDirect && discoveryInfo && (
{
{discoveryAutoFilledLabels ? (
- {t('自动填充字段')}:
- {' '}
- {discoveryAutoFilledLabels}
+ {t('自动填充字段')}: {discoveryAutoFilledLabels}
) : null}
{discoveryInfo.scopesSupported?.length ? (
- {t('Discovery scopes')}:
- {' '}
+ {t('Discovery 建议 scopes')}:{' '}
{discoveryInfo.scopesSupported.join(', ')}
) : null}
{discoveryInfo.claimsSupported?.length ? (
- {t('Discovery claims')}:
- {' '}
+ {t('Discovery 建议 claims')}:{' '}
{discoveryInfo.claimsSupported.join(', ')}
) : null}
@@ -718,62 +1026,109 @@ const CustomOAuthSetting = ({ serverAddress }) => {
({
- value: key,
- label: config.name,
- })),
- ]}
- />
-
-
-
-
-
-
- }
- onClick={handleFetchFromDiscovery}
- loading={discoveryLoading}
- block
- >
- {t('获取 Discovery 配置')}
-
-
-
-
-
-
- {
+ mergeFormValues({
+ kind: value,
+ jwt_identity_mode:
+ value === 'jwt_direct'
+ ? formValues.jwt_identity_mode || 'claims'
+ : formValues.jwt_identity_mode,
+ jwt_acquire_mode:
+ value === 'jwt_direct'
+ ? formValues.jwt_acquire_mode || 'direct_token'
+ : formValues.jwt_acquire_mode,
+ jwt_source:
+ value === 'jwt_direct'
+ ? formValues.jwt_source || 'query'
+ : formValues.jwt_source,
+ });
+ if (value === 'jwt_direct') {
+ setSelectedPreset('');
+ setBaseUrl('');
+ resetDiscoveryState();
+ }
+ }}
/>
+ {!isJWTDirect && (
+
+
+ ({
+ value: key,
+ label: config.name,
+ })),
+ ]}
+ />
+
+
+
+
+
+
+ }
+ onClick={handleFetchFromDiscovery}
+ loading={discoveryLoading}
+ block
+ >
+ {t('获取 Discovery 配置')}
+
+
+
+
+ )}
+ {!isJWTDirect && (
+
+
+
+
+
+ )}
+
{
{
{t(
@@ -828,109 +1185,561 @@ const CustomOAuthSetting = ({ serverAddress }) => {
-
-
-
+ {(!isJWTDirect || editingProvider) && (
+
+ {
+ setClientSecretDirty(true);
+ if (value) {
+ setClearClientSecret(false);
+ }
+ }}
+ rules={
+ editingProvider
+ ? []
+ : [
+ {
+ required: true,
+ message: t('请输入客户端密钥'),
+ },
+ ]
+ }
+ />
+ {editingProvider && (
+
+ {
+ setClearClientSecret(checked);
+ if (checked) {
+ setClientSecretDirty(true);
+ mergeFormValues({ client_secret: '' });
+ return;
+ }
+ setClientSecretDirty(false);
+ mergeFormValues({ client_secret: '' });
+ }}
+ />
+
+ {t('清空已保存的客户端密钥')}
+
+
+ )}
+
+ )}
- {t('OAuth 端点')}
+ {isJWTDirect ? t('JWT 入口与验签') : t('OAuth 端点')}
-
-
-
-
-
-
-
-
-
+ {!isJWTDirect && (
+
+
+
+
+
+
+
+
+ )}
+
+ {isJWTDirect && !isJWTTicketBasedMode && (
+
+
+
+ )}
+ {isJWTDirect && (
+ <>
+
+
+ option.value === 'claims',
+ )
+ : jwtIdentityModeOptions
+ }
+ onChange={(value) => {
+ if (
+ currentJWTAcquireMode === 'ticket_validate' &&
+ value !== 'claims'
+ ) {
+ showError(
+ t('Ticket Validation 模式仅支持 claims 身份解析方式'),
+ );
+ mergeFormValues({ jwt_identity_mode: 'claims' });
+ return;
+ }
+ mergeFormValues({ jwt_identity_mode: value });
+ }}
+ extraText={
+ isJWTUserInfoMode
+ ? t(
+ '当前模式下,后端会带 token 调用 User Info Endpoint,以响应 JSON 作为身份字段来源',
+ )
+ : t(
+ '当前模式下,后端会直接验证 JWT 签名并从 claims 中提取身份',
+ )
+ }
+ />
+
+
+ {
+ const nextValues = {
+ jwt_acquire_mode: value,
+ authorization_service_field:
+ formValues.authorization_service_field || 'service',
+ ticket_exchange_method:
+ formValues.ticket_exchange_method || 'GET',
+ ticket_exchange_payload_mode:
+ formValues.ticket_exchange_payload_mode || 'query',
+ ticket_exchange_ticket_field:
+ formValues.ticket_exchange_ticket_field || 'ticket',
+ };
+ if (value === 'ticket_validate') {
+ nextValues.jwt_identity_mode = 'claims';
+ Object.assign(
+ nextValues,
+ applyTicketValidateSuggestions(getLatestFormValues()),
+ );
+ }
+ mergeFormValues({
+ ...nextValues,
+ });
+ }}
+ extraText={
+ isJWTTicketExchange
+ ? t(
+ '当前模式下,浏览器回调接收 ticket,后端再调用票据交换接口换取 JWT',
+ )
+ : isJWTTicketValidateMode
+ ? t(
+ '当前模式下,浏览器回调接收 ticket,后端再调用票据校验接口,并直接从响应中提取身份字段',
+ )
+ : t(
+ '当前模式下,浏览器回调页直接接收 JWT 并提交给后端验签',
+ )
+ }
+ />
+
+ {isJWTTicketBasedMode && (
+
+
+
+ )}
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ {isJWTTicketBasedMode && (
+ <>
+
+ {t('票据处理配置')}
+
+
+ {isJWTTicketExchange
+ ? t(
+ '配置后端如何把浏览器回调得到的 ticket 换成 JWT。支持 query、form、json、multipart 四种请求方式,以及可选额外参数与请求头',
+ )
+ : t(
+ '配置后端如何把浏览器回调得到的 ticket 发送给校验接口。支持标准 CAS serviceValidate / p3/serviceValidate,也支持直接返回身份 JSON 的校验服务',
+ )}
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ {isJWTTicketExchange && (
+
+
+
+ )}
+
+
+
+
+
+
+
+
+ mergeFormValues({
+ ticket_exchange_extra_params: value,
+ })
+ }
+ label={t('额外参数 JSON(可选)')}
+ rows={4}
+ placeholder={`{
+ "appId": "new-api",
+ "channel": "web"
+}`}
+ extraText={t(
+ '仅支持 JSON 对象。值中可使用占位符:{ticket} {callback_url} {provider_slug} {state}',
+ )}
+ />
+
+
+
+ mergeFormValues({ ticket_exchange_headers: value })
+ }
+ label={t('请求头 JSON(可选)')}
+ rows={4}
+ placeholder={`{
+ "X-Provider": "{provider_slug}",
+ "X-State": "{state}"
+}`}
+ extraText={t(
+ '仅支持 JSON 对象。值中同样支持占位符:{ticket} {callback_url} {provider_slug} {state}',
+ )}
+ />
+
+
+ >
+ )}
+
+
+
+
+ mergeFormValues({ public_key: value })
+ }
+ label={t('公钥 PEM(可选)')}
+ rows={6}
+ placeholder={`-----BEGIN PUBLIC KEY-----
+MIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEAtest
+-----END PUBLIC KEY-----`}
+ extraText={t(
+ 'JWKS URL 与 Public Key 至少配置一项;都配置时优先使用 Public Key',
+ )}
+ showClear
+ />
+
+
+ >
+ )}
+
{t('字段映射')}
-
- {t('配置如何从用户信息 API 响应中提取用户数据,支持 JSONPath 语法')}
+
+ {isJWTDirect
+ ? t('配置如何从 JWT claims 中提取用户数据,支持 gjson 路径语法')
+ : t(
+ '配置如何从用户信息 API 响应中提取用户数据,支持 JSONPath 语法',
+ )}
{
@@ -948,20 +1757,134 @@ const CustomOAuthSetting = ({ serverAddress }) => {
+ {isJWTDirect && (
+ <>
+
+ {t('权限映射')}
+
+
+ {t(
+ 'group 只会命中系统现有分组;role 仅允许同步到 common 或 admin;guest 和 root 会被拒绝',
+ )}
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ mergeFormValues({ group_mapping: value })
+ }
+ label={t('分组映射 JSON(可选)')}
+ rows={4}
+ placeholder={`{
+ "engineering": "vip",
+ "support": "default"
+}`}
+ />
+
+
+
+ mergeFormValues({ role_mapping: value })
+ }
+ label={t('角色映射 JSON(可选)')}
+ rows={4}
+ placeholder={`{
+ "platform-admin": "admin",
+ "member": "user"
+}`}
+ />
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ >
+ )}
+
{
}}
>
-
-
-
-
-
+ {!isJWTDirect && (
+
+
+
+
+
+ )}
{t('准入策略')}
-
- {t('可选:基于用户信息 JSON 做组合条件准入,条件不满足时返回自定义提示')}
+
+ {t(
+ '可选:基于用户信息 JSON 做组合条件准入,条件不满足时返回自定义提示',
+ )}
mergeFormValues({ access_policy: value })}
+ onChange={(value) =>
+ mergeFormValues({ access_policy: value })
+ }
label={t('准入策略 JSON(可选)')}
rows={6}
placeholder={`{
@@ -1007,14 +1939,32 @@ const CustomOAuthSetting = ({ serverAddress }) => {
{"field": "active", "op": "eq", "value": true}
]
}`}
- extraText={t('支持逻辑 and/or 与嵌套 groups;操作符支持 eq/ne/gt/gte/lt/lte/in/not_in/contains/exists')}
+ extraText={
+ isJWTDirect
+ ? t(
+ '支持基于 JWT claims 做 and/or 组合准入;操作符支持 eq/ne/gt/gte/lt/lte/in/not_in/contains/exists',
+ )
+ : t(
+ '支持逻辑 and/or 与嵌套 groups;操作符支持 eq/ne/gt/gte/lt/lte/in/not_in/contains/exists',
+ )
+ }
showClear
/>
-
@@ -1025,17 +1975,31 @@ const CustomOAuthSetting = ({ serverAddress }) => {
mergeFormValues({ access_denied_message: value })}
+ onChange={(value) =>
+ mergeFormValues({ access_denied_message: value })
+ }
label={t('拒绝提示模板(可选)')}
- placeholder={t('例如:需要等级 {{required}},你当前等级 {{current}}')}
- extraText={t('可用变量:{{provider}} {{field}} {{op}} {{required}} {{current}} 以及 {{current.path}}')}
+ placeholder={t(
+ '例如:需要等级 {{required}},你当前等级 {{current}}',
+ )}
+ extraText={t(
+ '可用变量:{{provider}} {{field}} {{op}} {{required}} {{current}} 以及 {{current.path}}',
+ )}
showClear
/>
- applyDeniedTemplate('level_hint')}>
+ applyDeniedTemplate('level_hint')}
+ >
{t('填充模板:等级提示')}
- applyDeniedTemplate('org_hint')}>
+ applyDeniedTemplate('org_hint')}
+ >
{t('填充模板:组织提示')}
diff --git a/web/src/components/settings/personal/cards/AccountManagement.jsx b/web/src/components/settings/personal/cards/AccountManagement.jsx
index 29249caa1624..fea8d4042273 100644
--- a/web/src/components/settings/personal/cards/AccountManagement.jsx
+++ b/web/src/components/settings/personal/cards/AccountManagement.jsx
@@ -564,9 +564,12 @@ const AccountManagement = ({
type='primary'
theme='outline'
size='small'
+ disabled={provider.browser_login_supported === false}
onClick={() => handleBindCustomOAuth(provider)}
>
- {t('绑定')}
+ {provider.browser_login_supported === false
+ ? t('仅 API')
+ : t('绑定')}
)}
diff --git a/web/src/components/table/tokens/modals/EditTokenModal.jsx b/web/src/components/table/tokens/modals/EditTokenModal.jsx
index 93664580c38a..1e3a936c9b90 100644
--- a/web/src/components/table/tokens/modals/EditTokenModal.jsx
+++ b/web/src/components/table/tokens/modals/EditTokenModal.jsx
@@ -102,7 +102,7 @@ const EditTokenModal = (props) => {
const { success, message, data } = res.data;
if (success) {
const categories = getModelCategories(t);
- let localModelOptions = data.map((model) => {
+ let localModelOptions = (Array.isArray(data) ? data : []).map((model) => {
let icon = null;
for (const [key, category] of Object.entries(categories)) {
if (key !== 'all' && category.filter({ model_name: model })) {
@@ -130,10 +130,10 @@ const EditTokenModal = (props) => {
let res = await API.get(`/api/user/self/groups`);
const { success, message, data } = res.data;
if (success) {
- let localGroupOptions = Object.entries(data).map(([group, info]) => ({
- label: info.desc,
+ let localGroupOptions = Object.entries(data || {}).map(([group, info]) => ({
+ label: info?.desc || group,
value: group,
- ratio: info.ratio,
+ ratio: info?.ratio ?? 1,
}));
if (statusState?.status?.default_use_auto_group) {
if (localGroupOptions.some((group) => group.value === 'auto')) {
diff --git a/web/src/helpers/api.js b/web/src/helpers/api.js
index 88122a564cff..e06bf139f352 100644
--- a/web/src/helpers/api.js
+++ b/web/src/helpers/api.js
@@ -24,6 +24,7 @@ import {
isValidMessage,
} from './utils';
import axios from 'axios';
+import i18n from '../i18n/i18n';
import { MESSAGE_ROLES } from '../constants/playground.constants';
export let API = axios.create({
@@ -36,7 +37,6 @@ export let API = axios.create({
},
});
-
function redirectToOAuthUrl(url, options = {}) {
const { openInNewTab = false } = options;
const targetUrl = typeof url === 'string' ? url : url.toString();
@@ -49,6 +49,83 @@ function redirectToOAuthUrl(url, options = {}) {
window.location.assign(targetUrl);
}
+function getCustomProviderKind(provider) {
+ return provider?.kind || 'oauth_code';
+}
+
+function isTicketAcquireMode(mode) {
+ return mode === 'ticket_exchange' || mode === 'ticket_validate';
+}
+
+function supportsCustomProviderBrowserLogin(provider) {
+ if (provider?.browser_login_supported !== undefined) {
+ return Boolean(provider.browser_login_supported);
+ }
+ const providerKind = getCustomProviderKind(provider);
+ if (providerKind === 'jwt_direct') {
+ if (isTicketAcquireMode(provider?.jwt_acquire_mode || 'direct_token')) {
+ return Boolean(provider?.authorization_endpoint);
+ }
+ return Boolean(
+ provider?.authorization_endpoint &&
+ provider?.client_id &&
+ provider?.jwt_source !== 'body',
+ );
+ }
+ return Boolean(provider?.authorization_endpoint && provider?.client_id);
+}
+
+function ensureAbsoluteOAuthURL(url) {
+ if (typeof url !== 'string' || url.trim() === '') {
+ throw new Error(i18n.t('缺少授权端点 URL'));
+ }
+ if (!url.startsWith('http://') && !url.startsWith('https://')) {
+ throw new Error(i18n.t('授权端点必须是完整的 URL(以 http:// 或 https:// 开头)'));
+ }
+ return new URL(url);
+}
+
+function buildCustomJWTAuthorizationUrl(provider, state) {
+ const authUrl = ensureAbsoluteOAuthURL(provider.authorization_endpoint);
+ const acquireMode = provider.jwt_acquire_mode || 'direct_token';
+ const callbackUrl = new URL(
+ `/oauth/${provider.slug}`,
+ window.location.origin,
+ );
+
+ if (isTicketAcquireMode(acquireMode)) {
+ callbackUrl.searchParams.set('state', state);
+ authUrl.searchParams.set(
+ provider.authorization_service_field || 'service',
+ callbackUrl.toString(),
+ );
+ return authUrl;
+ }
+
+ const jwtSource = provider.jwt_source || 'query';
+
+ if (jwtSource === 'body') {
+ throw new Error(
+ i18n.t('当前浏览器登录暂不支持 form_post 模式,请改用 query 或 fragment'),
+ );
+ }
+ if (!provider.client_id) {
+ throw new Error(i18n.t('JWT 登录缺少 Client ID 配置'));
+ }
+
+ authUrl.searchParams.set('client_id', provider.client_id);
+ authUrl.searchParams.set('redirect_uri', callbackUrl.toString());
+ authUrl.searchParams.set('scope', provider.scopes || 'openid profile email');
+ authUrl.searchParams.set('state', state);
+ authUrl.searchParams.set('nonce', state);
+ authUrl.searchParams.set('response_type', 'id_token');
+ authUrl.searchParams.set(
+ 'response_mode',
+ jwtSource === 'fragment' ? 'fragment' : 'query',
+ );
+
+ return authUrl;
+}
function patchAPIInstance(instance) {
const originalGet = instance.get.bind(instance);
@@ -193,7 +270,8 @@ export const handleApiError = (error, response = null) => {
// 处理模型数据
export const processModelsData = (data, currentModel) => {
- const modelOptions = data.map((model) => ({
+ const normalizedModels = Array.isArray(data) ? data : [];
+ const modelOptions = normalizedModels.map((model) => ({
label: model,
value: model,
}));
@@ -211,13 +289,20 @@ export const processModelsData = (data, currentModel) => {
// 处理分组数据
export const processGroupsData = (data, userGroup) => {
- let groupOptions = Object.entries(data).map(([group, info]) => ({
- label:
- info.desc.length > 20 ? info.desc.substring(0, 20) + '...' : info.desc,
- value: group,
- ratio: info.ratio,
- fullLabel: info.desc,
- }));
+ const normalizedGroups =
+ data && typeof data === 'object' && !Array.isArray(data) ? data : {};
+ let groupOptions = Object.entries(normalizedGroups).map(([group, info]) => {
+ const description = info?.desc || group;
+ return {
+ label:
+ description.length > 20
+ ? description.substring(0, 20) + '...'
+ : description,
+ value: group,
+ ratio: info?.ratio ?? 1,
+ fullLabel: description,
+ };
+ });
if (groupOptions.length === 0) {
groupOptions = [
@@ -326,44 +411,41 @@ export async function onLinuxDOOAuthClicked(
* @param {boolean} options.shouldLogout - Whether to logout first
*/
export async function onCustomOAuthClicked(provider, options = {}) {
+ if (!supportsCustomProviderBrowserLogin(provider)) {
+ showError(i18n.t('当前身份提供商仅支持后端接口直连,暂不支持浏览器登录/绑定'));
+ return;
+ }
+
const state = await prepareOAuthState(options);
if (!state) return;
try {
- const redirect_uri = `${window.location.origin}/oauth/${provider.slug}`;
-
- // Check if authorization_endpoint is a full URL or relative path
- let authUrl;
- if (
- provider.authorization_endpoint.startsWith('http://') ||
- provider.authorization_endpoint.startsWith('https://')
- ) {
- authUrl = new URL(provider.authorization_endpoint);
- } else {
- // Relative path - this is a configuration error, show error message
- console.error(
- 'Custom OAuth authorization_endpoint must be a full URL:',
- provider.authorization_endpoint,
+ const providerKind = getCustomProviderKind(provider);
+ const authUrl =
+ providerKind === 'jwt_direct'
+ ? buildCustomJWTAuthorizationUrl(provider, state)
+ : ensureAbsoluteOAuthURL(provider.authorization_endpoint);
+
+ if (providerKind !== 'jwt_direct') {
+ authUrl.searchParams.set('client_id', provider.client_id);
+ authUrl.searchParams.set(
+ 'redirect_uri',
+ `${window.location.origin}/oauth/${provider.slug}`,
);
- showError(
- 'OAuth 配置错误:授权端点必须是完整的 URL(以 http:// 或 https:// 开头)',
+ authUrl.searchParams.set('response_type', 'code');
+ authUrl.searchParams.set(
+ 'scope',
+ provider.scopes || 'openid profile email',
);
- return;
+ authUrl.searchParams.set('state', state);
}
- authUrl.searchParams.set('client_id', provider.client_id);
- authUrl.searchParams.set('redirect_uri', redirect_uri);
- authUrl.searchParams.set('response_type', 'code');
- authUrl.searchParams.set(
- 'scope',
- provider.scopes || 'openid profile email',
- );
- authUrl.searchParams.set('state', state);
-
redirectToOAuthUrl(authUrl);
} catch (error) {
console.error('Failed to initiate custom OAuth:', error);
- showError('OAuth 登录失败:' + (error.message || '未知错误'));
+ showError(
+ i18n.t('OAuth 登录失败:') + (error.message || i18n.t('未知错误')),
+ );
}
}
diff --git a/web/src/i18n/locales/en.json b/web/src/i18n/locales/en.json
index 1e0671a8c812..43b2cdcdc052 100644
--- a/web/src/i18n/locales/en.json
+++ b/web/src/i18n/locales/en.json
@@ -3489,6 +3489,147 @@
"(当前仅支持易支付接口,默认使用上方服务器地址作为回调地址!)": "(Currently only supports Epay interface, the default callback address is the server address above!)",
",当前无生效订阅,将自动使用钱包": ", no active subscription. Wallet will be used automatically.",
",时间:": ",time:",
- ",点击更新": ", click Update"
+ ",点击更新": ", click Update",
+ "提示:端点映射仅用于模型广场展示,不会影响模型真实调用。如需配置真实调用,请前往「渠道管理」。": "Notice: Endpoint mapping is for Model Marketplace display only and does not affect real model invocation. To configure real invocation, please go to Channel Management.",
+ "购买订阅获得模型额度/次数": "Purchase a subscription to get model quota/usage",
+ "生产环境 RSA 私钥 Base64 (PKCS#8 DER)": "Production RSA private key Base64 (PKCS#8 DER)",
+ "沙盒环境 RSA 私钥 Base64 (PKCS#8 DER)": "Sandbox RSA private key Base64 (PKCS#8 DER)",
+ "生产环境 Waffo 公钥 Base64 (X.509 DER)": "Production Waffo public key Base64 (X.509 DER)",
+ "沙盒环境 Waffo 公钥 Base64 (X.509 DER)": "Sandbox Waffo public key Base64 (X.509 DER)",
+ "支付方式类型": "Pay Method Type",
+ "支付方式名称": "Pay Method Name",
+ "获取充值配置失败": "Failed to get topup configuration",
+ "获取充值配置异常": "Topup configuration error",
+ "分组相关设置": "Group Related Settings",
+ "保存分组相关设置": "Save Group Related Settings",
+ "此页面仅显示未设置价格或基础倍率的模型,设置后会自动从列表中移出": "This page only shows models without base pricing. After saving, configured models will be removed from this list automatically.",
+ "没有未设置定价的模型": "No unpriced models",
+ "当前没有未设置定价的模型": "There are currently no models without pricing",
+ "模型计费编辑器": "Model Pricing Editor",
+ "价格摘要": "Price Summary",
+ "当前提示": "Current Notes",
+ "这个界面默认按价格填写,保存时会自动换算回后端需要的倍率 JSON。": "This editor uses prices by default and converts them back into the ratio JSON required by the backend when saved.",
+ "当前未启用,需要时再打开即可。": "This field is currently disabled. Enable it when needed.",
+ "下面展示这个模型保存后会写入哪些后端字段,便于和原始 JSON 编辑框保持一致。": "The fields below show which backend values will be written after saving, so you can keep them aligned with the raw JSON editors.",
+ "补全价格已锁定": "Completion price is locked",
+ "后端固定倍率:{{ratio}}。该字段仅展示换算后的价格。": "Backend fixed ratio: {{ratio}}. This field only displays the converted price.",
+ "这些价格都是可选项,不填也可以。": "All of these prices are optional and can be left empty.",
+ "请先开启并填写音频输入价格。": "Enable and fill in the audio input price first.",
+ "输入模型名称,例如 gpt-4.1": "Enter a model name, for example gpt-4.1",
+ "当前模型同时存在按次价格和倍率配置,保存时会按当前计费方式覆盖。": "This model currently has both per-request pricing and ratio-based pricing. Saving will overwrite them according to the current billing mode.",
+ "当前模型存在未显式设置输入倍率的扩展倍率;填写输入价格后会自动换算为价格字段。": "This model has derived ratios without an explicit input ratio. Once you fill in the input price, they will be converted into price fields automatically.",
+ "按量计费下需要先填写输入价格,才能保存其它价格项。": "For per-token billing, fill in the input price before saving other price fields.",
+ "填写音频补全价格前,需要先填写音频输入价格。": "Fill in the audio input price before setting the audio completion price.",
+ "模型 {{name}} 缺少输入价格,无法计算补全/缓存/图片/音频价格对应的倍率": "Model {{name}} is missing an input price, so the ratios for completion, cache, image, and audio pricing cannot be calculated.",
+ "模型 {{name}} 缺少音频输入价格,无法计算音频补全倍率": "Model {{name}} is missing an audio input price, so the audio completion ratio cannot be calculated.",
+ "批量应用当前模型价格": "Batch Apply Current Model Pricing",
+ "请先选择一个作为模板的模型": "Please select a model to use as the template first",
+ "请先勾选需要批量设置的模型": "Please select the models you want to update in batch first",
+ "已将模型 {{name}} 的价格配置批量应用到 {{count}} 个模型": "Applied the pricing configuration of model {{name}} to {{count}} models in batch",
+ "将把当前编辑中的模型 {{name}} 的价格配置,批量应用到已勾选的 {{count}} 个模型。": "The pricing configuration of the currently edited model {{name}} will be applied to the {{count}} selected models.",
+ "适合同系列模型一起定价,例如把 gpt-5.1 的价格批量同步到 gpt-5.1-high、gpt-5.1-low 等模型。": "Useful for pricing model variants together, for example syncing the pricing of gpt-5.1 to gpt-5.1-high, gpt-5.1-low, and similar models.",
+ "已勾选": "Selected",
+ "当前编辑": "Editing",
+ "已勾选 {{count}} 个模型": "{{count}} models selected",
+ "计费方式": "Billing Mode",
+ "未设置价格": "Price not set",
+ "保存预览": "Save Preview",
+ "基础价格": "Base Pricing",
+ "扩展价格": "Additional Pricing",
+ "额外价格项": "Additional price items",
+ "补全价格": "Completion Price",
+ "缓存读取价格": "Input Cache Read Price",
+ "缓存创建价格": "Input Cache Creation Price",
+ "图片输入价格": "Image Input Price",
+ "音频输入价格": "Audio Input Price",
+ "音频输入价格:{{symbol}}{{price}} / 1M tokens": "Audio input price: {{symbol}}{{price}} / 1M tokens",
+ "音频补全价格": "Audio Completion Price",
+ "音频补全价格:{{symbol}}{{price}} / 1M tokens": "Audio completion price: {{symbol}}{{price}} / 1M tokens",
+ "适合 MJ / 任务类等按次收费模型。": "Suitable for MJ and other task-based models billed per request.",
+ "该模型补全倍率由后端固定为 {{ratio}}。补全价格不能在这里修改。": "This model's completion ratio is fixed to {{ratio}} by the backend. The completion price cannot be changed here.",
+ "Web 搜索调用 {{webSearchCallCount}} 次": "Web search called {{webSearchCallCount}} times",
+ "文件搜索调用 {{fileSearchCallCount}} 次": "File search called {{fileSearchCallCount}} times",
+ "实际结算金额:{{symbol}}{{total}}(已包含分组价格调整)": "Actual charge: {{symbol}}{{total}} (group pricing adjustment included)",
+ "图片倍率 {{imageRatio}}": "Image ratio {{imageRatio}}",
+ "音频倍率 {{audioRatio}}": "Audio ratio {{audioRatio}}",
+ "普通输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Standard input: {{tokens}} / 1M * model ratio {{modelRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "缓存输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 缓存倍率 {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Cached input: {{tokens}} / 1M * model ratio {{modelRatio}} * cache ratio {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "图片输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 图片倍率 {{imageRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Image input: {{tokens}} / 1M * model ratio {{modelRatio}} * image ratio {{imageRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "音频输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 音频倍率 {{audioRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Audio input: {{tokens}} / 1M * model ratio {{modelRatio}} * audio ratio {{audioRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 补全倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Output: {{tokens}} / 1M * model ratio {{modelRatio}} * completion ratio {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "Web 搜索:{{count}} / 1K * 单价 {{price}} * {{ratioType}} {{ratio}} = {{amount}}": "Web search: {{count}} / 1K * unit price {{price}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "文件搜索:{{count}} / 1K * 单价 {{price}} * {{ratioType}} {{ratio}} = {{amount}}": "File search: {{count}} / 1K * unit price {{price}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "图片生成:1 次 * 单价 {{price}} * {{ratioType}} {{ratio}} = {{amount}}": "Image generation: 1 call * unit price {{price}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "合计:{{total}}": "Total: {{total}}",
+ "模型倍率 {{modelRatio}},补全倍率 {{completionRatio}},音频倍率 {{audioRatio}},音频补全倍率 {{audioCompletionRatio}},{{cachePart}}{{ratioType}} {{ratio}}": "Model ratio {{modelRatio}}, completion ratio {{completionRatio}}, audio ratio {{audioRatio}}, audio completion ratio {{audioCompletionRatio}}, {{cachePart}}{{ratioType}} {{ratio}}",
+ "文字输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 补全倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Text output: {{tokens}} / 1M * model ratio {{modelRatio}} * completion ratio {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "音频输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 音频倍率 {{audioRatio}} * 音频补全倍率 {{audioCompletionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Audio output: {{tokens}} / 1M * model ratio {{modelRatio}} * audio ratio {{audioRatio}} * audio completion ratio {{audioCompletionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "合计:文字部分 {{textTotal}} + 音频部分 {{audioTotal}} = {{total}}": "Total: text {{textTotal}} + audio {{audioTotal}} = {{total}}",
+ "模型倍率 {{modelRatio}},输出倍率 {{completionRatio}},缓存倍率 {{cacheRatio}},{{ratioType}} {{ratio}}": "Model ratio {{modelRatio}}, output ratio {{completionRatio}}, cache ratio {{cacheRatio}}, {{ratioType}} {{ratio}}",
+ "缓存读取:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 缓存倍率 {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Cache read: {{tokens}} / 1M * model ratio {{modelRatio}} * cache ratio {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "缓存创建:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 缓存创建倍率 {{cacheCreationRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Cache creation: {{tokens}} / 1M * model ratio {{modelRatio}} * cache creation ratio {{cacheCreationRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "5m缓存创建:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 5m缓存创建倍率 {{cacheCreationRatio5m}} * {{ratioType}} {{ratio}} = {{amount}}": "5m cache creation: {{tokens}} / 1M * model ratio {{modelRatio}} * 5m cache creation ratio {{cacheCreationRatio5m}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "1h缓存创建:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 1h缓存创建倍率 {{cacheCreationRatio1h}} * {{ratioType}} {{ratio}} = {{amount}}": "1h cache creation: {{tokens}} / 1M * model ratio {{modelRatio}} * 1h cache creation ratio {{cacheCreationRatio1h}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 输出倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "Output: {{tokens}} / 1M * model ratio {{modelRatio}} * output ratio {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "空": "Empty",
+ "{{ratioType}} {{ratio}}x": "{{ratioType}} {{ratio}}x",
+ "模型价格:{{symbol}}{{price}}": "Model price: {{symbol}}{{price}}",
+ "模型价格 {{price}}": "Model price {{price}}",
+ "缓存读 {{price}} / 1M tokens": "Cache read {{price}} / 1M tokens",
+ "5m缓存创建 {{price}} / 1M tokens": "5m cache creation {{price}} / 1M tokens",
+ "1h缓存创建 {{price}} / 1M tokens": "1h cache creation {{price}} / 1M tokens",
+ "缓存创建 {{price}} / 1M tokens": "Cache creation {{price}} / 1M tokens",
+ "图片输入 {{price}} / 1M tokens": "Image input {{price}} / 1M tokens",
+ "输入 {{price}} / 1M tokens": "Input {{price}} / 1M tokens",
+ "缓存创建 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "Cache creation {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "5m缓存创建 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "5m cache creation {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "1h缓存创建 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "1h cache creation {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "(输入 {{nonImageInput}} tokens + 图片输入 {{imageInput}} tokens / 1M tokens * {{symbol}}{{price}}": "(Input {{nonImageInput}} tokens + Image input {{imageInput}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "图片输入价格:{{symbol}}{{total}} / 1M tokens": "Image input price: {{symbol}}{{total}} / 1M tokens",
+ "文字提示 {{input}} tokens / 1M tokens * {{symbol}}{{textInputPrice}} + 文字补全 {{completion}} tokens / 1M tokens * {{symbol}}{{textCompPrice}} + 音频提示 {{audioInput}} tokens / 1M tokens * {{symbol}}{{audioInputPrice}} + 音频补全 {{audioCompletion}} tokens / 1M tokens * {{symbol}}{{audioCompPrice}} * {{ratioType}} {{ratio}} = {{symbol}}{{total}}": "Text prompt {{input}} tokens / 1M tokens * {{symbol}}{{textInputPrice}} + Text completion {{completion}} tokens / 1M tokens * {{symbol}}{{textCompPrice}} + Audio prompt {{audioInput}} tokens / 1M tokens * {{symbol}}{{audioInputPrice}} + Audio completion {{audioCompletion}} tokens / 1M tokens * {{symbol}}{{audioCompPrice}} * {{ratioType}} {{ratio}} = {{symbol}}{{total}}",
+ "缓存读取价格:{{symbol}}{{total}} / 1M tokens": "Cache read price: {{symbol}}{{total}} / 1M tokens",
+ "补全 {{completion}} tokens * 输出倍率 {{completionRatio}}": "Completion {{completion}} tokens * Output ratio {{completionRatio}}",
+ "补全倍率 {{completionRatio}}": "Completion ratio {{completionRatio}}",
+ "输入价格:{{symbol}}{{price}} / 1M tokens": "Input Price: {{symbol}}{{price}} / 1M tokens",
+ "输出价格 {{symbol}}{{price}} / 1M tokens": "Output Price {{symbol}}{{price}} / 1M tokens",
+ "输出价格:{{symbol}}{{price}} / 1M tokens": "Output Price: {{symbol}}{{price}} / 1M tokens",
+ "输出价格:{{symbol}}{{total}} / 1M tokens": "Output Price: {{symbol}}{{total}} / 1M tokens",
+ "例如:gpt-4.1-nano,regex:^claude-.*$,regex:^sora-.*$": "Example: gpt-4.1-nano,regex:^claude-.*$,regex:^sora-.*$",
+ "支持精确匹配;使用 regex: 开头可按正则匹配。": "Supports exact matching. Use a regex: prefix for regex matching.",
+ "复制密钥": "Copy Key",
+ "复制连接信息": "Copy Connection String",
+ "检测到剪贴板中的连接信息": "Connection info detected in clipboard",
+ "自动填入": "Auto-fill",
+ "忽略": "Ignore",
+ "从剪贴板粘贴配置": "Paste Config",
+ "剪贴板中未检测到连接信息": "No connection info found in clipboard",
+ "连接信息已填入": "Connection info applied",
+ "无法读取剪贴板": "Cannot read clipboard",
+ "OAuth 2.0 / OIDC 授权码模式": "OAuth 2.0 / OIDC Authorization Code Flow",
+ "JWT 直连": "JWT Direct",
+ "OAuth 授权码": "OAuth Authorization Code",
+ "JWT 直连登录": "JWT Direct Login",
+ "添加身份提供商": "Add Identity Provider",
+ "编辑身份提供商": "Edit Identity Provider",
+ "浏览器回调 URL": "Browser Callback URL",
+ "配置自定义外部身份提供商,支持 OAuth Code Flow 和 JWT Direct 两种接入模式": "Configure custom external identity providers with OAuth Code Flow and JWT Direct modes",
+ "JWT Direct 支持 direct_token、ticket_exchange、ticket_validate 三种获取模式,并支持 claims 或 userinfo 两类身份解析方式": "JWT Direct supports direct_token, ticket_exchange, and ticket_validate acquisition modes, with claims or userinfo identity resolution",
+ "查询参数": "Query Parameters",
+ "URL 片段": "URL Fragment",
+ "请求体(仅 API)": "Request Body (API Only)",
+ "直接回调 JWT": "Direct JWT Callback",
+ "票据换取 JWT": "Exchange Ticket for JWT",
+ "票据校验(CAS serviceValidate)": "Ticket Validation (CAS serviceValidate)",
+ "本地验签并解析 JWT Claims": "Validate Locally and Parse JWT Claims",
+ "通过用户信息端点解析身份": "Resolve Identity via UserInfo Endpoint",
+ "查询字符串": "Query String",
+ "表单 URL 编码": "Form URL Encoded",
+ "JSON 请求体": "JSON Request Body",
+ "Multipart 表单": "Multipart Form",
+ "仅显式映射": "Explicit Mapping Only",
+ "映射优先,其次透传": "Mapping First, Then Passthrough",
+ "请填写 {{fieldLabel}}": "Please fill in {{fieldLabel}}",
+ "清空已保存的客户端密钥": "Clear saved client secret",
+ "票据处理模式必须填写有效的 Ticket Processing URL": "Ticket modes require a valid Ticket Processing URL"
}
}
diff --git a/web/src/i18n/locales/zh-CN.json b/web/src/i18n/locales/zh-CN.json
index 99a721ab81c9..7b7c2b498adf 100644
--- a/web/src/i18n/locales/zh-CN.json
+++ b/web/src/i18n/locales/zh-CN.json
@@ -2980,6 +2980,32 @@
"从剪贴板粘贴配置": "从剪贴板粘贴配置",
"剪贴板中未检测到连接信息": "剪贴板中未检测到连接信息",
"连接信息已填入": "连接信息已填入",
- "无法读取剪贴板": "无法读取剪贴板"
+ "无法读取剪贴板": "无法读取剪贴板",
+ "OAuth 2.0 / OIDC 授权码模式": "OAuth 2.0 / OIDC 授权码模式",
+ "JWT 直连": "JWT 直连",
+ "OAuth 授权码": "OAuth 授权码",
+ "JWT 直连登录": "JWT 直连登录",
+ "添加身份提供商": "添加身份提供商",
+ "编辑身份提供商": "编辑身份提供商",
+ "浏览器回调 URL": "浏览器回调 URL",
+ "配置自定义外部身份提供商,支持 OAuth Code Flow 和 JWT Direct 两种接入模式": "配置自定义外部身份提供商,支持 OAuth Code Flow 和 JWT Direct 两种接入模式",
+ "JWT Direct 支持 direct_token、ticket_exchange、ticket_validate 三种获取模式,并支持 claims 或 userinfo 两类身份解析方式": "JWT Direct 支持 direct_token、ticket_exchange、ticket_validate 三种获取模式,并支持 claims 或 userinfo 两类身份解析方式",
+ "查询参数": "查询参数",
+ "URL 片段": "URL 片段",
+ "请求体(仅 API)": "请求体(仅 API)",
+ "直接回调 JWT": "直接回调 JWT",
+ "票据换取 JWT": "票据换取 JWT",
+ "票据校验(CAS serviceValidate)": "票据校验(CAS serviceValidate)",
+ "本地验签并解析 JWT Claims": "本地验签并解析 JWT Claims",
+ "通过用户信息端点解析身份": "通过用户信息端点解析身份",
+ "查询字符串": "查询字符串",
+ "表单 URL 编码": "表单 URL 编码",
+ "JSON 请求体": "JSON 请求体",
+ "Multipart 表单": "Multipart 表单",
+ "仅显式映射": "仅显式映射",
+ "映射优先,其次透传": "映射优先,其次透传",
+ "请填写 {{fieldLabel}}": "请填写 {{fieldLabel}}",
+ "清空已保存的客户端密钥": "清空已保存的客户端密钥",
+ "票据处理模式必须填写有效的 Ticket Processing URL": "票据处理模式必须填写有效的 Ticket Processing URL"
}
}
diff --git a/web/src/i18n/locales/zh-TW.json b/web/src/i18n/locales/zh-TW.json
index 4eb73fb00a92..dda02bd622d9 100644
--- a/web/src/i18n/locales/zh-TW.json
+++ b/web/src/i18n/locales/zh-TW.json
@@ -3111,6 +3111,243 @@
"(当前仅支持易支付接口,默认使用上方服务器地址作为回调地址!)": "(當前僅支援易支付接口,預設使用上方伺服器位址作為回調位址!)",
",当前无生效订阅,将自动使用钱包": ",當前無生效訂閱,將自動使用錢包",
",时间:": ",時間:",
- ",点击更新": ",點擊更新"
+ ",点击更新": ",點擊更新",
+ "个已过期": "個已過期",
+ "订阅": "訂閱",
+ "至": "至",
+ "过期于": "過期於",
+ "作废于": "作廢於",
+ "购买套餐后即可享受模型权益": "購買訂閱後即可享受模型權益",
+ "限购": "限購",
+ "推荐": "推薦",
+ "已达到购买上限": "已達到購買上限",
+ "已达上限": "已達上限",
+ "立即订阅": "立即訂閱",
+ "暂无可购买套餐": "暫無可購買訂閱",
+ "该套餐未配置 Stripe": "該訂閱未設定 Stripe",
+ "已打开支付页面": "已打開支付頁面",
+ "支付失败": "支付失敗",
+ "该套餐未配置 Creem": "該訂閱未設定 Creem",
+ "已发起支付": "已發起支付",
+ "购买订阅套餐": "購買訂閱",
+ "套餐名称": "訂閱名稱",
+ "应付金额": "應付金額",
+ "支付": "支付",
+ "管理员未开启在线支付功能,请联系管理员配置。": "管理員未開啟在線支付功能,請聯繫管理員設定。",
+ "偏好设置": "偏好設定",
+ "界面语言和其他个人偏好": "界面語言和其他個人偏好",
+ "语言偏好": "語言偏好",
+ "选择您的首选界面语言,设置将自动保存并同步到所有设备": "選擇您的首選界面語言,設定將自動儲存並同步到所有設備",
+ "语言偏好已保存": "語言偏好已儲存",
+ "提示:语言偏好会同步到您登录的所有设备,并影响API返回的错误消息语言。": "提示:語言偏好會同步到您登錄的所有設備,並影響API返回的錯誤消息語言。",
+ "自定义 OAuth 提供商": "自訂 OAuth 提供商",
+ "配置自定义 OAuth 提供商,支持 GitHub Enterprise、GitLab、Gitea、Nextcloud、Keycloak、ORY 等兼容 OAuth 2.0 协议的身份提供商": "設定自訂 OAuth 提供商,支援 GitHub Enterprise、GitLab、Gitea、Nextcloud、Keycloak、ORY 等兼容 OAuth 2.0 協議的身份提供商",
+ "回调 URL 格式": "回調 URL 格式",
+ "添加提供商": "添加提供商",
+ "编辑提供商": "編輯提供商",
+ "选择预设...": "選擇設定檔...",
+ "输入基础 URL": "輸入基礎 URL",
+ "例如": "例如",
+ "提供商名称": "提供商名稱",
+ "标识符 (Slug)": "標識符 (Slug)",
+ "授权端点": "授權端點",
+ "令牌端点": "令牌端點",
+ "用户信息端点": "使用者資訊端點",
+ "用户 ID 字段": "使用者 ID 字段",
+ "支持 JSONPath,如 sub, id, data.user.id": "支援 JSONPath,如 sub, id, data.user.id",
+ "用户名字段": "使用者名字段",
+ "支持 JSONPath,如 preferred_username, login, data.user.username": "支援 JSONPath,如 preferred_username, login, data.user.username",
+ "显示名称字段": "顯示名稱字段",
+ "支持 JSONPath,如 name, display_name, data.user.name": "支援 JSONPath,如 name, display_name, data.user.name",
+ "邮箱字段": "信箱字段",
+ "支持 JSONPath,如 email, data.user.email": "支援 JSONPath,如 email, data.user.email",
+ "授权范围 (Scopes)": "授權範圍 (Scopes)",
+ "认证方式": "認證方式",
+ "参数传递": "參數傳遞",
+ "Basic Auth 头": "Basic Auth 頭",
+ "暂无自定义 OAuth 提供商": "暫無自訂 OAuth 提供商",
+ "确定要删除该提供商吗?": "確定要刪除該提供商嗎?",
+ "确定要解绑 {{name}} 吗?": "確定要解綁 {{name}} 嗎?",
+ "解绑成功": "解綁成功",
+ "{{name}} ID": "{{name}} ID",
+ "使用 {{name}} 继续": "使用 {{name}} 繼續",
+ "端点 URL 必须以 http:// 或 https:// 开头:": "端點 URL 必須以 http:// 或 https:// 開頭:",
+ "OAuth 配置错误:授权端点必须是完整的 URL(以 http:// 或 https:// 开头)": "OAuth 設定錯誤:授權端點必須是完整的 URL(以 http:// 或 https:// 開頭)",
+ "OAuth 登录失败:": "OAuth 登錄失敗:",
+ "必填:请输入服务器地址以自动生成完整端点 URL": "必填:請輸入伺服器位址以自動生成完整端點 URL",
+ "填写服务器地址后自动生成:": "填寫伺服器位址後自動生成:",
+ "自动生成:": "自動生成:",
+ "请先填写服务器地址,以自动生成完整的端点 URL": "請先填寫伺服器位址,以自動生成完整的端點 URL",
+ "端点 URL 必须是完整地址(以 http:// 或 https:// 开头)": "端點 URL 必須是完整位址(以 http:// 或 https:// 開頭)",
+ "未匹配到模型,按回车键可将「{{name}}」作为自定义模型名添加": "未匹配到模型,按下 Enter 鍵可將「{{name}}」作為自訂模型名稱新增",
+ "分组相关设置": "分組相關設定",
+ "保存分组相关设置": "保存分組相關設定",
+ "此页面仅显示未设置价格或基础倍率的模型,设置后会自动从列表中移出": "此頁面僅顯示未設定價格或基礎倍率的模型,設定後會自動從列表中移出",
+ "没有未设置定价的模型": "沒有未設定定價的模型",
+ "当前没有未设置定价的模型": "目前沒有未設定定價的模型",
+ "模型计费编辑器": "模型計費編輯器",
+ "价格摘要": "價格摘要",
+ "当前提示": "目前提示",
+ "这个界面默认按价格填写,保存时会自动换算回后端需要的倍率 JSON。": "這個介面預設按價格填寫,儲存時會自動換算回後端需要的倍率 JSON。",
+ "当前未启用,需要时再打开即可。": "目前未啟用,需要時再開啟即可。",
+ "下面展示这个模型保存后会写入哪些后端字段,便于和原始 JSON 编辑框保持一致。": "下方會顯示此模型儲存後將寫入哪些後端欄位,方便與原始 JSON 編輯框保持一致。",
+ "补全价格已锁定": "補全價格已鎖定",
+ "后端固定倍率:{{ratio}}。该字段仅展示换算后的价格。": "後端固定倍率:{{ratio}}。此欄位僅展示換算後的價格。",
+ "这些价格都是可选项,不填也可以。": "這些價格都是可選項,不填也可以。",
+ "请先开启并填写音频输入价格。": "請先開啟並填寫音訊輸入價格。",
+ "输入模型名称,例如 gpt-4.1": "輸入模型名稱,例如 gpt-4.1",
+ "当前模型同时存在按次价格和倍率配置,保存时会按当前计费方式覆盖。": "目前模型同時存在按次價格與倍率配置,儲存時會依目前計費方式覆蓋。",
+ "当前模型存在未显式设置输入倍率的扩展倍率;填写输入价格后会自动换算为价格字段。": "目前模型存在未明確設定輸入倍率的擴展倍率;填寫輸入價格後會自動換算為價格欄位。",
+ "按量计费下需要先填写输入价格,才能保存其它价格项。": "按量計費下需要先填寫輸入價格,才能儲存其它價格項。",
+ "填写音频补全价格前,需要先填写音频输入价格。": "填寫音訊補全價格前,需要先填寫音訊輸入價格。",
+ "模型 {{name}} 缺少输入价格,无法计算补全/缓存/图片/音频价格对应的倍率": "模型 {{name}} 缺少輸入價格,無法計算補全、快取、圖片與音訊價格對應的倍率",
+ "模型 {{name}} 缺少音频输入价格,无法计算音频补全倍率": "模型 {{name}} 缺少音訊輸入價格,無法計算音訊補全倍率",
+ "批量应用当前模型价格": "批量套用目前模型價格",
+ "请先选择一个作为模板的模型": "請先選擇一個作為範本的模型",
+ "请先勾选需要批量设置的模型": "請先勾選需要批量設定的模型",
+ "已将模型 {{name}} 的价格配置批量应用到 {{count}} 个模型": "已將模型 {{name}} 的價格配置批量套用到 {{count}} 個模型",
+ "将把当前编辑中的模型 {{name}} 的价格配置,批量应用到已勾选的 {{count}} 个模型。": "會把目前編輯中的模型 {{name}} 的價格配置,批量套用到已勾選的 {{count}} 個模型。",
+ "适合同系列模型一起定价,例如把 gpt-5.1 的价格批量同步到 gpt-5.1-high、gpt-5.1-low 等模型。": "適合同系列模型一起定價,例如把 gpt-5.1 的價格批量同步到 gpt-5.1-high、gpt-5.1-low 等模型。",
+ "已勾选": "已勾選",
+ "当前编辑": "目前編輯",
+ "已勾选 {{count}} 个模型": "已勾選 {{count}} 個模型",
+ "基础价格": "基礎價格",
+ "扩展价格": "擴展價格",
+ "额外价格项": "額外價格項",
+ "补全价格": "補全價格",
+ "缓存读取价格": "快取讀取價格",
+ "缓存创建价格": "快取建立價格",
+ "图片输入价格": "圖片輸入價格",
+ "音频输入价格": "音訊輸入價格",
+ "音频补全价格": "音訊補全價格",
+ "适合 MJ / 任务类等按次收费模型。": "適合 MJ / 任務類等按次收費模型。",
+ "该模型补全倍率由后端固定为 {{ratio}}。补全价格不能在这里修改。": "該模型補全倍率由後端固定為 {{ratio}}。補全價格不能在這裡修改。",
+ "计费显示模式": "計費顯示模式",
+ "价格模式(默认)": "價格模式(預設)",
+ "模型价格 {{symbol}}{{price}} / 次": "模型價格 {{symbol}}{{price}} / 次",
+ "按次 {{symbol}}{{price}} * {{ratioType}} {{ratio}} = {{symbol}}{{total}}": "按次 {{symbol}}{{price}} * {{ratioType}} {{ratio}} = {{symbol}}{{total}}",
+ "模型价格:{{symbol}}{{price}} / 次": "模型價格:{{symbol}}{{price}} / 次",
+ "按次:{{symbol}}{{price}}": "按次:{{symbol}}{{price}}",
+ "实际结算金额:{{symbol}}{{total}}(已包含分组价格调整)": "實際結算金額:{{symbol}}{{total}}(已包含分組價格調整)",
+ "缓存读取价格:{{symbol}}{{price}} / 1M tokens": "快取讀取價格:{{symbol}}{{price}} / 1M tokens",
+ "缓存读取价格 {{symbol}}{{price}} / 1M tokens": "快取讀取價格 {{symbol}}{{price}} / 1M tokens",
+ "缓存创建价格:{{symbol}}{{price}} / 1M tokens": "快取建立價格:{{symbol}}{{price}} / 1M tokens",
+ "缓存创建价格 {{symbol}}{{price}} / 1M tokens": "快取建立價格 {{symbol}}{{price}} / 1M tokens",
+ "5m缓存创建价格:{{symbol}}{{price}} / 1M tokens": "5m快取建立價格:{{symbol}}{{price}} / 1M tokens",
+ "5m缓存创建价格 {{symbol}}{{price}} / 1M tokens": "5m快取建立價格 {{symbol}}{{price}} / 1M tokens",
+ "1h缓存创建价格:{{symbol}}{{price}} / 1M tokens": "1h快取建立價格:{{symbol}}{{price}} / 1M tokens",
+ "1h缓存创建价格 {{symbol}}{{price}} / 1M tokens": "1h快取建立價格 {{symbol}}{{price}} / 1M tokens",
+ "图片输入价格:{{symbol}}{{price}} / 1M tokens": "圖片輸入價格:{{symbol}}{{price}} / 1M tokens",
+ "图片输入价格 {{symbol}}{{price}} / 1M tokens": "圖片輸入價格 {{symbol}}{{price}} / 1M tokens",
+ "输入价格 {{symbol}}{{price}} / 1M tokens": "輸入價格 {{symbol}}{{price}} / 1M tokens",
+ "音频输入价格:{{symbol}}{{price}} / 1M tokens": "音訊輸入價格:{{symbol}}{{price}} / 1M tokens",
+ "音频补全价格:{{symbol}}{{price}} / 1M tokens": "音訊補全價格:{{symbol}}{{price}} / 1M tokens",
+ "Web 搜索调用 {{webSearchCallCount}} 次": "Web 搜尋呼叫 {{webSearchCallCount}} 次",
+ "文件搜索调用 {{fileSearchCallCount}} 次": "檔案搜尋呼叫 {{fileSearchCallCount}} 次",
+ "图片倍率 {{imageRatio}}": "圖片倍率 {{imageRatio}}",
+ "音频倍率 {{audioRatio}}": "音訊倍率 {{audioRatio}}",
+ "普通输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "普通輸入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "缓存输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 缓存倍率 {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "快取輸入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 快取倍率 {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "图片输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 图片倍率 {{imageRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "圖片輸入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 圖片倍率 {{imageRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "音频输入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 音频倍率 {{audioRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "音訊輸入:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 音訊倍率 {{audioRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 补全倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "輸出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 補全倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "Web 搜索:{{count}} / 1K * 单价 {{price}} * {{ratioType}} {{ratio}} = {{amount}}": "Web 搜尋:{{count}} / 1K * 單價 {{price}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "文件搜索:{{count}} / 1K * 单价 {{price}} * {{ratioType}} {{ratio}} = {{amount}}": "檔案搜尋:{{count}} / 1K * 單價 {{price}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "图片生成:1 次 * 单价 {{price}} * {{ratioType}} {{ratio}} = {{amount}}": "圖片生成:1 次 * 單價 {{price}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "合计:{{total}}": "合計:{{total}}",
+ "模型倍率 {{modelRatio}},补全倍率 {{completionRatio}},音频倍率 {{audioRatio}},音频补全倍率 {{audioCompletionRatio}},{{cachePart}}{{ratioType}} {{ratio}}": "模型倍率 {{modelRatio}},補全倍率 {{completionRatio}},音訊倍率 {{audioRatio}},音訊補全倍率 {{audioCompletionRatio}},{{cachePart}}{{ratioType}} {{ratio}}",
+ "文字输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 补全倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "文字輸出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 補全倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "音频输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 音频倍率 {{audioRatio}} * 音频补全倍率 {{audioCompletionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "音訊輸出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 音訊倍率 {{audioRatio}} * 音訊補全倍率 {{audioCompletionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "合计:文字部分 {{textTotal}} + 音频部分 {{audioTotal}} = {{total}}": "合計:文字部分 {{textTotal}} + 音訊部分 {{audioTotal}} = {{total}}",
+ "模型倍率 {{modelRatio}},输出倍率 {{completionRatio}},缓存倍率 {{cacheRatio}},{{ratioType}} {{ratio}}": "模型倍率 {{modelRatio}},輸出倍率 {{completionRatio}},快取倍率 {{cacheRatio}},{{ratioType}} {{ratio}}",
+ "缓存读取:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 缓存倍率 {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "快取讀取:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 快取倍率 {{cacheRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "缓存创建:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 缓存创建倍率 {{cacheCreationRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "快取建立:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 快取建立倍率 {{cacheCreationRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "5m缓存创建:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 5m缓存创建倍率 {{cacheCreationRatio5m}} * {{ratioType}} {{ratio}} = {{amount}}": "5m快取建立:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 5m快取建立倍率 {{cacheCreationRatio5m}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "1h缓存创建:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 1h缓存创建倍率 {{cacheCreationRatio1h}} * {{ratioType}} {{ratio}} = {{amount}}": "1h快取建立:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 1h快取建立倍率 {{cacheCreationRatio1h}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "输出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 输出倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}": "輸出:{{tokens}} / 1M * 模型倍率 {{modelRatio}} * 輸出倍率 {{completionRatio}} * {{ratioType}} {{ratio}} = {{amount}}",
+ "空": "空",
+ "提示:端点映射仅用于模型广场展示,不会影响模型真实调用。如需配置真实调用,请前往「渠道管理」。": "提示:端點映射僅用於模型廣場展示,不會影響模型真實呼叫。如需配置真實呼叫,請前往「管道管理」。",
+ "购买订阅获得模型额度/次数": "購買訂閱取得模型額度/次數",
+ "生产环境 RSA 私钥 Base64 (PKCS#8 DER)": "正式環境 RSA 私鑰 Base64 (PKCS#8 DER)",
+ "沙盒环境 RSA 私钥 Base64 (PKCS#8 DER)": "沙盒環境 RSA 私鑰 Base64 (PKCS#8 DER)",
+ "生产环境 Waffo 公钥 Base64 (X.509 DER)": "正式環境 Waffo 公鑰 Base64 (X.509 DER)",
+ "沙盒环境 Waffo 公钥 Base64 (X.509 DER)": "沙盒環境 Waffo 公鑰 Base64 (X.509 DER)",
+ "支付方式类型": "付款方式類型",
+ "支付方式名称": "付款方式名稱",
+ "获取充值配置失败": "取得儲值設定失敗",
+ "获取充值配置异常": "儲值設定異常",
+ "{{ratioType}} {{ratio}}x": "{{ratioType}} {{ratio}}x",
+ "模型价格:{{symbol}}{{price}}": "模型價格:{{symbol}}{{price}}",
+ "模型价格 {{price}}": "模型價格 {{price}}",
+ "缓存读 {{price}} / 1M tokens": "快取讀 {{price}} / 1M tokens",
+ "5m缓存创建 {{price}} / 1M tokens": "5m快取建立 {{price}} / 1M tokens",
+ "1h缓存创建 {{price}} / 1M tokens": "1h快取建立 {{price}} / 1M tokens",
+ "缓存创建 {{price}} / 1M tokens": "快取建立 {{price}} / 1M tokens",
+ "图片输入 {{price}} / 1M tokens": "圖片輸入 {{price}} / 1M tokens",
+ "输入 {{price}} / 1M tokens": "輸入 {{price}} / 1M tokens",
+ "缓存 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "快取 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "缓存创建 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "快取建立 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "5m缓存创建 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "5m快取建立 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "1h缓存创建 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}": "1h快取建立 {{tokens}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "Key": "Key",
+ "Key 摘要": "Key 摘要",
+ "写": "寫",
+ "异步任务退款": "非同步任務退款",
+ "扣费": "扣費",
+ "根据 Anthropic 协定,/v1/messages 的输入 tokens 仅统计非缓存输入,不包含缓存读取与缓存写入 tokens。": "根據 Anthropic 協定,/v1/messages 的輸入 tokens 僅統計非快取輸入,不包含快取讀取與快取寫入 tokens。",
+ "渠道亲和性": "渠道親和性",
+ "由订阅抵扣": "由訂閱抵扣",
+ "缓存写": "快取寫",
+ "缓存读": "快取讀",
+ "规则": "規則",
+ "订阅抵扣": "訂閱抵扣",
+ "违规扣费": "違規扣費",
+ "退款": "退款",
+ "(输入 {{nonImageInput}} tokens + 图片输入 {{imageInput}} tokens / 1M tokens * {{symbol}}{{price}}": "(輸入 {{nonImageInput}} tokens + 圖片輸入 {{imageInput}} tokens / 1M tokens * {{symbol}}{{price}}",
+ "图片输入价格:{{symbol}}{{total}} / 1M tokens": "圖片輸入價格:{{symbol}}{{total}} / 1M tokens",
+ "文字提示 {{input}} tokens / 1M tokens * {{symbol}}{{textInputPrice}} + 文字补全 {{completion}} tokens / 1M tokens * {{symbol}}{{textCompPrice}} + 音频提示 {{audioInput}} tokens / 1M tokens * {{symbol}}{{audioInputPrice}} + 音频补全 {{audioCompletion}} tokens / 1M tokens * {{symbol}}{{audioCompPrice}} * {{ratioType}} {{ratio}} = {{symbol}}{{total}}": "文字提示 {{input}} tokens / 1M tokens * {{symbol}}{{textInputPrice}} + 文字補全 {{completion}} tokens / 1M tokens * {{symbol}}{{textCompPrice}} + 音訊提示 {{audioInput}} tokens / 1M tokens * {{symbol}}{{audioInputPrice}} + 音訊補全 {{audioCompletion}} tokens / 1M tokens * {{symbol}}{{audioCompPrice}} * {{ratioType}} {{ratio}} = {{symbol}}{{total}}",
+ "模型价格 {{symbol}}{{price}} / 次 * {{ratioType}} {{ratio}} = {{symbol}}{{total}}": "模型價格 {{symbol}}{{price}} / 次 * {{ratioType}} {{ratio}} = {{symbol}}{{total}}",
+ "缓存读取价格:{{symbol}}{{total}} / 1M tokens": "快取讀取價格:{{symbol}}{{total}} / 1M tokens",
+ "补全 {{completion}} tokens * 输出倍率 {{completionRatio}}": "補全 {{completion}} tokens * 輸出倍率 {{completionRatio}}",
+ "补全倍率 {{completionRatio}}": "補全倍率 {{completionRatio}}",
+ "输入价格:{{symbol}}{{price}} / 1M tokens": "輸入價格:{{symbol}}{{price}} / 1M tokens",
+ "输出价格 {{symbol}}{{price}} / 1M tokens": "輸出價格 {{symbol}}{{price}} / 1M tokens",
+ "输出价格:{{symbol}}{{price}} / 1M tokens": "輸出價格:{{symbol}}{{price}} / 1M tokens",
+ "输出价格:{{symbol}}{{total}} / 1M tokens": "輸出價格:{{symbol}}{{total}} / 1M tokens",
+ "复制密钥": "複製金鑰",
+ "复制连接信息": "複製連線資訊",
+ "检测到剪贴板中的连接信息": "偵測到剪貼簿中的連線資訊",
+ "自动填入": "自動填入",
+ "忽略": "忽略",
+ "从剪贴板粘贴配置": "從剪貼簿貼上設定",
+ "剪贴板中未检测到连接信息": "剪貼簿中未偵測到連線資訊",
+ "连接信息已填入": "連線資訊已填入",
+ "无法读取剪贴板": "無法讀取剪貼簿",
+ "OAuth 2.0 / OIDC 授权码模式": "OAuth 2.0 / OIDC 授權碼模式",
+ "JWT 直连": "JWT 直連",
+ "OAuth 授权码": "OAuth 授權碼",
+ "JWT 直连登录": "JWT 直連登入",
+ "添加身份提供商": "添加身份提供商",
+ "编辑身份提供商": "編輯身份提供商",
+ "浏览器回调 URL": "瀏覽器回呼 URL",
+ "配置自定义外部身份提供商,支持 OAuth Code Flow 和 JWT Direct 两种接入模式": "配置自訂外部身份提供商,支援 OAuth Code Flow 和 JWT Direct 兩種接入模式",
+ "JWT Direct 支持 direct_token、ticket_exchange、ticket_validate 三种获取模式,并支持 claims 或 userinfo 两类身份解析方式": "JWT Direct 支援 direct_token、ticket_exchange、ticket_validate 三種取得模式,並支援 claims 或 userinfo 兩類身份解析方式",
+ "查询参数": "查詢參數",
+ "URL 片段": "URL 片段",
+ "请求体(仅 API)": "請求體(僅 API)",
+ "直接回调 JWT": "直接回呼 JWT",
+ "票据换取 JWT": "票據換取 JWT",
+ "票据校验(CAS serviceValidate)": "票據校驗(CAS serviceValidate)",
+ "本地验签并解析 JWT Claims": "本地驗簽並解析 JWT Claims",
+ "通过用户信息端点解析身份": "透過使用者資訊端點解析身份",
+ "查询字符串": "查詢字串",
+ "表单 URL 编码": "表單 URL 編碼",
+ "JSON 请求体": "JSON 請求體",
+ "Multipart 表单": "Multipart 表單",
+ "仅显式映射": "僅顯式映射",
+ "映射优先,其次透传": "映射優先,其次透傳",
+ "请填写 {{fieldLabel}}": "請填寫 {{fieldLabel}}",
+ "清空已保存的客户端密钥": "清空已保存的客戶端密鑰",
+ "票据处理模式必须填写有效的 Ticket Processing URL": "票據處理模式必須填寫有效的 Ticket Processing URL"
}
}