mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
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:
@@ -20,6 +20,7 @@ const (
|
||||
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
|
||||
)
|
||||
|
||||
// OAuth State 缓存 Key 格式与过期时间
|
||||
@@ -35,8 +36,10 @@ const (
|
||||
)
|
||||
|
||||
type oauthStatePayload struct {
|
||||
SourceName string `json:"source_name"`
|
||||
Purpose string `json:"purpose"`
|
||||
SourceName string `json:"source_name"`
|
||||
Purpose string `json:"purpose"`
|
||||
UserID uint64 `json:"user_id,omitempty"`
|
||||
SessionHash string `json:"session_hash"`
|
||||
}
|
||||
|
||||
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
|
||||
@@ -54,3 +57,4 @@ func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -18,4 +18,5 @@ const (
|
||||
UsernameFromSourceFailed = "无法从认证源获取用户名"
|
||||
AuthSourceDisabled = "认证源未启用"
|
||||
InvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
||||
TokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
)
|
||||
|
||||
@@ -101,3 +101,15 @@ func LoginRequired() gin.HandlerFunc {
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
|
||||
func DisallowTokenAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth {
|
||||
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error_msg": TokenAuthNotAllowed, "data": nil})
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -25,7 +25,6 @@ import (
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
"github.com/google/uuid"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"golang.org/x/oauth2"
|
||||
"gorm.io/driver/sqlite"
|
||||
@@ -245,7 +244,7 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string)
|
||||
|
||||
// -----------------------------------------------------------------------------
|
||||
// Test Helpers
|
||||
func newMockOIDCClient(issuer, clientID, expectedState, sub, username, email, name string) *http.Client {
|
||||
func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, username, email, name string) *http.Client {
|
||||
cleanIssuer := strings.TrimRight(issuer, "/")
|
||||
return &http.Client{
|
||||
Transport: &mockRoundTripper{
|
||||
@@ -276,7 +275,11 @@ func newMockOIDCClient(issuer, clientID, expectedState, sub, username, email, na
|
||||
}, nil
|
||||
}
|
||||
if req.Method == http.MethodPost && (strings.Contains(urlStr, "/token") || strings.Contains(urlStr, "/access_token")) {
|
||||
idToken := generateMockIDToken(issuer, sub, clientID, expectedState, username, email, name)
|
||||
var stateVal string
|
||||
if expectedState != nil {
|
||||
stateVal = *expectedState
|
||||
}
|
||||
idToken := generateMockIDToken(issuer, sub, clientID, stateVal, username, email, name)
|
||||
body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
@@ -564,8 +567,29 @@ func TestAuthorize(t *testing.T) {
|
||||
util.SetHTTPClient(httpMock)
|
||||
router := setupTestRouter(dbConn, mockRedis, httpMock)
|
||||
|
||||
// Case 1: Active Source Authorize with purpose=bind
|
||||
w := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, nil)
|
||||
// Case 1a: Active Source Authorize with purpose=bind without login -> 401
|
||||
wUnauth := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, nil)
|
||||
if wUnauth.Code != http.StatusUnauthorized {
|
||||
t.Errorf("expected 401 for unauthorized bind authorize, got %d", wUnauth.Code)
|
||||
}
|
||||
|
||||
// Case 1b: Active Source Authorize with purpose=bind (authenticated)
|
||||
router.GET("/test-helper/login-777", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(UserIDKey, uint64(777))
|
||||
_ = session.Save()
|
||||
c.String(200, "ok")
|
||||
})
|
||||
wLogin := performRequest(router, http.MethodGet, "/test-helper/login-777", nil, nil, nil)
|
||||
var activeCookie *http.Cookie
|
||||
for _, cookie := range wLogin.Result().Cookies() {
|
||||
if cookie.Name == config.Config.App.SessionCookieName {
|
||||
activeCookie = cookie
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
w := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie})
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200, got %d, body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
@@ -573,6 +597,7 @@ func TestAuthorize(t *testing.T) {
|
||||
var resp struct {
|
||||
Data OAuthAuthorizeResponse `json:"data"`
|
||||
}
|
||||
|
||||
_ = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
parsedURL, _ := url.Parse(resp.Data.AuthorizeURL)
|
||||
@@ -611,25 +636,44 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
dbConn := setupTestDB(t)
|
||||
mockRedis := newMockRedisClient()
|
||||
seedTestAuthSource(t, dbConn)
|
||||
state := uuid.NewString()
|
||||
|
||||
var state string
|
||||
|
||||
// 1. Mock the outgoing HTTP client for token exchange and user info fetching
|
||||
httpMock := newMockOIDCClient(testIssuerURL, testClientID, state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User")
|
||||
httpMock := newMockOIDCClient(testIssuerURL, testClientID, &state, "88888", "test_oauth_user", "oauth@linux.do", "Oauth Test User")
|
||||
util.SetHTTPClient(httpMock)
|
||||
router := setupTestRouter(dbConn, mockRedis, httpMock)
|
||||
|
||||
// 2. Setup state in Redis
|
||||
payloadValue, _ := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: testSourceName,
|
||||
Purpose: OAuthPurposeLogin,
|
||||
})
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration)
|
||||
// Get Login URL first to initialize the session and generate the state
|
||||
wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||
if wLogin.Code != http.StatusOK {
|
||||
t.Fatalf("failed to get login URL: %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
|
||||
}
|
||||
}
|
||||
if anonymousCookie == nil {
|
||||
t.Fatal("session cookie not found after login URL generation")
|
||||
}
|
||||
|
||||
// 3. Trigger Callback (Login flow - new user)
|
||||
reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state)
|
||||
w := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
}, nil)
|
||||
}, []*http.Cookie{anonymousCookie})
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("callback failed with status %d, body: %s", w.Code, w.Body.String())
|
||||
@@ -674,20 +718,35 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
}
|
||||
|
||||
// 5. Test Callback (Login flow - existing user, username collision check)
|
||||
state2 := uuid.NewString()
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state2)), payloadValue, OAuthStateCacheKeyExpiration)
|
||||
|
||||
var state2 string
|
||||
// Callback with same username but different external ID (99999)
|
||||
httpMock2 := newMockOIDCClient(testIssuerURL, testClientID, state2, "99999", "test_oauth_user", "another@linux.do", "Another User")
|
||||
httpMock2 := newMockOIDCClient(testIssuerURL, testClientID, &state2, "99999", "test_oauth_user", "another@linux.do", "Another User")
|
||||
util.SetHTTPClient(httpMock2)
|
||||
|
||||
// Create another router for this mock client
|
||||
router2 := setupTestRouter(dbConn, mockRedis, httpMock2)
|
||||
|
||||
// Call login to get state2 and new anonymous session
|
||||
wLogin2 := performRequest(router2, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||
var loginUrlResp2 struct {
|
||||
Data OAuthAuthorizeResponse `json:"data"`
|
||||
}
|
||||
_ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp2)
|
||||
parsedURL2, _ := url.Parse(loginUrlResp2.Data.AuthorizeURL)
|
||||
state2 = parsedURL2.Query().Get("state")
|
||||
|
||||
var anonymousCookie2 *http.Cookie
|
||||
for _, cookie := range wLogin2.Result().Cookies() {
|
||||
if cookie.Name == config.Config.App.SessionCookieName {
|
||||
anonymousCookie2 = cookie
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state2)
|
||||
w3 := performRequest(router2, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
}, nil)
|
||||
}, []*http.Cookie{anonymousCookie2})
|
||||
|
||||
if w3.Code != http.StatusOK {
|
||||
t.Fatalf("callback for collision failed: %d, body: %s", w3.Code, w3.Body.String())
|
||||
@@ -712,21 +771,31 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
||||
dbConn.Where("key = ?", model.ConfigKeyRegistrationEnabled).Delete(&model.SystemConfig{})
|
||||
}()
|
||||
|
||||
state4 := uuid.NewString()
|
||||
payloadValue4, _ := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: testSourceName,
|
||||
Purpose: OAuthPurposeLogin,
|
||||
})
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state4)), payloadValue4, OAuthStateCacheKeyExpiration)
|
||||
|
||||
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
|
||||
var state4 string
|
||||
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, &state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
|
||||
util.SetHTTPClient(httpMock4)
|
||||
router4 := setupTestRouter(dbConn, mockRedis, httpMock4)
|
||||
|
||||
wLogin4 := performRequest(router4, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||
var loginUrlResp4 struct {
|
||||
Data OAuthAuthorizeResponse `json:"data"`
|
||||
}
|
||||
_ = json.Unmarshal(wLogin4.Body.Bytes(), &loginUrlResp4)
|
||||
parsedURL4, _ := url.Parse(loginUrlResp4.Data.AuthorizeURL)
|
||||
state4 = parsedURL4.Query().Get("state")
|
||||
|
||||
var anonymousCookie4 *http.Cookie
|
||||
for _, cookie := range wLogin4.Result().Cookies() {
|
||||
if cookie.Name == config.Config.App.SessionCookieName {
|
||||
anonymousCookie4 = cookie
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
reqBody4 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state4)
|
||||
w4 := performRequest(router4, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody4), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
}, nil)
|
||||
}, []*http.Cookie{anonymousCookie4})
|
||||
|
||||
if w4.Code != http.StatusOK {
|
||||
t.Fatalf("callback failed: %d, body: %s", w4.Code, w4.Body.String())
|
||||
@@ -773,30 +842,13 @@ func TestCallbackBind(t *testing.T) {
|
||||
OpenIDDiscoveryURL: "https://github.com",
|
||||
})
|
||||
|
||||
// Inject state in Redis with purpose=bind
|
||||
state := uuid.NewString()
|
||||
payloadValue, _ := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: "github",
|
||||
Purpose: OAuthPurposeBind,
|
||||
})
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration)
|
||||
|
||||
var state string
|
||||
// Mock OIDC discovery, JWKS, and Token exchange for custom source (GitHub)
|
||||
httpMock := newMockOIDCClient("https://github.com", "gh_client", state, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
|
||||
httpMock := newMockOIDCClient("https://github.com", "gh_client", &state, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
|
||||
util.SetHTTPClient(httpMock)
|
||||
router := setupTestRouter(dbConn, mockRedis, httpMock)
|
||||
|
||||
// Case 1: Bind attempt without session -> 401
|
||||
reqBody := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state)
|
||||
w1 := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
}, nil)
|
||||
|
||||
if w1.Code != http.StatusUnauthorized {
|
||||
t.Errorf("expected 401 for unauthenticated bind, got %d, body: %s", w1.Code, w1.Body.String())
|
||||
}
|
||||
|
||||
// Case 2: Bind success (authenticated)
|
||||
// Set up login helper
|
||||
router.GET("/test-helper/login-777", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(UserIDKey, uint64(777))
|
||||
@@ -813,10 +865,54 @@ func TestCallbackBind(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// Restore state key in redis
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration)
|
||||
// Generate OAuth authorize link (purpose=bind) to set state in Redis and Session
|
||||
wAuth := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie})
|
||||
if wAuth.Code != http.StatusOK {
|
||||
t.Fatalf("authorize failed: %d, body: %s", wAuth.Code, wAuth.Body.String())
|
||||
}
|
||||
// Extract the cookie from wAuth to get the session with the token!
|
||||
for _, cookie := range wAuth.Result().Cookies() {
|
||||
if cookie.Name == config.Config.App.SessionCookieName {
|
||||
activeCookie = cookie
|
||||
break
|
||||
}
|
||||
}
|
||||
var authResp struct {
|
||||
Data OAuthAuthorizeResponse `json:"data"`
|
||||
}
|
||||
_ = json.Unmarshal(wAuth.Body.Bytes(), &authResp)
|
||||
parsedURL, _ := url.Parse(authResp.Data.AuthorizeURL)
|
||||
state = parsedURL.Query().Get("state")
|
||||
|
||||
w2 := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
|
||||
// Case 1: Bind attempt without session -> 401
|
||||
reqBody := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state)
|
||||
w1 := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
}, nil)
|
||||
|
||||
if w1.Code != http.StatusUnauthorized {
|
||||
t.Errorf("expected 401 for unauthenticated bind, got %d, body: %s", w1.Code, w1.Body.String())
|
||||
}
|
||||
|
||||
// Case 2: Bind success (authenticated)
|
||||
// Re-run authorize since state is consumed/deleted during Callback attempt
|
||||
wAuth2 := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie})
|
||||
// Extract updated cookie from wAuth2
|
||||
for _, cookie := range wAuth2.Result().Cookies() {
|
||||
if cookie.Name == config.Config.App.SessionCookieName {
|
||||
activeCookie = cookie
|
||||
break
|
||||
}
|
||||
}
|
||||
var authResp2 struct {
|
||||
Data OAuthAuthorizeResponse `json:"data"`
|
||||
}
|
||||
_ = json.Unmarshal(wAuth2.Body.Bytes(), &authResp2)
|
||||
parsedURL2, _ := url.Parse(authResp2.Data.AuthorizeURL)
|
||||
state = parsedURL2.Query().Get("state")
|
||||
|
||||
reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state)
|
||||
w2 := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
}, []*http.Cookie{activeCookie})
|
||||
|
||||
@@ -865,14 +961,31 @@ func TestCallbackBind(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
state3 := uuid.NewString()
|
||||
mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state3)), payloadValue, OAuthStateCacheKeyExpiration)
|
||||
|
||||
var state3 string
|
||||
// Re-sign token for new state (since state serves as OIDC Nonce)
|
||||
httpMock3 := newMockOIDCClient("https://github.com", "gh_client", state3, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
|
||||
httpMock3 := newMockOIDCClient("https://github.com", "gh_client", &state3, "github_user_123", "github_tester", "tester@github.com", "GitHub Tester")
|
||||
util.SetHTTPClient(httpMock3)
|
||||
router3 := setupTestRouter(dbConn, mockRedis, httpMock3)
|
||||
|
||||
// Generate state3 and SessionHash using activeCookie2
|
||||
wAuth3 := performRequest(router3, http.MethodGet, "/api/v1/oauth/github/authorize?purpose=bind", nil, nil, []*http.Cookie{activeCookie2})
|
||||
if wAuth3.Code != http.StatusOK {
|
||||
t.Fatalf("authorize failed: %d, body: %s", wAuth3.Code, wAuth3.Body.String())
|
||||
}
|
||||
// Extract the cookie to get the updated session token
|
||||
for _, cookie := range wAuth3.Result().Cookies() {
|
||||
if cookie.Name == config.Config.App.SessionCookieName {
|
||||
activeCookie2 = cookie
|
||||
break
|
||||
}
|
||||
}
|
||||
var authResp3 struct {
|
||||
Data OAuthAuthorizeResponse `json:"data"`
|
||||
}
|
||||
_ = json.Unmarshal(wAuth3.Body.Bytes(), &authResp3)
|
||||
parsedURL3, _ := url.Parse(authResp3.Data.AuthorizeURL)
|
||||
state3 = parsedURL3.Query().Get("state")
|
||||
|
||||
reqBody3 := fmt.Sprintf(`{"state":"%s","code":"test_code"}`, state3)
|
||||
w3 := performRequest(router3, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody3), map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
@@ -881,6 +994,7 @@ func TestCallbackBind(t *testing.T) {
|
||||
if w3.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 for already bound account, got %d, body: %s", w3.Code, w3.Body.String())
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func TestExternalAccountsListAndDelete(t *testing.T) {
|
||||
|
||||
@@ -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()))
|
||||
|
||||
@@ -443,3 +443,115 @@ func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) {
|
||||
t.Errorf("expected login success, got error %q", successResp.ErrorMsg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// 1. Seed a user
|
||||
const (
|
||||
userID = uint64(500)
|
||||
username = "tokenuser"
|
||||
password = "tokenpassword123"
|
||||
)
|
||||
now := time.Now()
|
||||
userRecord := model.User{
|
||||
ID: userID,
|
||||
Username: username,
|
||||
Nickname: "Token User",
|
||||
Email: "tokenuser@example.com",
|
||||
IsActive: true,
|
||||
IsAdmin: true, // Make them an admin so we can test with is_admin=true token requests if needed
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := userRecord.SetEncryptedPassword(password); err != nil {
|
||||
t.Fatalf("set encrypted password failed: %v", err)
|
||||
}
|
||||
if err := dbConn.Create(&userRecord).Error; err != nil {
|
||||
t.Fatalf("create test user failed: %v", err)
|
||||
}
|
||||
|
||||
// Seed an active AccessToken for this user
|
||||
tokenStr, err := model.GenerateTokenString()
|
||||
if err != nil {
|
||||
t.Fatalf("generate token string failed: %v", err)
|
||||
}
|
||||
tokenHash := model.HashToken(tokenStr)
|
||||
tokenRecord := model.AccessToken{
|
||||
UserID: userID,
|
||||
Name: "Test Token",
|
||||
TokenHash: tokenHash,
|
||||
MaskedToken: model.MaskTokenString(tokenStr),
|
||||
IsAdmin: false,
|
||||
}
|
||||
if err := dbConn.Create(&tokenRecord).Error; err != nil {
|
||||
t.Fatalf("create test access token failed: %v", err)
|
||||
}
|
||||
|
||||
// 2. Set up router with access-token routes and oauth middlewares
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
store := cookie.NewStore([]byte("test_session_secret"))
|
||||
r.Use(sessions.Sessions("test_session_id", store))
|
||||
|
||||
apiV1Router := r.Group("/api/v1")
|
||||
userRouter := apiV1Router.Group("/user")
|
||||
tokenRouter := userRouter.Group("/access-tokens")
|
||||
tokenRouter.Use(oauth.LoginRequired(), oauth.DisallowTokenAuth())
|
||||
{
|
||||
tokenRouter.GET("", ListAccessTokens)
|
||||
tokenRouter.POST("", CreateAccessToken)
|
||||
tokenRouter.DELETE("/:id", DeleteAccessToken)
|
||||
tokenRouter.POST("/:id/rotate", RotateAccessToken)
|
||||
}
|
||||
|
||||
// 3. Test that accessing using an Access Token fails with 403 Forbidden
|
||||
req, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
req.Header.Set("X-Access-Token", tokenStr)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusForbidden {
|
||||
t.Errorf("expected status 403 Forbidden when accessing with Access Token, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("decode response failed: %v", err)
|
||||
}
|
||||
if resp.ErrorMsg != oauth.TokenAuthNotAllowed {
|
||||
t.Errorf("expected error message %q, got %q", oauth.TokenAuthNotAllowed, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// 4. Test that accessing using a Session succeeds
|
||||
sessionCookieStore := cookie.NewStore([]byte("test_session_secret"))
|
||||
rSession := gin.New()
|
||||
rSession.Use(sessions.Sessions("test_session_id", sessionCookieStore))
|
||||
rSession.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), oauth.DisallowTokenAuth(), ListAccessTokens)
|
||||
|
||||
// We can login/register or just mock the session handler to set user ID
|
||||
rSession.GET("/mock-login", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, userID)
|
||||
session.Set(oauth.UserNameKey, username)
|
||||
_ = session.Save()
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
wMock := httptest.NewRecorder()
|
||||
reqMock, _ := http.NewRequest(http.MethodGet, "/mock-login", nil)
|
||||
rSession.ServeHTTP(wMock, reqMock)
|
||||
cookieVal := wMock.Header().Get("Set-Cookie")
|
||||
|
||||
reqSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil)
|
||||
reqSession.Header.Set("Cookie", cookieVal)
|
||||
wSession := httptest.NewRecorder()
|
||||
rSession.ServeHTTP(wSession, reqSession)
|
||||
|
||||
if wSession.Code != http.StatusOK {
|
||||
t.Errorf("expected status 200 OK when accessing with Session, got %d. Body: %s", wSession.Code, wSession.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -188,7 +188,7 @@ func registerRoutes(r *gin.Engine) {
|
||||
|
||||
// Access Token
|
||||
tokenRouter := userRouter.Group("/access-tokens")
|
||||
tokenRouter.Use(oauth.LoginRequired())
|
||||
tokenRouter.Use(oauth.LoginRequired(), oauth.DisallowTokenAuth())
|
||||
{
|
||||
tokenRouter.GET("", user.ListAccessTokens)
|
||||
tokenRouter.POST("", user.CreateAccessToken)
|
||||
|
||||
Reference in New Issue
Block a user