//go:build test package auth // Package auth contains tests for JWT generation, verification, and JTI revocation. // // Test Coverage: // - generateJTI: UUID v4 format validation, uniqueness // - GenerateToken: returns non-empty token and JTI, JTI matches claim // - VerifyToken: returns correct user_id/role/JTI, rejects revoked/missing JTI // - RevokeJTI: adds to revoked set, IsJTIRevoked reflects changes // - CleanupRevokedJTIs: removes expired JTIs, keeps valid ones import ( "context" "strings" "testing" "time" "crussell/clock" "crussell/db" "crussell/testutils/fixtures" "crussell/testutils/testtx" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // ============================================================================= // generateJTI Tests // ============================================================================= // TestGenerateJTI_Format verifies that generateJTI returns a valid UUID v4 string // in the format: 8-4-4-4-12 hexadecimal digits with version and variant bits set. func TestGenerateJTI_Format(t *testing.T) { jti, err := generateJTI() if err != nil { t.Fatalf("generateJTI() failed: %v", err) } // UUID v4 format: xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx parts := strings.Split(jti, "-") if len(parts) != 5 { t.Errorf("expected 5 parts in UUID, got %d: %s", len(parts), jti) } // Check each part length: 8-4-4-4-12 expectedLengths := []int{8, 4, 4, 4, 12} for i, part := range parts { if len(part) != expectedLengths[i] { t.Errorf("part %d expected length %d, got %d (jti=%s)", i, expectedLengths[i], len(part), jti) } } // Check version nibble (4xxx) in the third group if len(parts[2]) > 0 && parts[2][0] != '4' { t.Errorf("expected version 4 UUID, got version %c in jti=%s", parts[2][0], jti) } // Check variant bits (8xxx, 9xxx, axxx, or bxxx) in the fourth group if len(parts[3]) > 0 { c := parts[3][0] if c != '8' && c != '9' && c != 'a' && c != 'b' { t.Errorf("expected variant bits 10xx, got %c in jti=%s", c, jti) } } } // TestGenerateJTI_Unique generates 100 JTIs and verifies all are unique. func TestGenerateJTI_Unique(t *testing.T) { seen := make(map[string]bool) for i := 0; i < 100; i++ { jti, err := generateJTI() if err != nil { t.Fatalf("generateJTI() failed at iteration %d: %v", i, err) } if seen[jti] { t.Errorf("duplicate JTI generated at iteration %d: %s", i, jti) } seen[jti] = true } if len(seen) != 100 { t.Errorf("expected 100 unique JTIs, got %d", len(seen)) } } // ============================================================================= // GenerateToken Tests // ============================================================================= // TestGenerateToken_ReturnsJTI verifies that GenerateToken returns both a // non-empty token string and a non-empty JTI string. func TestGenerateToken_ReturnsJTI(t *testing.T) { token, jti, err := GenerateToken("user-001", "verified_email") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } if token == "" { t.Error("expected non-empty token") } if jti == "" { t.Error("expected non-empty JTI") } } // TestGenerateToken_JTIInClaims decodes the generated token and verifies // the "jti" claim matches the returned JTI value. func TestGenerateToken_JTIInClaims(t *testing.T) { token, jti, err := GenerateToken("user-002", "admin") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } decoded, err := TokenAuth.Decode(token) if err != nil { t.Fatalf("failed to decode token: %v", err) } var claimJTI string if err := decoded.Get("jti", &claimJTI); err != nil { t.Fatalf("failed to get jti claim: %v", err) } if claimJTI != jti { t.Errorf("expected jti claim %q, got %q", jti, claimJTI) } } // ============================================================================= // VerifyToken Tests // ============================================================================= // TestVerifyToken_ReturnsJTI creates a token, verifies it, and checks the // returned user_id, role, and JTI match the expected values. func TestVerifyToken_ReturnsJTI(t *testing.T) { ctx, _ := testtx.SetupTestTx(t) token, jti, err := GenerateToken("user-003", "verified_email") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } userID, role, returnedJTI, err := VerifyToken(token, ctx) if err != nil { t.Fatalf("VerifyToken() failed: %v", err) } if userID != "user-003" { t.Errorf("expected userID 'user-003', got %q", userID) } if role != "verified_email" { t.Errorf("expected role 'verified_email', got %q", role) } if returnedJTI != jti { t.Errorf("expected JTI %q, got %q", jti, returnedJTI) } } // TestVerifyToken_RevokedJTI creates a token, revokes its JTI, and verifies // that VerifyToken returns an error containing "token revoked". func TestVerifyToken_RevokedJTI(t *testing.T) { ctx, _ := testtx.SetupTestTx(t) token, jti, err := GenerateToken("user-004", "verified_email") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } RevokeJTI(ctx, jti, clock.Now().Add(30*24*time.Hour)) _, _, _, err = VerifyToken(token, ctx) if err == nil { t.Fatal("expected error for revoked JTI, got nil") } if !strings.Contains(err.Error(), "token revoked") { t.Errorf("expected 'token revoked' error, got: %v", err) } } // TestVerifyToken_MissingJTI creates a token without a "jti" claim (using // TokenAuth.Encode directly) and verifies that VerifyToken returns an error // containing "invalid jti claim". func TestVerifyToken_MissingJTI(t *testing.T) { _, tokenString, err := TokenAuth.Encode(map[string]interface{}{ "user_id": "user-005", "role": "verified_email", "exp": clock.Now().Add(30 * 24 * time.Hour).Unix(), }) if err != nil { t.Fatalf("failed to create token without JTI: %v", err) } _, _, _, err = VerifyToken(tokenString, context.Background()) if err == nil { t.Fatal("expected error for missing JTI, got nil") } if !strings.Contains(err.Error(), "invalid jti claim") { t.Errorf("expected 'invalid jti claim' error, got: %v", err) } } // TestIsJTIRevoked_QueryError verifies that IsJTIRevoked returns false when // the context is cancelled (causing the QueryRow to fail). func TestIsJTIRevoked_QueryError(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() if IsJTIRevoked(ctx, "test-jti") { t.Error("expected false on cancelled context query error") } } // ============================================================================= // RevokeJTI / IsJTIRevoked Tests // ============================================================================= // TestRevokeJTI_AddsToSet verifies that calling RevokeJTI adds the JTI to the // revoked set, and IsJTIRevoked returns true for it. func TestRevokeJTI_AddsToSet(t *testing.T) { ctx, _ := testtx.SetupTestTx(t) _, jti, err := GenerateToken("user-006", "verified_email") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } if IsJTIRevoked(ctx, jti) { t.Fatal("JTI should not be revoked before calling RevokeJTI") } RevokeJTI(ctx, jti, clock.Now().Add(30*24*time.Hour)) if !IsJTIRevoked(ctx, jti) { t.Error("expected IsJTIRevoked to return true after RevokeJTI") } } // TestIsJTIRevoked_NonExistent verifies that checking a non-existent JTI // returns false. func TestIsJTIRevoked_NonExistent(t *testing.T) { ctx, _ := testtx.SetupTestTx(t) if IsJTIRevoked(ctx, "nonexistent-jti-12345") { t.Error("expected IsJTIRevoked to return false for non-existent JTI") } } // ============================================================================= // CleanupRevokedJTIs Tests // ============================================================================= // TestCleanupRevokedJTIs_RemovesExpired adds a JTI with a past expiry time, // runs CleanupRevokedJTIs, and verifies the JTI is removed from the set. func TestCleanupRevokedJTIs_RemovesExpired(t *testing.T) { ctx, tx := testtx.SetupTestTx(t) _, jti, err := GenerateToken("user-007", "verified_email") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } // Add with future expiry so IsJTIRevoked sees it RevokeJTI(ctx, jti, clock.Now().Add(1*time.Hour)) if !IsJTIRevoked(ctx, jti) { t.Fatal("JTI should be in revoked set after RevokeJTI") } // Directly update the DB to set expiry in the past _, err = tx.Exec(ctx, "UPDATE revoked_jtis SET expires_at = NOW() - INTERVAL '1 hour' WHERE jti = $1", jti) if err != nil { t.Fatalf("failed to expire JTI: %v", err) } CleanupRevokedJTIs(ctx) if IsJTIRevoked(ctx, jti) { t.Error("expected expired JTI to be removed after cleanup") } } // TestCleanupRevokedJTIs_KeepsValid adds a JTI with a future expiry time, // runs CleanupRevokedJTIs, and verifies the JTI is still in the set. func TestCleanupRevokedJTIs_KeepsValid(t *testing.T) { ctx, _ := testtx.SetupTestTx(t) _, jti, err := GenerateToken("user-008", "verified_email") if err != nil { t.Fatalf("GenerateToken() failed: %v", err) } // Add with future expiry RevokeJTI(ctx, jti, clock.Now().Add(30*24*time.Hour)) if !IsJTIRevoked(ctx, jti) { t.Fatal("JTI should be in revoked set before cleanup") } CleanupRevokedJTIs(ctx) if !IsJTIRevoked(ctx, jti) { t.Error("expected valid (future expiry) JTI to remain after cleanup") } } // ============================================================================= // generateRefreshTokenString Tests // ============================================================================= // TestGenerateRefreshTokenString_Format verifies that generateRefreshTokenString // returns a 64-character hex string. func TestGenerateRefreshTokenString_Format(t *testing.T) { token, err := generateRefreshTokenString() if err != nil { t.Fatalf("generateRefreshTokenString() failed: %v", err) } if len(token) != 64 { t.Errorf("expected 64-character hex string, got length %d: %s", len(token), token) } // Verify all characters are valid lowercase hex for _, c := range token { if !((c >= '0' && c <= '9') || (c >= 'a' && c <= 'f')) { t.Errorf("non-hex character %c in token %s", c, token) break } } } // TestGenerateRefreshTokenString_Unique generates 100 refresh token strings // and verifies all are unique. func TestGenerateRefreshTokenString_Unique(t *testing.T) { seen := make(map[string]bool) for i := 0; i < 100; i++ { token, err := generateRefreshTokenString() if err != nil { t.Fatalf("generateRefreshTokenString() failed at iteration %d: %v", i, err) } if seen[token] { t.Errorf("duplicate refresh token at iteration %d: %s", i, token) } seen[token] = true } if len(seen) != 100 { t.Errorf("expected 100 unique tokens, got %d", len(seen)) } } // ============================================================================= // GenerateRefreshToken Tests // ============================================================================= // TestGenerateRefreshToken_Success calls GenerateRefreshToken and verifies a // row was inserted in the refresh_tokens table with the correct user_id and role. func TestGenerateRefreshToken_Success(t *testing.T) { ctx, tx := testtx.SetupTestTx(t) userID, err := fixtures.CreateTestUser(tx) if err != nil { t.Fatalf("failed to create test user: %v", err) } token, err := GenerateRefreshToken(ctx, userID, "verified_email") if err != nil { t.Fatalf("GenerateRefreshToken() failed: %v", err) } if token == "" { t.Error("expected non-empty token") } // Verify row was inserted in refresh_tokens var dbUserID string var dbRole string err = tx.QueryRow(ctx, `SELECT user_id, role FROM refresh_tokens WHERE token_hash = encode(sha256($1::bytea), 'hex')`, token).Scan(&dbUserID, &dbRole) if err != nil { t.Fatalf("failed to query refresh_tokens: %v", err) } if dbUserID != userID { t.Errorf("expected user_id %q, got %q", userID, dbUserID) } if dbRole != "verified_email" { t.Errorf("expected role 'verified_email', got %q", dbRole) } } // ============================================================================= // VerifyRefreshToken Tests // ============================================================================= // TestVerifyRefreshToken_Success generates a refresh token, verifies it, and // asserts the returned userID and role match. Then confirms the token was // consumed (second call fails with "invalid or expired"). func TestVerifyRefreshToken_Success(t *testing.T) { ctx, tx := testtx.SetupTestTx(t) userID, err := fixtures.CreateTestUser(tx) if err != nil { t.Fatalf("failed to create test user: %v", err) } token, err := GenerateRefreshToken(ctx, userID, "verified_email") if err != nil { t.Fatalf("GenerateRefreshToken() failed: %v", err) } // First verify should succeed retUserID, retRole, err := VerifyRefreshToken(ctx, token) if err != nil { t.Fatalf("VerifyRefreshToken() failed: %v", err) } if retUserID != userID { t.Errorf("expected user_id %q, got %q", userID, retUserID) } if retRole != "verified_email" { t.Errorf("expected role 'verified_email', got %q", retRole) } // Second verify with same token must fail (rotation — token consumed) _, _, err = VerifyRefreshToken(ctx, token) if err == nil { t.Fatal("expected error for consumed token, got nil") } if !strings.Contains(err.Error(), "invalid or expired") { t.Errorf("expected 'invalid or expired' error, got: %v", err) } } // TestVerifyRefreshToken_Rotation verifies the token rotation mechanism: // first call succeeds, second call with the same token fails. func TestVerifyRefreshToken_Rotation(t *testing.T) { ctx, tx := testtx.SetupTestTx(t) userID, err := fixtures.CreateTestUser(tx) if err != nil { t.Fatalf("failed to create test user: %v", err) } token, err := GenerateRefreshToken(ctx, userID, "verified_email") if err != nil { t.Fatalf("GenerateRefreshToken() failed: %v", err) } // First call should succeed _, _, err = VerifyRefreshToken(ctx, token) if err != nil { t.Fatalf("first verification should succeed, got: %v", err) } // Second call with the same token must fail _, _, err = VerifyRefreshToken(ctx, token) if err == nil { t.Fatal("expected error for rotated token, got nil") } if !strings.Contains(err.Error(), "invalid or expired") { t.Errorf("expected 'invalid or expired' error, got: %v", err) } } // TestVerifyRefreshToken_InvalidToken calls VerifyRefreshToken with a fake // token string and expects it to fail with "invalid or expired". func TestVerifyRefreshToken_InvalidToken(t *testing.T) { ctx, _ := testtx.SetupTestTx(t) _, _, err := VerifyRefreshToken(ctx, "this-is-a-completely-fake-token-string") if err == nil { t.Fatal("expected error for invalid token, got nil") } if !strings.Contains(err.Error(), "invalid or expired") { t.Errorf("expected 'invalid or expired' error, got: %v", err) } } // TestVerifyToken_EmptyJTI creates a token with an empty "jti" claim and // verifies that VerifyToken returns an error containing "invalid jti claim". func TestVerifyToken_EmptyJTI(t *testing.T) { _, tokenStr, err := TokenAuth.Encode(map[string]any{ "user_id": "user-test", "role": "verified_email", "jti": "", "exp": clock.Now().Add(1 * time.Hour).Unix(), }) require.NoError(t, err) _, _, _, err = VerifyToken(tokenStr, context.Background()) require.Error(t, err) assert.Contains(t, err.Error(), "invalid jti claim") } // ============================================================================= // Nil db.Conn Tests // ============================================================================= // TestRevokeJTI_NilConn verifies that RevokeJTI does not panic when db.Conn is nil. func TestRevokeJTI_NilConn(t *testing.T) { savedConn := db.Conn db.Conn = nil t.Cleanup(func() { db.Conn = savedConn }) // Should not panic when db.Conn is nil RevokeJTI(context.Background(), "test-jti", time.Now()) } // TestIsJTIRevoked_NilConn verifies that IsJTIRevoked returns false when db.Conn is nil. func TestIsJTIRevoked_NilConn(t *testing.T) { savedConn := db.Conn db.Conn = nil t.Cleanup(func() { db.Conn = savedConn }) if IsJTIRevoked(context.Background(), "test-jti") { t.Error("expected IsJTIRevoked to return false when db.Conn is nil") } } // TestCleanupRevokedJTIs_NilConn verifies CleanupRevokedJTIs returns (0, nil) when db.Conn is nil. func TestCleanupRevokedJTIs_NilConn(t *testing.T) { savedConn := db.Conn db.Conn = nil t.Cleanup(func() { db.Conn = savedConn }) n, err := CleanupRevokedJTIs(context.Background()) if err != nil { t.Errorf("expected no error, got: %v", err) } if n != 0 { t.Errorf("expected 0 rows, got %d", n) } } // Note: crypto error paths in generateJTI() and generateRefreshTokenString() // are unreachable on Go 1.26+ because crypto/rand.Read() calls runtime.fatal() // instead of returning an error (see https://go.dev/issue/66821). The error // return exists for backward compatibility with older Go versions. // ============================================================================= // VerifyToken Edge Case Tests // ============================================================================= // TestVerifyToken_MissingUserID creates a token without user_id claim and verifies // VerifyToken returns an error containing "invalid user_id claim". func TestVerifyToken_MissingUserID(t *testing.T) { _, tokenString, err := TokenAuth.Encode(map[string]any{ "role": "admin", "jti": "test-jti-001", "exp": clock.Now().Add(1 * time.Hour).Unix(), }) if err != nil { t.Fatalf("failed to encode token: %v", err) } _, _, _, err = VerifyToken(tokenString, context.Background()) if err == nil { t.Fatal("expected error for missing user_id claim, got nil") } if !strings.Contains(err.Error(), "invalid user_id claim") { t.Errorf("expected 'invalid user_id claim' error, got: %v", err) } } // TestVerifyToken_WrongUserIDType creates a token with user_id as an integer (wrong type) // and verifies VerifyToken returns an error containing "invalid user_id claim". func TestVerifyToken_WrongUserIDType(t *testing.T) { _, tokenString, err := TokenAuth.Encode(map[string]any{ "user_id": 12345, "role": "admin", "jti": "test-jti-002", "exp": clock.Now().Add(1 * time.Hour).Unix(), }) if err != nil { t.Fatalf("failed to encode token: %v", err) } _, _, _, err = VerifyToken(tokenString, context.Background()) if err == nil { t.Fatal("expected error for wrong user_id type, got nil") } if !strings.Contains(err.Error(), "invalid user_id claim") { t.Errorf("expected 'invalid user_id claim' error, got: %v", err) } } // TestVerifyToken_MissingRole creates a token without role claim and verifies // VerifyToken returns an error containing "invalid role claim". func TestVerifyToken_MissingRole(t *testing.T) { _, tokenString, err := TokenAuth.Encode(map[string]any{ "user_id": "user-001", "jti": "test-jti-003", "exp": clock.Now().Add(1 * time.Hour).Unix(), }) if err != nil { t.Fatalf("failed to encode token: %v", err) } _, _, _, err = VerifyToken(tokenString, context.Background()) if err == nil { t.Fatal("expected error for missing role claim, got nil") } if !strings.Contains(err.Error(), "invalid role claim") { t.Errorf("expected 'invalid role claim' error, got: %v", err) } } // TestVerifyToken_WrongRoleType creates a token with role as an integer (wrong type) // and verifies VerifyToken returns an error containing "invalid role claim". func TestVerifyToken_WrongRoleType(t *testing.T) { _, tokenString, err := TokenAuth.Encode(map[string]any{ "user_id": "user-001", "role": 12345, "jti": "test-jti-004", "exp": clock.Now().Add(1 * time.Hour).Unix(), }) if err != nil { t.Fatalf("failed to encode token: %v", err) } _, _, _, err = VerifyToken(tokenString, context.Background()) if err == nil { t.Fatal("expected error for wrong role type, got nil") } if !strings.Contains(err.Error(), "invalid role claim") { t.Errorf("expected 'invalid role claim' error, got: %v", err) } }