76e4f5790f
Health Endpoint * health
410 lines
13 KiB
Go
410 lines
13 KiB
Go
package cognitoauth
|
|
|
|
import (
|
|
"encoding/base64"
|
|
"fmt"
|
|
"log/slog"
|
|
"os"
|
|
"reflect"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
// mockConfigProvider implements the ConfigProvider interface for testing purposes.
|
|
// It provides predefined values for configuration methods and no-op implementations
|
|
// for setters, allowing tests to run without requiring actual AWS Cognito credentials.
|
|
type mockConfigProvider struct {
|
|
region string
|
|
userPoolID string
|
|
}
|
|
|
|
func (m *mockConfigProvider) GetAuthRegion() string { return m.region }
|
|
func (m *mockConfigProvider) SetAuthRegion(s string) { m.region = s }
|
|
func (m *mockConfigProvider) GetAuthUserPoolID() string { return m.userPoolID }
|
|
func (m *mockConfigProvider) SetAuthUserPoolID(s string) { m.userPoolID = s }
|
|
func (m *mockConfigProvider) GetAuthClientSecret() string { return "mock-secret" }
|
|
func (m *mockConfigProvider) SetAuthClientSecret(s string) {}
|
|
func (m *mockConfigProvider) GetAuthDomain() string { return "mock-domain" }
|
|
func (m *mockConfigProvider) SetAuthDomain(s string) {}
|
|
func (m *mockConfigProvider) GetAuthClientID() string { return "mock-client-id" }
|
|
func (m *mockConfigProvider) SetAuthClientID(s string) {}
|
|
func (m *mockConfigProvider) GetAuthRedirectURI() string { return "http://localhost/callback" }
|
|
func (m *mockConfigProvider) SetAuthRedirectURI(s string) {}
|
|
func (m *mockConfigProvider) GetAuthTokenURL() string { return "http://localhost/token" }
|
|
func (m *mockConfigProvider) SetAuthTokenURL(s string) {}
|
|
func (m *mockConfigProvider) GetAuthJwksURL() string { return "http://localhost/jwks" }
|
|
func (m *mockConfigProvider) SetAuthJwksURL(s string) {}
|
|
func (m *mockConfigProvider) GetAuthURL() string { return "http://localhost/auth" }
|
|
func (m *mockConfigProvider) SetAuthURL(s string) {}
|
|
func (m *mockConfigProvider) GetAuthLoginPath() string { return "/login" }
|
|
func (m *mockConfigProvider) SetAuthLoginPath(s string) {}
|
|
func (m *mockConfigProvider) GetAuthCallbackPath() string { return "/callback" }
|
|
func (m *mockConfigProvider) SetAuthCallbackPath(s string) {}
|
|
func (m *mockConfigProvider) GetAuthHomePath() string { return "/home" }
|
|
func (m *mockConfigProvider) SetAuthHomePath(s string) {}
|
|
func (m *mockConfigProvider) GetAuthLogoutPath() string { return "/logout" }
|
|
func (m *mockConfigProvider) SetAuthLogoutPath(s string) {}
|
|
func (m *mockConfigProvider) GetHealthPath() string { return "/health" }
|
|
func (m *mockConfigProvider) SetHealthPath(s string) {}
|
|
func (m *mockConfigProvider) GetAuthLogger() *slog.Logger { return slog.Default() }
|
|
func (m *mockConfigProvider) SetAuthLogger(*slog.Logger) {}
|
|
func (m *mockConfigProvider) GetAuthRoutePermissions() map[string][]string {
|
|
return map[string][]string{"/test": {"admin"}}
|
|
}
|
|
func (m *mockConfigProvider) SetAuthRoutePermissions(map[string][]string) {}
|
|
func (m *mockConfigProvider) InitializeAuthConfig(string, *slog.Logger) error {
|
|
return nil
|
|
}
|
|
|
|
// Helper function to generate test JWT and JWK
|
|
|
|
// TestGetUserGroups verifies the GetUserGroups function correctly extracts user group
|
|
// information from JWT claims. It tests:
|
|
// - Extraction from "cognito:groups" claim as an array
|
|
// - Extraction from "groups" claim as an array
|
|
// - Extraction from "custom:groups" claim as a string
|
|
// - Error handling when no group information is present
|
|
// - Handling of mixed data types in group arrays
|
|
func TestGetUserGroups(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
claims map[string]interface{}
|
|
expected []string
|
|
expectError bool
|
|
}{
|
|
{
|
|
name: "cognito:groups as string array",
|
|
claims: map[string]interface{}{
|
|
"cognito:groups": []interface{}{"admin", "user"},
|
|
},
|
|
expected: []string{"admin", "user"},
|
|
expectError: false,
|
|
},
|
|
{
|
|
name: "groups as string array",
|
|
claims: map[string]interface{}{
|
|
"groups": []string{"developer", "tester"},
|
|
},
|
|
expected: []string{"developer", "tester"},
|
|
expectError: false,
|
|
},
|
|
{
|
|
name: "custom:groups as string",
|
|
claims: map[string]interface{}{
|
|
"custom:groups": "single-group",
|
|
},
|
|
expected: []string{"single-group"},
|
|
expectError: false,
|
|
},
|
|
{
|
|
name: "no groups in claims",
|
|
claims: map[string]interface{}{},
|
|
expected: nil,
|
|
expectError: true,
|
|
},
|
|
{
|
|
name: "groups with mixed types",
|
|
claims: map[string]interface{}{
|
|
"cognito:groups": []interface{}{"admin", 123, true},
|
|
},
|
|
expected: []string{"admin", "123", "true"},
|
|
expectError: false,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
groups, err := GetUserGroups(tt.claims)
|
|
|
|
if tt.expectError && err == nil {
|
|
t.Errorf("GetUserGroups() expected error, got nil")
|
|
}
|
|
|
|
if !tt.expectError && err != nil {
|
|
t.Errorf("GetUserGroups() unexpected error: %v", err)
|
|
}
|
|
|
|
if !reflect.DeepEqual(groups, tt.expected) {
|
|
t.Errorf("GetUserGroups() got = %v, want %v", groups, tt.expected)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestCreateCodeChallenge verifies the createCodeChallenge function correctly
|
|
// generates PKCE code challenges from code verifiers according to OAuth 2.0 PKCE
|
|
// (RFC 7636) specifications. It tests multiple input/output pairs to ensure the
|
|
// SHA-256 hashing and Base64-URL encoding are implemented correctly.
|
|
func TestCreateCodeChallenge(t *testing.T) {
|
|
tests := []struct {
|
|
verifier string
|
|
challenge string
|
|
}{
|
|
{
|
|
verifier: "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk",
|
|
challenge: "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM",
|
|
},
|
|
{
|
|
verifier: "abcdefghijklmnopqrstuvwxyz0123456789",
|
|
challenge: "AR_CmU450lEUFUD4emkJKz8iqGdn9yg95-7ts4l77fY",
|
|
},
|
|
{
|
|
verifier: "test-verifier",
|
|
challenge: "JBbiqONGWPaAmwXk_8bT6UnlPfrn65D32eZlJS-zGG0",
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.verifier, func(t *testing.T) {
|
|
challenge := createCodeChallenge(tt.verifier)
|
|
if challenge != tt.challenge {
|
|
t.Errorf("createCodeChallenge(%q) = %q, want %q", tt.verifier, challenge, tt.challenge)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestGenerateRandomString verifies the generateRandomString function produces
|
|
// cryptographically secure random strings of the requested length. It tests:
|
|
// - Generation of strings with various lengths
|
|
// - Correct string length
|
|
// - Base64-URL encoding validity
|
|
// - String uniqueness (randomness)
|
|
func TestGenerateRandomString(t *testing.T) {
|
|
lengths := []int{32, 43, 64, 128}
|
|
|
|
for _, length := range lengths {
|
|
t.Run(fmt.Sprintf("length=%d", length), func(t *testing.T) {
|
|
// Generate multiple strings to check for uniqueness
|
|
strings := make([]string, 5)
|
|
for i := 0; i < 5; i++ {
|
|
s, err := generateRandomString(length)
|
|
if err != nil {
|
|
t.Errorf("generateRandomString(%d) unexpected error: %v", length, err)
|
|
}
|
|
|
|
// Check length
|
|
if len(s) != length {
|
|
t.Errorf("generateRandomString(%d) = %q (length %d), want length %d",
|
|
length, s, len(s), length)
|
|
}
|
|
|
|
// Check if it's base64url valid
|
|
_, err = base64.RawURLEncoding.DecodeString(s)
|
|
if err != nil {
|
|
t.Errorf("generateRandomString(%d) = %q, not valid base64url: %v",
|
|
length, s, err)
|
|
}
|
|
|
|
strings[i] = s
|
|
}
|
|
|
|
// Check for duplicates (very unlikely but possible)
|
|
for i := 0; i < len(strings); i++ {
|
|
for j := i + 1; j < len(strings); j++ {
|
|
if strings[i] == strings[j] {
|
|
t.Errorf("generateRandomString(%d) generated duplicate strings: %q",
|
|
length, strings[i])
|
|
}
|
|
}
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestPKCESessionManagement verifies the PKCE session management functions work correctly.
|
|
// It tests:
|
|
// - Storing and retrieving valid code verifier sessions
|
|
// - Error handling when retrieving non-existent sessions
|
|
// - Proper expiration of sessions after their timeout period
|
|
// - Automatic cleanup of expired sessions during store operations
|
|
// - Persistence of valid sessions during cleanup
|
|
func TestPKCESessionManagement(t *testing.T) {
|
|
// Setup logger
|
|
logHandler := slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{
|
|
Level: slog.LevelDebug,
|
|
})
|
|
logger := slog.New(logHandler)
|
|
|
|
// Test storing and retrieving a valid session
|
|
t.Run("store and retrieve valid session", func(t *testing.T) {
|
|
state := "test-state-valid"
|
|
codeVerifier := "test-verifier-valid"
|
|
|
|
// Store the session
|
|
storePKCESession(state, codeVerifier, logger)
|
|
|
|
// Retrieve it
|
|
retrievedVerifier, err := getCodeVerifier(state, logger)
|
|
|
|
if err != nil {
|
|
t.Errorf("getCodeVerifier(%q) unexpected error: %v", state, err)
|
|
}
|
|
|
|
if retrievedVerifier != codeVerifier {
|
|
t.Errorf("getCodeVerifier(%q) = %q, want %q", state, retrievedVerifier, codeVerifier)
|
|
}
|
|
})
|
|
|
|
// Test retrieving non-existent session
|
|
t.Run("retrieve non-existent session", func(t *testing.T) {
|
|
state := "non-existent-state"
|
|
|
|
_, err := getCodeVerifier(state, logger)
|
|
|
|
if err == nil {
|
|
t.Errorf("getCodeVerifier(%q) expected error for non-existent state", state)
|
|
}
|
|
})
|
|
|
|
// Test expired session
|
|
t.Run("expired session", func(t *testing.T) {
|
|
state := "test-state-expired"
|
|
codeVerifier := "test-verifier-expired"
|
|
|
|
// Manually create an expired session
|
|
pkceSessionMap.Lock()
|
|
pkceSessionMap.sessions[state] = &PKCESession{
|
|
CodeVerifier: codeVerifier,
|
|
CreatedAt: time.Now().Add(-10 * time.Minute),
|
|
ExpiresAt: time.Now().Add(-5 * time.Minute), // Expired 5 minutes ago
|
|
}
|
|
pkceSessionMap.Unlock()
|
|
|
|
_, err := getCodeVerifier(state, logger)
|
|
|
|
if err == nil {
|
|
t.Errorf("getCodeVerifier(%q) expected error for expired session", state)
|
|
}
|
|
})
|
|
|
|
// Test cleanup of expired sessions during store
|
|
t.Run("cleanup expired sessions", func(t *testing.T) {
|
|
// Create several expired sessions
|
|
pkceSessionMap.Lock()
|
|
for i := 0; i < 5; i++ {
|
|
state := fmt.Sprintf("expired-state-%d", i)
|
|
pkceSessionMap.sessions[state] = &PKCESession{
|
|
CodeVerifier: fmt.Sprintf("expired-verifier-%d", i),
|
|
CreatedAt: time.Now().Add(-10 * time.Minute),
|
|
ExpiresAt: time.Now().Add(-1 * time.Minute), // Expired 1 minute ago
|
|
}
|
|
}
|
|
initialCount := len(pkceSessionMap.sessions)
|
|
pkceSessionMap.Unlock()
|
|
|
|
// Store a new session, which should trigger cleanup
|
|
newState := "new-test-state"
|
|
newVerifier := "new-test-verifier"
|
|
storePKCESession(newState, newVerifier, logger)
|
|
|
|
// Check if expired sessions were cleaned up
|
|
pkceSessionMap.RLock()
|
|
cleanedCount := len(pkceSessionMap.sessions)
|
|
pkceSessionMap.RUnlock()
|
|
|
|
if cleanedCount >= initialCount {
|
|
t.Errorf("storePKCESession() did not clean up expired sessions: before=%d, after=%d",
|
|
initialCount, cleanedCount)
|
|
}
|
|
|
|
// Verify the new session is still retrievable
|
|
retrievedVerifier, err := getCodeVerifier(newState, logger)
|
|
if err != nil || retrievedVerifier != newVerifier {
|
|
t.Errorf("New session not retrievable after cleanup: err=%v, verifier=%q",
|
|
err, retrievedVerifier)
|
|
}
|
|
})
|
|
}
|
|
|
|
// TestExtractUserInfo verifies the ExtractUserInfo function correctly extracts user
|
|
// information from JWT claims into a UserInfo struct. It tests:
|
|
// - Extraction of complete user information (username, email, groups)
|
|
// - Handling of missing email field
|
|
// - Handling of missing username field
|
|
// - Handling of missing groups field
|
|
// - Handling of completely empty claims
|
|
func TestExtractUserInfo(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
claims map[string]interface{}
|
|
want UserInfo
|
|
}{
|
|
{
|
|
name: "complete user info",
|
|
claims: map[string]interface{}{
|
|
"cognito:username": "testuser",
|
|
"email": "test@example.com",
|
|
"cognito:groups": []interface{}{"admin", "developer"},
|
|
},
|
|
want: UserInfo{
|
|
Username: "testuser",
|
|
Email: "test@example.com",
|
|
Groups: []string{"admin", "developer"},
|
|
},
|
|
},
|
|
{
|
|
name: "missing email",
|
|
claims: map[string]interface{}{
|
|
"cognito:username": "testuser",
|
|
"cognito:groups": []interface{}{"admin"},
|
|
},
|
|
want: UserInfo{
|
|
Username: "testuser",
|
|
Email: "",
|
|
Groups: []string{"admin"},
|
|
},
|
|
},
|
|
{
|
|
name: "missing username",
|
|
claims: map[string]interface{}{
|
|
"email": "test@example.com",
|
|
"cognito:groups": []interface{}{"admin"},
|
|
},
|
|
want: UserInfo{
|
|
Username: "",
|
|
Email: "test@example.com",
|
|
Groups: []string{"admin"},
|
|
},
|
|
},
|
|
{
|
|
name: "missing groups",
|
|
claims: map[string]interface{}{
|
|
"cognito:username": "testuser",
|
|
"email": "test@example.com",
|
|
},
|
|
want: UserInfo{
|
|
Username: "testuser",
|
|
Email: "test@example.com",
|
|
Groups: nil, // GetUserGroups will return error, so Groups will be nil
|
|
},
|
|
},
|
|
{
|
|
name: "empty claims",
|
|
claims: map[string]interface{}{},
|
|
want: UserInfo{
|
|
Username: "",
|
|
Email: "",
|
|
Groups: nil,
|
|
},
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
got := ExtractUserInfo(tt.claims)
|
|
|
|
if got.Username != tt.want.Username {
|
|
t.Errorf("ExtractUserInfo() username = %v, want %v", got.Username, tt.want.Username)
|
|
}
|
|
|
|
if got.Email != tt.want.Email {
|
|
t.Errorf("ExtractUserInfo() email = %v, want %v", got.Email, tt.want.Email)
|
|
}
|
|
|
|
// We don't compare Groups directly because GetUserGroups may return an error
|
|
// for the "missing groups" test case, resulting in nil Groups
|
|
})
|
|
}
|
|
}
|