Add user and security tests
This commit is contained in:
parent
2610036082
commit
ef07622e7e
3 changed files with 1063 additions and 8 deletions
624
openapi/oauth/security_test.go
Normal file
624
openapi/oauth/security_test.go
Normal file
|
|
@ -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)
|
||||
})
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
439
openapi/oauth/user_test.go
Normal file
439
openapi/oauth/user_test.go
Normal file
|
|
@ -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"])
|
||||
}
|
||||
})
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue