From ef07622e7e3b2448a8f64eccae275fd5965f87b6 Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 18 Jul 2025 15:48:50 +0800 Subject: [PATCH] Add user and security tests --- openapi/oauth/security_test.go | 624 +++++++++++++++++++++++++++++++++ openapi/oauth/user.go | 8 - openapi/oauth/user_test.go | 439 +++++++++++++++++++++++ 3 files changed, 1063 insertions(+), 8 deletions(-) create mode 100644 openapi/oauth/security_test.go create mode 100644 openapi/oauth/user_test.go diff --git a/openapi/oauth/security_test.go b/openapi/oauth/security_test.go new file mode 100644 index 00000000..cc966b20 --- /dev/null +++ b/openapi/oauth/security_test.go @@ -0,0 +1,624 @@ +package oauth + +import ( + "context" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/openapi/oauth/types" +) + +// ============================================================================= +// PKCE Code Challenge Tests +// ============================================================================= + +func TestGenerateCodeChallenge(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + codeVerifier := "test_code_verifier_123456789" + + t.Run("S256 method", func(t *testing.T) { + challenge, err := service.GenerateCodeChallenge(ctx, codeVerifier, "S256") + assert.NoError(t, err) + assert.NotEmpty(t, challenge) + assert.NotEqual(t, codeVerifier, challenge) + + // Should be base64 URL encoded + assert.NotContains(t, challenge, "=") + assert.NotContains(t, challenge, "+") + assert.NotContains(t, challenge, "/") + }) + + t.Run("plain method", func(t *testing.T) { + challenge, err := service.GenerateCodeChallenge(ctx, codeVerifier, "plain") + assert.NoError(t, err) + assert.Equal(t, codeVerifier, challenge) + }) + + t.Run("unsupported method", func(t *testing.T) { + challenge, err := service.GenerateCodeChallenge(ctx, codeVerifier, "unsupported") + assert.Error(t, err) + assert.Empty(t, challenge) + assert.Contains(t, err.Error(), "unsupported code challenge method") + }) + + t.Run("consistency check", func(t *testing.T) { + // Same verifier should generate same challenge + challenge1, err := service.GenerateCodeChallenge(ctx, codeVerifier, "S256") + assert.NoError(t, err) + + challenge2, err := service.GenerateCodeChallenge(ctx, codeVerifier, "S256") + assert.NoError(t, err) + + assert.Equal(t, challenge1, challenge2) + }) +} + +func TestValidateCodeChallenge(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + codeVerifier := "test_code_verifier_123456789" + + t.Run("valid S256 challenge", func(t *testing.T) { + challenge, err := service.GenerateCodeChallenge(ctx, codeVerifier, "S256") + require.NoError(t, err) + + err = service.ValidateCodeChallenge(ctx, codeVerifier, challenge, "S256") + assert.NoError(t, err) + }) + + t.Run("valid plain challenge", func(t *testing.T) { + challenge, err := service.GenerateCodeChallenge(ctx, codeVerifier, "plain") + require.NoError(t, err) + + err = service.ValidateCodeChallenge(ctx, codeVerifier, challenge, "plain") + assert.NoError(t, err) + }) + + t.Run("invalid S256 challenge", func(t *testing.T) { + err := service.ValidateCodeChallenge(ctx, codeVerifier, "invalid_challenge", "S256") + assert.Error(t, err) + assert.Contains(t, err.Error(), "code challenge verification failed") + }) + + t.Run("invalid plain challenge", func(t *testing.T) { + err := service.ValidateCodeChallenge(ctx, "wrong_verifier", codeVerifier, "plain") + assert.Error(t, err) + assert.Contains(t, err.Error(), "code challenge verification failed") + }) + + t.Run("unsupported method", func(t *testing.T) { + err := service.ValidateCodeChallenge(ctx, codeVerifier, "challenge", "unsupported") + assert.Error(t, err) + assert.Contains(t, err.Error(), "unsupported code challenge method") + }) +} + +// ============================================================================= +// State Parameter Tests +// ============================================================================= + +func TestGenerateStateParameter(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + clientID := testClients[0].ClientID + + t.Run("generate valid state parameter", func(t *testing.T) { + stateParam, err := service.GenerateStateParameter(ctx, clientID) + assert.NoError(t, err) + assert.NotNil(t, stateParam) + assert.NotEmpty(t, stateParam.Value) + assert.Equal(t, clientID, stateParam.ClientID) + assert.True(t, stateParam.ExpiresAt.After(time.Now())) + }) + + t.Run("generate unique state parameters", func(t *testing.T) { + state1, err := service.GenerateStateParameter(ctx, clientID) + assert.NoError(t, err) + + state2, err := service.GenerateStateParameter(ctx, clientID) + assert.NoError(t, err) + + assert.NotEqual(t, state1.Value, state2.Value) + }) + + t.Run("state parameter format", func(t *testing.T) { + stateParam, err := service.GenerateStateParameter(ctx, clientID) + assert.NoError(t, err) + + // Should be base64 URL encoded + assert.NotContains(t, stateParam.Value, "=") + assert.NotContains(t, stateParam.Value, "+") + assert.NotContains(t, stateParam.Value, "/") + assert.True(t, len(stateParam.Value) > 0) + }) + + t.Run("empty client ID", func(t *testing.T) { + stateParam, err := service.GenerateStateParameter(ctx, "") + assert.NoError(t, err) + assert.NotNil(t, stateParam) + assert.Empty(t, stateParam.ClientID) + }) +} + +func TestValidateStateParameter(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + clientID := testClients[0].ClientID + + t.Run("validate valid state parameter", func(t *testing.T) { + stateParam, err := service.GenerateStateParameter(ctx, clientID) + require.NoError(t, err) + + result, err := service.ValidateStateParameter(ctx, stateParam.Value, clientID) + assert.NoError(t, err) + assert.True(t, result.Valid) + assert.Empty(t, result.Errors) + }) + + t.Run("validate non-existent state parameter", func(t *testing.T) { + result, err := service.ValidateStateParameter(ctx, "non_existent_state", clientID) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "State parameter not found") + }) + + t.Run("validate state parameter with wrong client", func(t *testing.T) { + stateParam, err := service.GenerateStateParameter(ctx, clientID) + require.NoError(t, err) + + wrongClientID := testClients[1].ClientID + result, err := service.ValidateStateParameter(ctx, stateParam.Value, wrongClientID) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "State parameter not found") + }) + + t.Run("validate expired state parameter", func(t *testing.T) { + // Create state parameter with very short lifetime + originalConfig := service.config.Security.StateParameterLifetime + service.config.Security.StateParameterLifetime = 1 * time.Nanosecond + + stateParam, err := service.GenerateStateParameter(ctx, clientID) + require.NoError(t, err) + + // Restore original config + service.config.Security.StateParameterLifetime = originalConfig + + // Wait for expiration + time.Sleep(2 * time.Millisecond) + + result, err := service.ValidateStateParameter(ctx, stateParam.Value, clientID) + assert.NoError(t, err) + assert.False(t, result.Valid) + // Note: The implementation might store the state parameter in cache, + // so it could return different error messages depending on cache state + assert.NotEmpty(t, result.Errors) + }) + + t.Run("validate empty state parameter", func(t *testing.T) { + result, err := service.ValidateStateParameter(ctx, "", clientID) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "State parameter not found") + }) +} + +// ============================================================================= +// Redirect URI Validation Tests +// ============================================================================= + +func TestValidateRedirectURI(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + + t.Run("valid redirect URI", func(t *testing.T) { + redirectURI := "https://example.com/callback" + registeredURIs := []string{ + "https://example.com/callback", + "https://example.com/other", + } + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.True(t, result.Valid) + assert.Empty(t, result.Errors) + }) + + t.Run("invalid redirect URI", func(t *testing.T) { + redirectURI := "https://malicious.com/callback" + registeredURIs := []string{ + "https://example.com/callback", + "https://example.com/other", + } + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "Redirect URI not found in registered URIs") + }) + + t.Run("no registered URIs", func(t *testing.T) { + redirectURI := "https://example.com/callback" + registeredURIs := []string{} + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "No registered URIs provided") + }) + + t.Run("nil registered URIs", func(t *testing.T) { + redirectURI := "https://example.com/callback" + var registeredURIs []string + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "No registered URIs provided") + }) + + t.Run("empty redirect URI", func(t *testing.T) { + redirectURI := "" + registeredURIs := []string{"https://example.com/callback"} + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "Redirect URI not found in registered URIs") + }) + + t.Run("exact match required", func(t *testing.T) { + redirectURI := "https://example.com/callback/extra" + registeredURIs := []string{"https://example.com/callback"} + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.Contains(t, result.Errors, "Redirect URI not found in registered URIs") + }) +} + +func TestValidateRedirectURIForClient(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + clientID := testClients[0].ClientID + validRedirectURI := testClients[0].RedirectURIs[0] + + t.Run("valid redirect URI for client", func(t *testing.T) { + result, err := service.ValidateRedirectURIForClient(ctx, clientID, validRedirectURI) + assert.NoError(t, err) + assert.True(t, result.Valid) + assert.Empty(t, result.Errors) + }) + + t.Run("invalid redirect URI for client", func(t *testing.T) { + invalidRedirectURI := "https://malicious.com/callback" + result, err := service.ValidateRedirectURIForClient(ctx, clientID, invalidRedirectURI) + assert.NoError(t, err) + assert.False(t, result.Valid) + assert.NotEmpty(t, result.Errors) + }) + + t.Run("non-existent client", func(t *testing.T) { + result, err := service.ValidateRedirectURIForClient(ctx, "non-existent-client", validRedirectURI) + assert.Error(t, err) + assert.Nil(t, result) + }) +} + +// ============================================================================= +// Pushed Authorization Request Tests +// ============================================================================= + +func TestPushAuthorizationRequest(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + clientID := testClients[0].ClientID + redirectURI := testClients[0].RedirectURIs[0] + + t.Run("successful pushed authorization request", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: clientID, + RedirectURI: redirectURI, + ResponseType: types.ResponseTypeCode, + Scope: "openid profile", + State: "test_state", + } + + response, err := service.PushAuthorizationRequest(ctx, request) + assert.NoError(t, err) + assert.NotNil(t, response) + assert.NotEmpty(t, response.RequestURI) + assert.True(t, strings.HasPrefix(response.RequestURI, "urn:ietf:params:oauth:request_uri:")) + assert.Equal(t, 600, response.ExpiresIn) + }) + + t.Run("invalid client", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: "invalid-client", + RedirectURI: redirectURI, + ResponseType: types.ResponseTypeCode, + Scope: "openid profile", + State: "test_state", + } + + response, err := service.PushAuthorizationRequest(ctx, request) + assert.Error(t, err) + assert.Nil(t, response) + + errorResp, ok := err.(*types.ErrorResponse) + assert.True(t, ok) + assert.Equal(t, types.ErrorInvalidClient, errorResp.Code) + }) + + t.Run("invalid redirect URI", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: clientID, + RedirectURI: "https://malicious.com/callback", + ResponseType: types.ResponseTypeCode, + Scope: "openid profile", + State: "test_state", + } + + response, err := service.PushAuthorizationRequest(ctx, request) + assert.Error(t, err) + assert.Nil(t, response) + + errorResp, ok := err.(*types.ErrorResponse) + assert.True(t, ok) + assert.Equal(t, types.ErrorInvalidRequest, errorResp.Code) + }) + + t.Run("invalid scope", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: clientID, + RedirectURI: redirectURI, + ResponseType: types.ResponseTypeCode, + Scope: "invalid_scope", + State: "test_state", + } + + response, err := service.PushAuthorizationRequest(ctx, request) + assert.Error(t, err) + assert.Nil(t, response) + + errorResp, ok := err.(*types.ErrorResponse) + assert.True(t, ok) + assert.Equal(t, types.ErrorInvalidScope, errorResp.Code) + }) + + t.Run("request without scope", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: clientID, + RedirectURI: redirectURI, + ResponseType: types.ResponseTypeCode, + State: "test_state", + } + + response, err := service.PushAuthorizationRequest(ctx, request) + assert.NoError(t, err) + assert.NotNil(t, response) + assert.NotEmpty(t, response.RequestURI) + }) + + t.Run("request URI uniqueness", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: clientID, + RedirectURI: redirectURI, + ResponseType: types.ResponseTypeCode, + Scope: "openid profile", + State: "test_state", + } + + response1, err := service.PushAuthorizationRequest(ctx, request) + assert.NoError(t, err) + + response2, err := service.PushAuthorizationRequest(ctx, request) + assert.NoError(t, err) + + assert.NotEqual(t, response1.RequestURI, response2.RequestURI) + }) +} + +// ============================================================================= +// Helper Method Tests +// ============================================================================= + +func TestSecurityHelperMethods(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + clientID := testClients[0].ClientID + + t.Run("state parameter key generation", func(t *testing.T) { + state := "test_state" + key := service.stateParameterKey(clientID, state) + + assert.NotEmpty(t, key) + assert.Contains(t, key, "oauth:state") + assert.Contains(t, key, clientID) + assert.Contains(t, key, state) + }) + + t.Run("pushed auth request key generation", func(t *testing.T) { + requestURI := "test_request_uri" + key := service.pushedAuthRequestKey(requestURI) + + assert.NotEmpty(t, key) + assert.Contains(t, key, "oauth:par") + assert.Contains(t, key, requestURI) + }) + + t.Run("request URI generation", func(t *testing.T) { + requestURI := service.generateRequestURI() + + assert.NotEmpty(t, requestURI) + assert.True(t, strings.HasPrefix(requestURI, "urn:ietf:params:oauth:request_uri:")) + + // Should be base64 URL encoded + parts := strings.Split(requestURI, ":") + assert.True(t, len(parts) >= 4) + encodedPart := parts[len(parts)-1] + assert.NotContains(t, encodedPart, "=") + assert.NotContains(t, encodedPart, "+") + assert.NotContains(t, encodedPart, "/") + }) + + t.Run("request URI uniqueness", func(t *testing.T) { + uri1 := service.generateRequestURI() + uri2 := service.generateRequestURI() + + assert.NotEqual(t, uri1, uri2) + }) +} + +// ============================================================================= +// Integration Tests +// ============================================================================= + +func TestSecurityIntegration(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + clientID := testClients[0].ClientID + + t.Run("complete PKCE flow", func(t *testing.T) { + codeVerifier := "test_code_verifier_123456789" + + // Generate code challenge + challenge, err := service.GenerateCodeChallenge(ctx, codeVerifier, "S256") + assert.NoError(t, err) + + // Validate code challenge + err = service.ValidateCodeChallenge(ctx, codeVerifier, challenge, "S256") + assert.NoError(t, err) + + // Test with wrong verifier + err = service.ValidateCodeChallenge(ctx, "wrong_verifier", challenge, "S256") + assert.Error(t, err) + }) + + t.Run("complete state parameter flow", func(t *testing.T) { + // Generate state parameter + stateParam, err := service.GenerateStateParameter(ctx, clientID) + assert.NoError(t, err) + + // Validate state parameter + result, err := service.ValidateStateParameter(ctx, stateParam.Value, clientID) + assert.NoError(t, err) + assert.True(t, result.Valid) + + // Test with wrong client + wrongClientID := testClients[1].ClientID + result, err = service.ValidateStateParameter(ctx, stateParam.Value, wrongClientID) + assert.NoError(t, err) + assert.False(t, result.Valid) + }) + + t.Run("complete pushed authorization flow", func(t *testing.T) { + // Create pushed authorization request + request := &types.PushedAuthorizationRequest{ + ClientID: clientID, + RedirectURI: testClients[0].RedirectURIs[0], + ResponseType: types.ResponseTypeCode, + Scope: "openid profile", + State: "test_state", + } + + // Push authorization request + response, err := service.PushAuthorizationRequest(ctx, request) + assert.NoError(t, err) + assert.NotNil(t, response) + assert.NotEmpty(t, response.RequestURI) + + // Request URI should be stored and retrievable + key := service.pushedAuthRequestKey(response.RequestURI) + data, ok := service.store.Get(key) + assert.True(t, ok) + assert.NotNil(t, data) + }) +} + +// ============================================================================= +// Edge Cases and Error Handling +// ============================================================================= + +func TestSecurityEdgeCases(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + + t.Run("PKCE with empty code verifier", func(t *testing.T) { + challenge, err := service.GenerateCodeChallenge(ctx, "", "S256") + assert.NoError(t, err) + assert.NotEmpty(t, challenge) + + err = service.ValidateCodeChallenge(ctx, "", challenge, "S256") + assert.NoError(t, err) + }) + + t.Run("PKCE with very long code verifier", func(t *testing.T) { + longVerifier := strings.Repeat("a", 1000) + challenge, err := service.GenerateCodeChallenge(ctx, longVerifier, "S256") + assert.NoError(t, err) + assert.NotEmpty(t, challenge) + + err = service.ValidateCodeChallenge(ctx, longVerifier, challenge, "S256") + assert.NoError(t, err) + }) + + t.Run("state parameter with special characters", func(t *testing.T) { + clientID := "client-with-special-chars-!@#$%" + stateParam, err := service.GenerateStateParameter(ctx, clientID) + assert.NoError(t, err) + assert.NotNil(t, stateParam) + + result, err := service.ValidateStateParameter(ctx, stateParam.Value, clientID) + assert.NoError(t, err) + assert.True(t, result.Valid) + }) + + t.Run("redirect URI with query parameters", func(t *testing.T) { + redirectURI := "https://example.com/callback?param=value" + registeredURIs := []string{redirectURI} + + result, err := service.ValidateRedirectURI(ctx, redirectURI, registeredURIs) + assert.NoError(t, err) + assert.True(t, result.Valid) + }) + + t.Run("pushed authorization request with empty fields", func(t *testing.T) { + request := &types.PushedAuthorizationRequest{ + ClientID: testClients[0].ClientID, + RedirectURI: testClients[0].RedirectURIs[0], + ResponseType: "", + Scope: "", + State: "", + } + + response, err := service.PushAuthorizationRequest(ctx, request) + assert.NoError(t, err) + assert.NotNil(t, response) + assert.NotEmpty(t, response.RequestURI) + }) +} diff --git a/openapi/oauth/user.go b/openapi/oauth/user.go index f0f247af..0cc06a48 100644 --- a/openapi/oauth/user.go +++ b/openapi/oauth/user.go @@ -8,11 +8,3 @@ import ( func (s *Service) UserInfo(ctx context.Context, accessToken string) (interface{}, error) { return s.userProvider.GetUserByAccessToken(ctx, accessToken) } - -// Additional user-related helper methods can be added here as needed -// For example: -// - User profile management -// - User consent handling -// - User authentication verification -// - User scope validation -// etc. diff --git a/openapi/oauth/user_test.go b/openapi/oauth/user_test.go new file mode 100644 index 00000000..299f1145 --- /dev/null +++ b/openapi/oauth/user_test.go @@ -0,0 +1,439 @@ +package oauth + +import ( + "context" + "fmt" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// ============================================================================= +// UserInfo Tests +// ============================================================================= + +func TestUserInfo(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + userProvider := service.GetUserProvider() + + t.Run("get user info with valid access token", func(t *testing.T) { + // Create a valid access token for the first test user + testUser := testUsers[0] + accessToken := "valid_access_token_123" + + // Store token data in user provider + tokenData := map[string]interface{}{ + "token": accessToken, + "user_id": testUser.ID, + "subject": testUser.Subject, + "username": testUser.Username, + "email": testUser.Email, + "first_name": testUser.FirstName, + "last_name": testUser.LastName, + "full_name": testUser.FullName, + "scopes": testUser.Scopes, + "status": testUser.Status, + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + // Store the token data + err := userProvider.StoreToken(accessToken, tokenData, time.Hour) + require.NoError(t, err) + + // Get user info using the access token + userInfo, err := service.UserInfo(ctx, accessToken) + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify the user info contains expected data + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, testUser.Subject, userInfoMap["subject"]) + assert.Equal(t, testUser.Username, userInfoMap["username"]) + assert.Equal(t, testUser.Email, userInfoMap["email"]) + } + }) + + t.Run("get user info with invalid access token", func(t *testing.T) { + invalidToken := "invalid_access_token_xyz" + + userInfo, err := service.UserInfo(ctx, invalidToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("get user info with non-existent access token", func(t *testing.T) { + nonExistentToken := "non_existent_token_abc" + + userInfo, err := service.UserInfo(ctx, nonExistentToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("get user info with empty access token", func(t *testing.T) { + emptyToken := "" + + userInfo, err := service.UserInfo(ctx, emptyToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("get user info with expired access token", func(t *testing.T) { + testUser := testUsers[1] + expiredToken := "expired_access_token_456" + + // Store expired token data + tokenData := map[string]interface{}{ + "token": expiredToken, + "user_id": testUser.ID, + "subject": testUser.Subject, + "username": testUser.Username, + "email": testUser.Email, + "scopes": testUser.Scopes, + "status": testUser.Status, + "exp": time.Now().Add(-time.Hour).Unix(), // Expired 1 hour ago + "iat": time.Now().Add(-2 * time.Hour).Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(expiredToken, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, expiredToken) + // UserInfo method returns user data regardless of token expiry + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify user info contains expected data + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, testUser.Subject, userInfoMap["subject"]) + assert.Equal(t, testUser.Username, userInfoMap["username"]) + } + }) + + t.Run("get user info with inactive user", func(t *testing.T) { + // Use the inactive test user + inactiveUser := testUsers[4] // inactive.user + inactiveToken := "inactive_user_token_789" + + tokenData := map[string]interface{}{ + "token": inactiveToken, + "user_id": inactiveUser.ID, + "subject": inactiveUser.Subject, + "username": inactiveUser.Username, + "email": inactiveUser.Email, + "scopes": inactiveUser.Scopes, + "status": inactiveUser.Status, // inactive + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(inactiveToken, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, inactiveToken) + // UserInfo method returns user data regardless of user status + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify user info contains expected data + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, inactiveUser.Subject, userInfoMap["subject"]) + assert.Equal(t, inactiveUser.Username, userInfoMap["username"]) + assert.Equal(t, inactiveUser.Status, userInfoMap["status"]) + } + }) + + t.Run("get user info with limited scope user", func(t *testing.T) { + // Use the limited scope test user + limitedUser := testUsers[5] // limited.user + limitedToken := "limited_scope_token_101" + + tokenData := map[string]interface{}{ + "token": limitedToken, + "user_id": limitedUser.ID, + "subject": limitedUser.Subject, + "username": limitedUser.Username, + "email": limitedUser.Email, + "scopes": limitedUser.Scopes, // Only openid + "status": limitedUser.Status, + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(limitedToken, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, limitedToken) + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify limited user info + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, limitedUser.Subject, userInfoMap["subject"]) + assert.Equal(t, limitedUser.Username, userInfoMap["username"]) + // Should only have basic scopes + if scopes, ok := userInfoMap["scopes"].([]string); ok { + assert.Contains(t, scopes, "openid") + assert.Len(t, scopes, 1) + } + } + }) + + t.Run("get user info with admin user", func(t *testing.T) { + // Use the admin test user + adminUser := testUsers[0] // admin + adminToken := "admin_token_202" + + tokenData := map[string]interface{}{ + "token": adminToken, + "user_id": adminUser.ID, + "subject": adminUser.Subject, + "username": adminUser.Username, + "email": adminUser.Email, + "first_name": adminUser.FirstName, + "last_name": adminUser.LastName, + "full_name": adminUser.FullName, + "scopes": adminUser.Scopes, + "status": adminUser.Status, + "email_verified": adminUser.EmailVerified, + "mobile_verified": adminUser.MobileVerified, + "two_factor_enabled": adminUser.TwoFactorEnabled, + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(adminToken, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, adminToken) + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify admin user info + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, adminUser.Subject, userInfoMap["subject"]) + assert.Equal(t, adminUser.Username, userInfoMap["username"]) + assert.Equal(t, adminUser.Email, userInfoMap["email"]) + assert.True(t, userInfoMap["email_verified"].(bool)) + assert.True(t, userInfoMap["two_factor_enabled"].(bool)) + + // Should have admin scopes + if scopes, ok := userInfoMap["scopes"].([]string); ok { + assert.Contains(t, scopes, "admin") + assert.Contains(t, scopes, "openid") + assert.Contains(t, scopes, "profile") + assert.Contains(t, scopes, "email") + } + } + }) +} + +// ============================================================================= +// Integration Tests +// ============================================================================= + +func TestUserInfoIntegration(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + userProvider := service.GetUserProvider() + + t.Run("complete user info flow", func(t *testing.T) { + // Use different test users for comprehensive testing + testCases := []struct { + name string + user *TestUser + tokenSuffix string + }{ + {"regular_user", testUsers[1], "regular"}, + {"verified_user", testUsers[2], "verified"}, + {"secure_user", testUsers[6], "secure"}, + {"api_user", testUsers[7], "api"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + token := "integration_token_" + tc.tokenSuffix + + tokenData := map[string]interface{}{ + "token": token, + "user_id": tc.user.ID, + "subject": tc.user.Subject, + "username": tc.user.Username, + "email": tc.user.Email, + "scopes": tc.user.Scopes, + "status": tc.user.Status, + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(token, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, token) + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify basic user info structure + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, tc.user.Subject, userInfoMap["subject"]) + assert.Equal(t, tc.user.Username, userInfoMap["username"]) + assert.Equal(t, tc.user.Email, userInfoMap["email"]) + assert.Equal(t, tc.user.Status, userInfoMap["status"]) + } + }) + } + }) + + t.Run("concurrent user info requests", func(t *testing.T) { + // Test concurrent access to user info + const numRequests = 10 + + // Create tokens for concurrent testing + tokens := make([]string, numRequests) + for i := 0; i < numRequests; i++ { + tokens[i] = fmt.Sprintf("concurrent_token_%d", i) + testUser := testUsers[i%len(testUsers)] + + tokenData := map[string]interface{}{ + "token": tokens[i], + "user_id": testUser.ID, + "subject": testUser.Subject, + "username": testUser.Username, + "email": testUser.Email, + "scopes": testUser.Scopes, + "status": testUser.Status, + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(tokens[i], tokenData, time.Hour) + require.NoError(t, err) + } + + // Make concurrent requests + results := make(chan error, numRequests) + for i := 0; i < numRequests; i++ { + go func(token string) { + userInfo, err := service.UserInfo(ctx, token) + if err != nil { + results <- err + return + } + if userInfo == nil { + results <- fmt.Errorf("user info is nil") + return + } + results <- nil + }(tokens[i]) + } + + // Collect results + for i := 0; i < numRequests; i++ { + err := <-results + assert.NoError(t, err) + } + }) +} + +// ============================================================================= +// Edge Cases and Error Handling +// ============================================================================= + +func TestUserInfoEdgeCases(t *testing.T) { + service, _, _, cleanup := setupOAuthTestEnvironment(t) + defer cleanup() + + ctx := context.Background() + userProvider := service.GetUserProvider() + + t.Run("malformed token data", func(t *testing.T) { + malformedToken := "malformed_token_data" + + // Store malformed token data + tokenData := map[string]interface{}{ + "token": malformedToken, + "user_id": "invalid_user_id", + "subject": nil, // Invalid subject + "username": "", // Empty username + "exp": "not_a_number", // Invalid expiration + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(malformedToken, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, malformedToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("very long access token", func(t *testing.T) { + // Create a very long token + longToken := "very_long_token_" + strings.Repeat("a", 1000) + + userInfo, err := service.UserInfo(ctx, longToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("special characters in token", func(t *testing.T) { + specialToken := "special_token_!@#$%^&*()_+{}[]|\\:;\"'<>?,./`~" + + userInfo, err := service.UserInfo(ctx, specialToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("token with only whitespace", func(t *testing.T) { + whitespaceToken := " \t\n\r " + + userInfo, err := service.UserInfo(ctx, whitespaceToken) + assert.Error(t, err) + assert.Nil(t, userInfo) + }) + + t.Run("token with minimal valid data", func(t *testing.T) { + minimalToken := "minimal_token_999" + testUser := testUsers[9] // test.user + + // Store minimal token data + tokenData := map[string]interface{}{ + "token": minimalToken, + "user_id": testUser.ID, + "subject": testUser.Subject, + "username": testUser.Username, + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Unix(), + "token_type": "Bearer", + } + + err := userProvider.StoreToken(minimalToken, tokenData, time.Hour) + require.NoError(t, err) + + userInfo, err := service.UserInfo(ctx, minimalToken) + assert.NoError(t, err) + assert.NotNil(t, userInfo) + + // Verify minimal user info + if userInfoMap, ok := userInfo.(map[string]interface{}); ok { + assert.Equal(t, testUser.Subject, userInfoMap["subject"]) + assert.Equal(t, testUser.Username, userInfoMap["username"]) + } + }) +}