From b23b17ac5334ab6e7ae0e25ab35c744cfab6e2fb Mon Sep 17 00:00:00 2001 From: jay brown Date: Tue, 22 Apr 2025 10:01:37 -0700 Subject: [PATCH] auth tests --- internal/cognitoauth/jwks_test.go | 326 ++++++++++++++ internal/cognitoauth/middleware_test.go | 315 ++++++++++++++ internal/cognitoauth/permission_utils_test.go | 221 ++++++++++ internal/cognitoauth/token_test.go | 407 ++++++++++++++++++ scripts/tests.yml | 2 +- 5 files changed, 1270 insertions(+), 1 deletion(-) create mode 100644 internal/cognitoauth/jwks_test.go create mode 100644 internal/cognitoauth/middleware_test.go create mode 100644 internal/cognitoauth/permission_utils_test.go create mode 100644 internal/cognitoauth/token_test.go diff --git a/internal/cognitoauth/jwks_test.go b/internal/cognitoauth/jwks_test.go new file mode 100644 index 00000000..f071bb63 --- /dev/null +++ b/internal/cognitoauth/jwks_test.go @@ -0,0 +1,326 @@ +package cognitoauth + +import ( + "crypto/rand" + "crypto/rsa" + "encoding/json" + "log/slog" + "net/http" + "net/http/httptest" + "os" + "testing" + "time" + + "github.com/lestrrat-go/jwx/v2/jwa" + "github.com/lestrrat-go/jwx/v2/jwk" +) + +// generateTestJWKS creates a test JSON Web Key Set (JWKS) for testing purposes. +// It performs the following operations: +// - Generates a 2048-bit RSA key pair for cryptographic operations +// - Creates a JWK (JSON Web Key) from the public key portion +// - Sets key properties including ID ("test-key-id"), algorithm (RS256), and usage ("sig") +// - Creates a JWK Set and adds the key to it +// - Returns both the JWK Set object and its JSON string representation +func generateTestJWKS(t *testing.T) (jwk.Set, string) { + // Generate RSA key for testing + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("failed to generate RSA key: %v", err) + } + + // Create JWK from RSA key + key, err := jwk.FromRaw(privateKey.PublicKey) + if err != nil { + t.Fatalf("failed to create JWK: %v", err) + } + + // Set key properties + if err := key.Set(jwk.KeyIDKey, "test-key-id"); err != nil { + t.Fatalf("failed to set key ID: %v", err) + } + if err := key.Set(jwk.AlgorithmKey, jwa.RS256); err != nil { + t.Fatalf("failed to set algorithm: %v", err) + } + if err := key.Set(jwk.KeyUsageKey, "sig"); err != nil { + t.Fatalf("failed to set key usage: %v", err) + } + + // Create JWK Set with the key + keySet := jwk.NewSet() + if err := keySet.AddKey(key); err != nil { + t.Fatalf("failed to add key to key set: %v", err) + } + + // Serialize to JSON + jwksJSON, err := json.Marshal(keySet) + if err != nil { + t.Fatalf("failed to marshal JWK Set: %v", err) + } + + return keySet, string(jwksJSON) +} + +// TestGetJWKS tests the GetJWKS function through multiple scenarios including +// initial fetching, caching, cache expiration, and error handling. +func TestGetJWKS(t *testing.T) { + // Setup logger + logHandler := slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + Level: slog.LevelDebug, + }) + logger := slog.New(logHandler) + + // Generate test JWK set for mock responses + testKeySet, jwksJSON := generateTestJWKS(t) + + // Mock server to return JWKS + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, err := w.Write([]byte(jwksJSON)) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer server.Close() + + // Reset JWKS cache before tests + jwksCache = &JWKSCache{ + ExpiresAt: time.Now(), // Initial state is expired + } + + // Test: First fetch - should make HTTP request + // Tests initial JWKS retrieval and verifies: + // - Successful HTTP request and response processing + // - Correct number of keys in the retrieved set + // - Proper caching of the retrieved key set + // - Correct cache expiration time setting + t.Run("first fetch", func(t *testing.T) { + keySet, err := GetJWKS(server.URL, logger) + if err != nil { + t.Fatalf("GetJWKS() error = %v", err) + } + + // Verify key set contains expected keys + if keySet.Len() != testKeySet.Len() { + t.Errorf("GetJWKS() returned %d keys, want %d", keySet.Len(), testKeySet.Len()) + } + + // Check that it's cached + if jwksCache.KeySet == nil { + t.Errorf("KeySet was not cached after fetch") + } + if time.Now().After(jwksCache.ExpiresAt) { + t.Errorf("Cache expiration not set correctly") + } + }) + + // Test: Subsequent fetch - should use cache + // Verifies caching behavior: + // - Tests that subsequent requests use cached data instead of making HTTP requests + // - Tests the cache works even when the server is unavailable/failing + // - Validates that the cached key count matches expected values + t.Run("cached fetch", func(t *testing.T) { + // Create a mock server that fails - to verify we're using the cache + failServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + _, err := w.Write([]byte("This server should not be called")) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer failServer.Close() + + // We should still get valid keys because we use the cache + keySet, err := GetJWKS(failServer.URL, logger) + if err != nil { + t.Fatalf("GetJWKS() with cache error = %v", err) + } + + // Verify key set contains expected keys + if keySet.Len() != testKeySet.Len() { + t.Errorf("GetJWKS() with cache returned %d keys, want %d", keySet.Len(), testKeySet.Len()) + } + }) + + // Test: Cache expiry - should fetch again + // Tests cache expiration behavior: + // - Verifies new keys are fetched after cache expiration + // - Ensures the system properly handles refreshing expired cache data + // - Validates that updated keys are received correctly + t.Run("cache expiry", func(t *testing.T) { + // Force cache expiration + jwksCache.mutex.Lock() + jwksCache.ExpiresAt = time.Now().Add(-1 * time.Hour) // Expired one hour ago + jwksCache.mutex.Unlock() + + // New server with updated keys + updatedKeySet, updatedJSON := generateTestJWKS(t) + updateServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, err := w.Write([]byte(updatedJSON)) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer updateServer.Close() + + // Now we should fetch new keys + keySet, err := GetJWKS(updateServer.URL, logger) + if err != nil { + t.Fatalf("GetJWKS() after expiry error = %v", err) + } + + // Verify key set contains expected keys (should match updated set) + if keySet.Len() != updatedKeySet.Len() { + t.Errorf("GetJWKS() after expiry returned %d keys, want %d", keySet.Len(), updatedKeySet.Len()) + } + }) + + // Test: Server error + // Tests error handling when server fails: + // - Validates behavior with 500 Internal Server Error responses + // - Verifies appropriate error is returned to caller + // - Tests system behavior with empty/reset cache + t.Run("server error", func(t *testing.T) { + // Reset cache + jwksCache = &JWKSCache{ + ExpiresAt: time.Now(), // Initial state is expired + } + + // Create a mock server that returns an error + errorServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + _, err := w.Write([]byte("Internal Server Error")) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer errorServer.Close() + + // Should return an error + _, err := GetJWKS(errorServer.URL, logger) + if err == nil { + t.Errorf("GetJWKS() with server error should return error, got nil") + } + }) + + // Test: Invalid JSON response + // Tests handling of malformed responses: + // - Verifies error handling for invalid JSON responses + // - Ensures system stability when receiving corrupted data + // - Validates error reporting for bad responses + t.Run("invalid json", func(t *testing.T) { + // Reset cache + jwksCache = &JWKSCache{ + ExpiresAt: time.Now(), // Initial state is expired + } + + // Create a mock server that returns invalid JSON + invalidServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, err := w.Write([]byte("This is not valid JSON")) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer invalidServer.Close() + + // Should return an error + _, err := GetJWKS(invalidServer.URL, logger) + if err == nil { + t.Errorf("GetJWKS() with invalid JSON should return error, got nil") + } + }) +} + +// TestFetchJWKS tests the fetchJWKS function directly with various scenarios +// including successful fetches, timeouts, invalid URLs, and error responses. +func TestFetchJWKS(t *testing.T) { + // Generate test JWK set for mock responses + _, jwksJSON := generateTestJWKS(t) + + // Test: Successful fetch + // Verifies basic JWKS retrieval functionality: + // - Tests successful HTTP request and response handling + // - Validates that response body matches expected JSON + // - Confirms proper return of fetched content + t.Run("successful fetch", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, err := w.Write([]byte(jwksJSON)) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer server.Close() + + result, err := fetchJWKS(server.URL) + if err != nil { + t.Fatalf("fetchJWKS() error = %v", err) + } + + if result != jwksJSON { + t.Errorf("fetchJWKS() = %v, want %v", result, jwksJSON) + } + }) + + // Test: Timeout + // Validates timeout handling: + // - Tests behavior when server response exceeds client timeout limit + // - Verifies appropriate error handling for timed-out requests + // - Ensures the function fails properly in timeout scenarios + t.Run("timeout", func(t *testing.T) { + // Create a server that sleeps longer than the client timeout + slowServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + time.Sleep(JWKS_CLIENT_TIMEOUT + 1*time.Second) + _, err := w.Write([]byte(jwksJSON)) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer slowServer.Close() + + _, err := fetchJWKS(slowServer.URL) + if err == nil { + t.Errorf("fetchJWKS() with timeout should return error, got nil") + } + }) + + // Test: Invalid URL + // Tests error handling for bad URLs: + // - Verifies behavior with non-existent or invalid URLs + // - Ensures appropriate error responses for network failures + // - Validates error handling for connection issues + t.Run("invalid url", func(t *testing.T) { + _, err := fetchJWKS("http://invalid-url-that-does-not-exist.example.com") + if err == nil { + t.Errorf("fetchJWKS() with invalid URL should return error, got nil") + } + }) + + // Test: HTTP error status + // Validates handling of HTTP error responses: + // - Tests behavior with various HTTP error status codes (e.g., 404) + // - Verifies response body handling regardless of status code + // - Tests that function returns body content even for error status codes + t.Run("http error status", func(t *testing.T) { + errorServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + _, err := w.Write([]byte("Not Found")) + if err != nil { + t.Errorf("Failed to write response: %v", err) + } + })) + defer errorServer.Close() + + result, err := fetchJWKS(errorServer.URL) + if err != nil { + t.Fatalf("fetchJWKS() unexpected error = %v", err) + } + + // Function doesn't check status code, just reads body + if result != "Not Found" { + t.Errorf("fetchJWKS() = %v, want %v", result, "Not Found") + } + }) +} diff --git a/internal/cognitoauth/middleware_test.go b/internal/cognitoauth/middleware_test.go new file mode 100644 index 00000000..56cfc663 --- /dev/null +++ b/internal/cognitoauth/middleware_test.go @@ -0,0 +1,315 @@ +package cognitoauth + +import ( + "encoding/json" + "log/slog" + "net/http" + "net/http/httptest" + "os" + "testing" + "time" + + "github.com/labstack/echo/v4" +) + +// TestJWTAuthMiddleware tests the JWTAuthMiddleware function to ensure it correctly processes +// authentication tokens from both headers and cookies. It verifies that: +// - Public paths (/login, /callback, /home) are skipped without authentication +// - The /logout path properly clears authentication cookies +// - Requests with Authorization headers are passed through +// - Requests with auth_token cookies have their tokens moved to Authorization headers +// - Requests without authentication are redirected to the login page +func TestJWTAuthMiddleware(t *testing.T) { + // Setup + e := echo.New() + config := &mockConfigProvider{ + region: "us-east-2", + userPoolID: "us-east-2_testpool", + } + + // Setup logger + logHandler := slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + Level: slog.LevelDebug, + }) + logger := slog.New(logHandler) + config.SetAuthLogger(logger) + + // Create middleware + middleware := JWTAuthMiddleware(config) + + // Test handler that simply returns success + testHandler := func(c echo.Context) error { + return c.String(http.StatusOK, "success") + } + + // Test cases + tests := []struct { + name string + path string + setupRequest func(*http.Request) + expectStatus int + expectLocation string + expectHeader string + }{ + { + name: "skip login path", + path: "/login", + setupRequest: func(req *http.Request) { + // No setup needed + }, + expectStatus: http.StatusOK, + }, + { + name: "skip callback path", + path: "/callback", + setupRequest: func(req *http.Request) { + // No setup needed + }, + expectStatus: http.StatusOK, + }, + { + name: "skip home path", + path: "/home", + setupRequest: func(req *http.Request) { + // No setup needed + }, + expectStatus: http.StatusOK, + }, + { + name: "logout path sets empty cookie", + path: "/logout", + setupRequest: func(req *http.Request) { + // No setup needed + }, + expectStatus: http.StatusOK, + }, + { + name: "use authorization header", + path: "/api/test", + setupRequest: func(req *http.Request) { + req.Header.Set("Authorization", "Bearer test-token") + }, + expectStatus: http.StatusOK, + }, + { + name: "use cookie if no authorization header", + path: "/api/test", + setupRequest: func(req *http.Request) { + cookie := &http.Cookie{ + Name: "auth_token", + Value: "cookie-token", + Path: "/", + } + req.AddCookie(cookie) + }, + expectStatus: http.StatusOK, + expectHeader: "Bearer cookie-token", + }, + { + name: "redirect to login if no token", + path: "/api/test", + setupRequest: func(req *http.Request) { + // No auth header or cookie + }, + expectStatus: http.StatusFound, // 302 Found (redirect) + expectLocation: "/login", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, tt.path, nil) + rec := httptest.NewRecorder() + + // Setup request (add headers, cookies, etc.) + tt.setupRequest(req) + + // Create Echo context + c := e.NewContext(req, rec) + c.SetPath(tt.path) + + // Run middleware + err := middleware(testHandler)(c) + if err != nil { + t.Errorf("Middleware returned error: %v", err) + } + + // Check status code + if rec.Code != tt.expectStatus { + t.Errorf("Expected status code %d, got %d", tt.expectStatus, rec.Code) + } + + // Check location header (for redirects) + if tt.expectLocation != "" { + location := rec.Header().Get("Location") + if location != tt.expectLocation { + t.Errorf("Expected Location header %q, got %q", tt.expectLocation, location) + } + } + + // Check Authorization header (if token taken from cookie) + if tt.expectHeader != "" { + authHeader := req.Header.Get("Authorization") + if authHeader != tt.expectHeader { + t.Errorf("Expected Authorization header %q, got %q", tt.expectHeader, authHeader) + } + } + + // Special case for logout - check cookie + if tt.path == "/logout" { + cookies := rec.Result().Cookies() + found := false + for _, cookie := range cookies { + if cookie.Name == "auth_token" { + found = true + if cookie.Value != "" { + t.Errorf("Expected empty auth_token cookie, got %q", cookie.Value) + } + if !cookie.Expires.Before(time.Now()) { + t.Errorf("Expected auth_token cookie to be expired, got %v", cookie.Expires) + } + } + } + if !found { + t.Error("Expected auth_token cookie to be set for logout") + } + } + }) + } +} + +// TestTokenValidationMiddleware tests the TokenValidationMiddleware function to verify that it: +// - Allows access to public paths (/home, /logout) without validation +// - Properly handles OAuth callbacks by processing code and state parameters +// - Initiates login by redirecting to Cognito when accessing the login path +// - Returns 401 Unauthorized for protected paths without Authorization header +// - Attempts to validate tokens for protected paths with Authorization headers +func TestTokenValidationMiddleware(t *testing.T) { + // Setup + e := echo.New() + config := &mockConfigProvider{ + region: "us-east-2", + userPoolID: "us-east-2_testpool", + } + + // Setup logger + logHandler := slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + Level: slog.LevelDebug, + }) + logger := slog.New(logHandler) + config.SetAuthLogger(logger) + + // Create middleware + middleware := TokenValidationMiddleware(config) + + // Test handler that simply returns success + testHandler := func(c echo.Context) error { + return c.String(http.StatusOK, "success") + } + + // Test cases + tests := []struct { + name string + path string + setupRequest func(*http.Request) + expectStatus int + checkBody func(*testing.T, []byte) + }{ + { + name: "skip public paths - home", + path: "/home", + setupRequest: func(req *http.Request) { + // No setup needed + }, + expectStatus: http.StatusOK, + }, + { + name: "skip public paths - logout", + path: "/logout", + setupRequest: func(req *http.Request) { + // No setup needed + }, + expectStatus: http.StatusOK, + }, + { + name: "handle oauth callback", + path: "/callback", + setupRequest: func(req *http.Request) { + q := req.URL.Query() + q.Add("code", "test-code") + q.Add("state", "test-state") + req.URL.RawQuery = q.Encode() + }, + // This will attempt to call handleOAuthCallback which will fail + // but we're just testing the routing logic + expectStatus: http.StatusBadRequest, + }, + { + name: "handle login initiation", + path: "/login", + setupRequest: func(req *http.Request) { + // No setup needed + }, + // The login initiation actually redirects to the Cognito login page + expectStatus: http.StatusFound, // 302 Found (redirect) + // We can't check the exact URL since it depends on generateRandomString output + // But we can check that a Location header exists + checkBody: func(t *testing.T, body []byte) { + // No body check needed, we'll verify the status code + }, + }, + { + name: "missing authorization header", + path: "/api/test", + setupRequest: func(req *http.Request) { + // No auth header + }, + expectStatus: http.StatusUnauthorized, + checkBody: func(t *testing.T, body []byte) { + var response map[string]string + err := json.Unmarshal(body, &response) + if err != nil { + t.Errorf("Failed to parse response: %v", err) + } + if response["error"] != "Authorization required" { + t.Errorf("Expected error 'Authorization required', got %q", response["error"]) + } + }, + }, + { + name: "invalid token", + path: "/api/test", + setupRequest: func(req *http.Request) { + req.Header.Set("Authorization", "Bearer invalid-token") + }, + expectStatus: http.StatusInternalServerError, // Since we can't mock JWKS fetching + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, tt.path, nil) + rec := httptest.NewRecorder() + + // Setup request + tt.setupRequest(req) + + // Create Echo context + c := e.NewContext(req, rec) + c.SetPath(tt.path) + + // Run middleware + _ = middleware(testHandler)(c) + + // Check status code + if rec.Code != tt.expectStatus { + t.Errorf("Expected status code %d, got %d", tt.expectStatus, rec.Code) + } + + // Check response body if needed + if tt.checkBody != nil { + tt.checkBody(t, rec.Body.Bytes()) + } + }) + } +} diff --git a/internal/cognitoauth/permission_utils_test.go b/internal/cognitoauth/permission_utils_test.go new file mode 100644 index 00000000..10b5b5e6 --- /dev/null +++ b/internal/cognitoauth/permission_utils_test.go @@ -0,0 +1,221 @@ +package cognitoauth + +import ( + "log/slog" + "os" + "reflect" + "testing" +) + +// TestMatchRoute verifies the matchRoute function correctly identifies when a request path +// matches a route pattern. It tests various scenarios including: +// - Exact path matching +// - Path parameter matching (e.g., /:id) +// - Multiple parameter matching +// - Wildcard matching +// - Path mismatches +// - Partial matching behavior +// - Special characters in paths +// - Edge cases like empty paths and root paths +func TestMatchRoute(t *testing.T) { + tests := []struct { + name string + requestPath string + routePattern string + expected bool + }{ + { + name: "exact match", + requestPath: "/api/users", + routePattern: "/api/users", + expected: true, + }, + { + name: "parameter match", + requestPath: "/api/users/123", + routePattern: "/api/users/:id", + expected: true, + }, + { + name: "multiple parameters", + requestPath: "/api/users/123/orders/456", + routePattern: "/api/users/:userId/orders/:orderId", + expected: true, + }, + { + name: "wildcard match", + requestPath: "/api/v1/docs/user-guide.pdf", + routePattern: "/api/v1/docs/*", + expected: true, + }, + { + name: "no match", + requestPath: "/api/orders", + routePattern: "/api/users", + expected: false, + }, + { + name: "partial match should fail", + requestPath: "/api/users/123", + routePattern: "/api/users", + expected: false, + }, + { + name: "longer request path should fail", + requestPath: "/api/users/123/details", + routePattern: "/api/users/:id", + expected: false, + }, + { + name: "regex characters in path", + requestPath: "/api/users/john.doe+test@example.com", + routePattern: "/api/users/:email", + expected: true, + }, + { + name: "empty request path", + requestPath: "", + routePattern: "/api/users", + expected: false, + }, + { + name: "empty route pattern", + requestPath: "/api/users", + routePattern: "", + expected: false, + }, + { + name: "root path", + requestPath: "/", + routePattern: "/", + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := matchRoute(tt.requestPath, tt.routePattern) + if result != tt.expected { + t.Errorf("matchRoute(%q, %q) = %v, want %v", + tt.requestPath, tt.routePattern, result, tt.expected) + } + }) + } +} + +// TestCheckPermissions verifies the checkPermissions function correctly enforces +// access control based on user groups and route permissions. It tests: +// - Admin access to various routes +// - Group-specific access permissions +// - Denying access when user lacks required group membership +// - Public route access without authentication +// - Multiple group membership scenarios +// - Routes with non-existent required groups +// - Path parameter matching with permissions +// - Default behavior for unmatched routes +func TestCheckPermissions(t *testing.T) { + // Setup logger + logHandler := slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + Level: slog.LevelDebug, + }) + logger := slog.New(logHandler) + + // Setup route permissions + routePermissions := map[string][]string{ + "/api/users": {"admin", "user-manager"}, + "/api/orders": {"admin", "order-manager"}, + "/api/products": {"admin", "product-manager"}, + "/api/reports/*": {"admin", "analyst"}, + "/api/users/:id": {"admin", "user-manager", "self-service"}, + "/public/*": {}, // Public route - empty permissions means anyone can access + "/api/restricted": {"non-existent-group"}, // No one can access this + } + + tests := []struct { + name string + path string + userGroups []string + expectedAllowed bool + expectedGroups []string + }{ + { + name: "admin can access user route", + path: "/api/users", + userGroups: []string{"admin"}, + expectedAllowed: true, + expectedGroups: []string{"admin", "user-manager"}, + }, + { + name: "user-manager can access user route", + path: "/api/users", + userGroups: []string{"user-manager"}, + expectedAllowed: true, + expectedGroups: []string{"admin", "user-manager"}, + }, + { + name: "order-manager cannot access user route", + path: "/api/users", + userGroups: []string{"order-manager"}, + expectedAllowed: false, + expectedGroups: []string{"admin", "user-manager"}, + }, + { + name: "empty user groups cannot access protected route", + path: "/api/users", + userGroups: []string{}, + expectedAllowed: false, + expectedGroups: []string{"admin", "user-manager"}, + }, + { + name: "anyone can access public route", + path: "/public/info", + userGroups: []string{}, + expectedAllowed: true, + expectedGroups: []string{}, + }, + { + name: "user with multiple groups can access based on any matching group", + path: "/api/reports/sales", + userGroups: []string{"user-manager", "analyst"}, + expectedAllowed: true, + expectedGroups: []string{"admin", "analyst"}, + }, + { + name: "no one can access route with non-existent group requirement", + path: "/api/restricted", + userGroups: []string{"admin", "user-manager", "order-manager"}, + expectedAllowed: false, + expectedGroups: []string{"non-existent-group"}, + }, + { + name: "self-service can access specific user route", + path: "/api/users/123", + userGroups: []string{"self-service"}, + expectedAllowed: true, + expectedGroups: []string{"admin", "user-manager", "self-service"}, + }, + { + name: "unmatched route defaults to allowed", + path: "/unknown/route", + userGroups: []string{}, + expectedAllowed: true, + expectedGroups: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + allowed, groups := checkPermissions(tt.path, tt.userGroups, routePermissions, logger) + + if allowed != tt.expectedAllowed { + t.Errorf("checkPermissions(%q, %v) allowed = %v, want %v", + tt.path, tt.userGroups, allowed, tt.expectedAllowed) + } + + if !reflect.DeepEqual(groups, tt.expectedGroups) { + t.Errorf("checkPermissions(%q, %v) groups = %v, want %v", + tt.path, tt.userGroups, groups, tt.expectedGroups) + } + }) + } +} diff --git a/internal/cognitoauth/token_test.go b/internal/cognitoauth/token_test.go new file mode 100644 index 00000000..f05136db --- /dev/null +++ b/internal/cognitoauth/token_test.go @@ -0,0 +1,407 @@ +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) 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 + }) + } +} diff --git a/scripts/tests.yml b/scripts/tests.yml index d4d9ecea..ab8c8f62 100644 --- a/scripts/tests.yml +++ b/scripts/tests.yml @@ -10,7 +10,7 @@ vars: PKG: "./pkg/..." CMD: "./cmd/..." # yamllint disable-line rule:line-length - EXCLUDED_FILES: ".gen.go|internal/serviceconfig/observability/prometheus/generator/main.go|internal/cognitoauth/auth.go|internal/cognitoauth/handler.go|internal/cognitoauth/jwks.go|internal/cognitoauth/middleware.go|internal/cognitoauth/models.go|internal/cognitoauth/token.go|internal/cognitoauth/permission_utils.go|api/queryAPI/authHandlers.go|api/queryAPI/homehandler.go" + EXCLUDED_FILES: ".gen.go|internal/serviceconfig/observability/prometheus/generator/main.go|internal/cognitoauth/middleware.go|internal/cognitoauth/token.go|internal/cognitoauth/auth.go|internal/cognitoauth/handler.go|internal/cognitoauth/models.go|api/queryAPI/authHandlers.go|api/queryAPI/homehandler.go" includes: docker: