diff --git a/internal/apps/oauth/constants.go b/internal/apps/oauth/constants.go index 3fde9096..bb066a61 100644 --- a/internal/apps/oauth/constants.go +++ b/internal/apps/oauth/constants.go @@ -16,10 +16,6 @@ const ( UserObjKey = "user_obj" TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权 TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限 - PendingOAuthSourceIDKey = "pending_oauth_source_id" - PendingOAuthExternalIDKey = "pending_oauth_external_id" - PendingOAuthExternalUsernameKey = "pending_oauth_external_username" - PendingOAuthEmailKey = "pending_oauth_email" SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials ) diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index 1fd99ab1..dd254e31 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -1068,3 +1068,116 @@ func TestExternalAccountsListAndDelete(t *testing.T) { t.Error("binding record was not deleted from DB") } } + +func TestOIDCPolicyEnforcement(t *testing.T) { + initializeTestConfig() + dbConn := setupTestDB(t) + mockRedis := newMockRedisClient() + seedTestAuthSource(t, dbConn) // seeds testSourceName ("linuxdo") active=true + + // Set up mock client & router + var state string + httpMock := newMockOIDCClient(testIssuerURL, testClientID, &state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User") + util.SetHTTPClient(httpMock) + router := setupTestRouter(dbConn, mockRedis, httpMock) + + // --- 1. Test GetLoginURL enforcement --- + // Disable globally + dbConn.Create(&model.SystemConfig{ + Key: model.ConfigKeyOIDCLoginEnabled, + Value: "false", + }) + mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) + if wLoginDisabled.Code != http.StatusBadRequest { + t.Errorf("expected 400 when OIDC globally disabled, got %d", wLoginDisabled.Code) + } + + // Re-enable globally, but deactivate source + dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") + mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) + + wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) + if wSourceInactive.Code != http.StatusBadRequest { + t.Errorf("expected 400 when OIDC source is inactive, got %d", wSourceInactive.Code) + } + + // --- 2. Test Authorize enforcement --- + // Deactivate globally again + dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") + mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) + + wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil) + if wAuthDisabled.Code != http.StatusBadRequest { + t.Errorf("expected 400 when OIDC globally disabled in Authorize, got %d", wAuthDisabled.Code) + } + + // --- 3. Test Callback enforcement --- + // Set up a valid state beforehand (when enabled) + dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") + mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) + + wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) + + if wLogin.Code != http.StatusOK { + t.Fatalf("failed to setup login: %s", wLogin.Body.String()) + } + var loginUrlResp struct { + Data OAuthAuthorizeResponse `json:"data"` + } + _ = json.Unmarshal(wLogin.Body.Bytes(), &loginUrlResp) + parsedURL, _ := url.Parse(loginUrlResp.Data.AuthorizeURL) + state = parsedURL.Query().Get("state") + + var anonymousCookie *http.Cookie + for _, cookie := range wLogin.Result().Cookies() { + if cookie.Name == config.Config.App.SessionCookieName { + anonymousCookie = cookie + break + } + } + + // Now disable OIDC globally and attempt callback + dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") + mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) + wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ + "Content-Type": "application/json", + }, []*http.Cookie{anonymousCookie}) + if wCallbackDisabled.Code != http.StatusBadRequest { + t.Errorf("expected 400 for callback when OIDC globally disabled, got %d, body: %s", wCallbackDisabled.Code, wCallbackDisabled.Body.String()) + } + + // Enable globally but deactivate source and attempt callback + dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") + mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) + dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) + + // Since callback deletes state, we need to generate state again + dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) + wLogin2 := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) + _ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp) + parsedURL, _ = url.Parse(loginUrlResp.Data.AuthorizeURL) + state = parsedURL.Query().Get("state") + var anonymousCookie2 *http.Cookie + for _, cookie := range wLogin2.Result().Cookies() { + if cookie.Name == config.Config.App.SessionCookieName { + anonymousCookie2 = cookie + break + } + } + + // Deactivate source + dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) + reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) + wCallbackSourceInactive := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{ + "Content-Type": "application/json", + }, []*http.Cookie{anonymousCookie2}) + if wCallbackSourceInactive.Code != http.StatusBadRequest { + t.Errorf("expected 400 for callback when OIDC source deactivated, got %d, body: %s", wCallbackSourceInactive.Code, wCallbackSourceInactive.Body.String()) + } +} + diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index c08c282f..5651580d 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -90,6 +90,15 @@ func hashSessionToken(token string) string { return hex.EncodeToString(h.Sum(nil)) } +func isOIDCLoginEnabled(ctx context.Context) bool { + enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) + if err != nil { + return true + } + return enabled +} + + func resolveAuthSource(sourceName string) (*model.AuthSource, error) { name := strings.TrimSpace(strings.ToLower(sourceName)) @@ -348,12 +357,24 @@ func GetLoginSources(c *gin.Context) { // @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败" // @Router /api/v1/oauth/login [get] func GetLoginURL(c *gin.Context) { + ctx := c.Request.Context() + if !isOIDCLoginEnabled(ctx) { + c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + return + } + source, err := resolveAuthSource(c.Query("source")) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } + if !source.IsActive { + c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + return + } + + session := sessions.Default(c) token, isNew := ensureSessionToken(session) if isNew { @@ -418,11 +439,18 @@ func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state stri // @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败" // @Router /api/v1/oauth/{source}/authorize [get] func Authorize(c *gin.Context) { + ctx := c.Request.Context() + if !isOIDCLoginEnabled(ctx) { + c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + return + } + source, err := resolveAuthSource(c.Param("source")) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } + if !source.IsActive { c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) return @@ -533,12 +561,23 @@ func Callback(c *gin.Context) { + if !isOIDCLoginEnabled(ctx) { + c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + return + } + source, err := resolveAuthSource(payload.SourceName) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) return } + if !source.IsActive { + c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled)) + return + } + + redirectURL, err := getFrontendLoginRedirectURL(ctx) if err != nil { c.JSON(http.StatusBadRequest, util.Err(err.Error())) @@ -633,19 +672,11 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A } if !registrationEnabled { - session := sessions.Default(c) - session.Set(PendingOAuthSourceIDKey, source.ID) - session.Set(PendingOAuthExternalIDKey, userInfo.Sub) - session.Set(PendingOAuthExternalUsernameKey, userInfo.Username) - session.Set(PendingOAuthEmailKey, userInfo.Email) - if err := session.Save(); err != nil { - c.JSON(http.StatusInternalServerError, util.Err(err.Error())) - return model.User{}, false - } c.JSON(http.StatusOK, util.OK(buildCallbackResult(nil, "need_bind"))) return model.User{}, false } + username, uniqueErr := uniqueUsername(ctx, userInfo.Username) if uniqueErr != nil { c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error())) diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index ea0e4bd7..26e72e2a 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -241,50 +241,7 @@ func validateRegisterEmailVerification(ctx context.Context, req *registerRequest return nil } -// completePendingOAuthBinding 完成登录后的 OAuth 待绑定绑定流程 -func completePendingOAuthBinding(ctx context.Context, session sessions.Session, user *model.User) { - pendingSourceID := session.Get(oauth.PendingOAuthSourceIDKey) - pendingExternalID := session.Get(oauth.PendingOAuthExternalIDKey) - pendingExternalUsername := session.Get(oauth.PendingOAuthExternalUsernameKey) - pendingEmail := session.Get(oauth.PendingOAuthEmailKey) - if pendingSourceID == nil || pendingExternalID == nil { - return - } - - var sourceID uint64 - switch v := pendingSourceID.(type) { - case uint64: - sourceID = v - case int: - if v >= 0 { - sourceID = uint64(v) - } - case float64: - if v >= 0 && v <= 18446744073709551615.0 { - sourceID = uint64(v) - } - } - externalID, _ := pendingExternalID.(string) - externalUsername, _ := pendingExternalUsername.(string) - email, _ := pendingEmail.(string) - - if sourceID != 0 && externalID != "" { - _ = model.BindExternalAccount(ctx, &model.ExternalAccount{ - AuthSourceID: sourceID, - UserID: user.ID, - ExternalID: externalID, - ExternalUsername: externalUsername, - Email: email, - }) - } - - session.Delete(oauth.PendingOAuthSourceIDKey) - session.Delete(oauth.PendingOAuthExternalIDKey) - session.Delete(oauth.PendingOAuthExternalUsernameKey) - session.Delete(oauth.PendingOAuthEmailKey) - _ = session.Save() -} type updateProfileRequest struct { Nickname string `json:"nickname"` diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index b55185da..1e632a9d 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -164,9 +164,6 @@ func Login(c *gin.Context) { return } - // 检查是否有未完成 of OAuth/OIDC 绑定 - completePendingOAuthBinding(ctx, session, &user) - c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&user, needChangePassword))) }