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:
ryan
2026-06-13 10:06:49 +08:00
parent 895788974c
commit 5412c385dc
5 changed files with 153 additions and 59 deletions
-4
View File
@@ -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
)
+113
View File
@@ -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())
}
}
+40 -9
View File
@@ -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()))
-43
View File
@@ -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"`
-3
View File
@@ -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)))
}