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 }) => { @@ -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 }) => { { 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, - })), - ]} - /> - - - - - -
- -
- - - - - { + 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, + })), + ]} + /> + + + + + +
+ +
+ + + )} + {!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 /> - - 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" } }