diff --git a/internal/apps/oauth/constants.go b/internal/apps/oauth/constants.go index 8abe0261..3fde9096 100644 --- a/internal/apps/oauth/constants.go +++ b/internal/apps/oauth/constants.go @@ -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 } + diff --git a/internal/apps/oauth/errs.go b/internal/apps/oauth/errs.go index 09da5a2c..b4dd32f5 100644 --- a/internal/apps/oauth/errs.go +++ b/internal/apps/oauth/errs.go @@ -18,4 +18,5 @@ const ( UsernameFromSourceFailed = "无法从认证源获取用户名" AuthSourceDisabled = "认证源未启用" InvalidExternalAccountBindingID = "绑定记录 ID 无效" + TokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials ) diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 8c0bcd42..aecc04cf 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -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() + } +} + diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index 5dcebee3..1fd99ab1 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -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) { diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index d5bc4698..c08c282f 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -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())) diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index 8784a1e3..3bc2055a 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -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()) + } +} + diff --git a/internal/router/router.go b/internal/router/router.go index c5d6bd9b..5d8fcfc9 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -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)