mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 01:36:37 +08:00
fix(oauth): remove pending oauth auto-binding and enforce oidc policies
- Complete removal of completePendingOAuthBinding logic to prevent unintended account takeovers (AUTH-ROUTE-1). - Add strict OIDC policy checks (global switch and source active states) across authorization and callback paths (AUTH-POLICY-1). - Fix OIDC test cases to properly clear the Redis-backed system config cache using composite keys.
This commit is contained in:
@@ -16,10 +16,6 @@ const (
|
|||||||
UserObjKey = "user_obj"
|
UserObjKey = "user_obj"
|
||||||
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
|
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
|
||||||
TokenAdminKey = "token_admin" // 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
|
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1068,3 +1068,116 @@ func TestExternalAccountsListAndDelete(t *testing.T) {
|
|||||||
t.Error("binding record was not deleted from DB")
|
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())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -90,6 +90,15 @@ func hashSessionToken(token string) string {
|
|||||||
return hex.EncodeToString(h.Sum(nil))
|
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) {
|
func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
|
||||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||||
@@ -348,12 +357,24 @@ func GetLoginSources(c *gin.Context) {
|
|||||||
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
||||||
// @Router /api/v1/oauth/login [get]
|
// @Router /api/v1/oauth/login [get]
|
||||||
func GetLoginURL(c *gin.Context) {
|
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"))
|
source, err := resolveAuthSource(c.Query("source"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !source.IsActive {
|
||||||
|
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
token, isNew := ensureSessionToken(session)
|
token, isNew := ensureSessionToken(session)
|
||||||
if isNew {
|
if isNew {
|
||||||
@@ -418,11 +439,18 @@ func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state stri
|
|||||||
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
||||||
// @Router /api/v1/oauth/{source}/authorize [get]
|
// @Router /api/v1/oauth/{source}/authorize [get]
|
||||||
func Authorize(c *gin.Context) {
|
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"))
|
source, err := resolveAuthSource(c.Param("source"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if !source.IsActive {
|
if !source.IsActive {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
||||||
return
|
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)
|
source, err := resolveAuthSource(payload.SourceName)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !source.IsActive {
|
||||||
|
c.JSON(http.StatusBadRequest, util.Err(AuthSourceDisabled))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
redirectURL, err := getFrontendLoginRedirectURL(ctx)
|
redirectURL, err := getFrontendLoginRedirectURL(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
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 {
|
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")))
|
c.JSON(http.StatusOK, util.OK(buildCallbackResult(nil, "need_bind")))
|
||||||
return model.User{}, false
|
return model.User{}, false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
|
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
|
||||||
if uniqueErr != nil {
|
if uniqueErr != nil {
|
||||||
c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error()))
|
c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error()))
|
||||||
|
|||||||
@@ -241,50 +241,7 @@ func validateRegisterEmailVerification(ctx context.Context, req *registerRequest
|
|||||||
return nil
|
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 {
|
type updateProfileRequest struct {
|
||||||
Nickname string `json:"nickname"`
|
Nickname string `json:"nickname"`
|
||||||
|
|||||||
@@ -164,9 +164,6 @@ func Login(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 检查是否有未完成 of OAuth/OIDC 绑定
|
|
||||||
completePendingOAuthBinding(ctx, session, &user)
|
|
||||||
|
|
||||||
c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&user, needChangePassword)))
|
c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&user, needChangePassword)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user