mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-28 05:46:36 +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"
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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()))
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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)))
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user