fix(oauth): secure OAuth state session binding to prevent account takeover

Bind OAuth state payloads to the initiating session token and user ID.
Verifies session token hash continuity during callback, and validates that
the user ID completing the binding flow matches the user ID that initiated it.
This commit is contained in:
ryan
2026-06-13 09:55:07 +08:00
parent f48426dbf8
commit 895788974c
7 changed files with 388 additions and 64 deletions
+87 -6
View File
@@ -5,14 +5,15 @@ package oauth
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"strconv"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
@@ -73,6 +74,23 @@ func GetUserIDFromContext(c *gin.Context) uint64 {
return GetUserIDFromSession(session)
}
func ensureSessionToken(s sessions.Session) (string, bool) {
token, ok := s.Get(SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
s.Set(SessionTokenKey, token)
return token, true
}
return token, false
}
func hashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
@@ -335,10 +353,25 @@ func GetLoginURL(c *gin.Context) {
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
return
}
session := sessions.Default(c)
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
return
}
}
userID := GetUserIDFromSession(session)
sessionHash := hashSessionToken(token)
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: OAuthPurposeLogin,
SourceName: source.Name,
Purpose: OAuthPurposeLogin,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
@@ -349,6 +382,7 @@ func GetLoginURL(c *gin.Context) {
return
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
@@ -397,10 +431,30 @@ func Authorize(c *gin.Context) {
if purpose != OAuthPurposeBind {
purpose = OAuthPurposeLogin
}
session := sessions.Default(c)
userID := GetUserIDFromSession(session)
if purpose == OAuthPurposeBind && userID == 0 {
c.JSON(http.StatusUnauthorized, util.Err(common.UnAuthorized))
return
}
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
return
}
}
sessionHash := hashSessionToken(token)
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: purpose,
SourceName: source.Name,
Purpose: purpose,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
@@ -410,6 +464,7 @@ func Authorize(c *gin.Context) {
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
return
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
@@ -452,6 +507,32 @@ func Callback(c *gin.Context) {
return
}
session := sessions.Default(c)
currentUserID := GetUserIDFromSession(session)
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
c.JSON(http.StatusUnauthorized, util.Err(common.UnAuthorized))
return
}
token, ok := session.Get(SessionTokenKey).(string)
if !ok || token == "" {
c.JSON(http.StatusBadRequest, util.Err("invalid session context"))
return
}
if hashSessionToken(token) != payload.SessionHash {
c.JSON(http.StatusBadRequest, util.Err("session mismatch for oauth state"))
return
}
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
c.JSON(http.StatusBadRequest, util.Err("user context mismatch for oauth binding"))
return
}
source, err := resolveAuthSource(payload.SourceName)
if err != nil {
c.JSON(http.StatusBadRequest, util.Err(err.Error()))