Files
query-orchestration/internal/cognitoauth/middleware_test.go
T
Jay Brown 2ff6e05ffc Merged in feature/permit-integration-1 (pull request #172)
Permit integration - part 1

* WIP permit integration

* slim down build

* Merge branch 'main' of bitbucket.org:aarete/query-orchestration into feature/permit-integration-1

* omit swagger from auth

* clean build

* comment

* part 1 completed

this is the initial permitio parts and tests without integration. Also testing.md since we have a new test target `task test:permitio`

* feature flag for no jwt validation

* permit integration

* fix ci unit tests

* ci fix

* build fix

* fix home handler

* fix redirect

* test fix

* update docs for auth
2025-07-11 19:27:14 +00:00

315 lines
8.2 KiB
Go

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: "continue to next middleware if no token",
path: "/api/test",
setupRequest: func(req *http.Request) {
// No auth header or cookie
},
expectStatus: http.StatusOK, // Should continue to test handler
},
}
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())
}
})
}
}