diff --git a/LICENSE b/LICENSE index c71f04c..e47911e 100644 --- a/LICENSE +++ b/LICENSE @@ -1,6 +1,6 @@ BSD 3-Clause License -Copyright (c) 2025, Lixen Wraith +Copyright (c) 2026, Lixen Wraith Redistribution and use in source and binary forms, with or without modification, are permitted provided that the following conditions are met: diff --git a/README.md b/README.md index 56c6e68..bd2a772 100644 --- a/README.md +++ b/README.md @@ -53,5 +53,14 @@ server.AddCredential(cred) ## Testing ```bash -go test -v ./ +go test ./... -race -count=1 +go test ./... -run '^$' -bench . -benchmem + +# fuzz targets (run individually) +go test -run '^$' -fuzz FuzzParsePHC -fuzztime 60s +go test -run '^$' -fuzz FuzzVerifyPassword -fuzztime 60s +go test -run '^$' -fuzz FuzzImportCredential -fuzztime 60s +go test -run '^$' -fuzz FuzzValidateHS256Token -fuzztime 60s +go test -run '^$' -fuzz FuzzParseBasicAuth -fuzztime 30s ``` + diff --git a/argon2.go b/argon2.go index 50089c7..7e0c900 100644 --- a/argon2.go +++ b/argon2.go @@ -1,9 +1,9 @@ package auth import ( - "crypto/hmac" "crypto/rand" "crypto/sha256" + "crypto/subtle" "encoding/base64" "fmt" "strings" @@ -25,6 +25,17 @@ const ( MaxPHCHashLen = 256 ) +// Execution budget for KDF parameters taken from an encoded record. +// parsePHC accepts the full PHC range (m <= 4 GiB, t <= 1000, p <= 255) because +// those values are well-formed; running them is a different decision. One +// crafted record would otherwise allocate 4 GiB. Applies to any Argon2 run +// whose parameters were chosen by a peer rather than by this process. +const ( + MaxVerifyArgonMemory = 256 * 1024 // KiB + MaxVerifyArgonTime = 16 + MaxVerifyArgonThreads = 16 +) + // argonParams holds configurable Argon2id parameters type argonParams struct { time uint32 @@ -86,9 +97,7 @@ func HashPassword(password string, opts ...Option) (string, error) { } salt := make([]byte, params.saltLen) - if _, err := rand.Read(salt); err != nil { - return "", fmt.Errorf("%w: %v", ErrSaltGenerationFailed, err) - } + rand.Read(salt) // cryptographically secure random bytes hash := argon2.IDKey([]byte(password), salt, params.time, params.memory, params.threads, params.keyLen) @@ -139,117 +148,128 @@ func credentialFromSaltedPassword(username string, saltedPassword, salt []byte, // PHC format for Argon2id. It validates structure, parameters, and encoding, // but does not verify a password against the hash. func ValidatePHCHashFormat(phcHash string) error { - // Cap total input before any splitting or base64 decoding - if len(phcHash) > MaxPHCHashLen { - return fmt.Errorf("%w: encoded hash exceeds %d bytes", ErrPHCInvalidFormat, MaxPHCHashLen) - } - - parts := strings.Split(phcHash, "$") - if len(parts) != 6 { - return fmt.Errorf("%w: expected 6 parts, got %d", ErrPHCInvalidFormat, len(parts)) - } - - // Validate empty parts[0] (PHC format starts with $) - if parts[0] != "" { - return fmt.Errorf("%w: hash must start with $", ErrPHCInvalidFormat) - } - - // Validate algorithm identifier - if parts[1] != "argon2id" { - return fmt.Errorf("%w: unsupported algorithm %q, expected argon2id", ErrPHCInvalidFormat, parts[1]) - } - - // Validate version - var version int - n, err := fmt.Sscanf(parts[2], "v=%d", &version) - if err != nil || n != 1 { - return fmt.Errorf("%w: invalid version format", ErrPHCInvalidFormat) - } - if version != argon2.Version { - return fmt.Errorf("%w: unsupported version %d, expected %d", ErrPHCInvalidFormat, version, argon2.Version) - } - - // Validate parameters - var memory, time uint32 - var threads uint8 - n, err = fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads) - if err != nil || n != 3 { - return fmt.Errorf("%w: failed to parse parameters", ErrPHCInvalidFormat) - } - - // Validate parameter ranges - if time == 0 || memory == 0 || threads == 0 { - return fmt.Errorf("%w: parameters must be non-zero", ErrPHCInvalidFormat) - } - if memory > 4*1024*1024 { // 4GB limit - return fmt.Errorf("%w: memory parameter exceeds maximum (4GB)", ErrPHCInvalidFormat) - } - if time > 1000 { // Reasonable upper bound - return fmt.Errorf("%w: time parameter exceeds maximum (1000)", ErrPHCInvalidFormat) - } - if threads > 255 { // uint8 max, but practically much lower - return fmt.Errorf("%w: threads parameter exceeds maximum (255)", ErrPHCInvalidFormat) - } - - // Validate salt encoding - salt, err := base64.RawStdEncoding.DecodeString(parts[4]) - if err != nil { - return fmt.Errorf("%w: %v", ErrPHCInvalidSalt, err) - } - if len(salt) < 8 { // Minimum safe salt length - return fmt.Errorf("%w: salt too short (%d bytes)", ErrPHCInvalidSalt, len(salt)) - } - if len(salt) > MaxArgonSaltLen { - return fmt.Errorf("%w: salt too long (%d bytes)", ErrPHCInvalidSalt, len(salt)) - } - - // Validate hash encoding - hash, err := base64.RawStdEncoding.DecodeString(parts[5]) - if err != nil { - return fmt.Errorf("%w: %v", ErrPHCInvalidHash, err) - } - if len(hash) < 16 { // Minimum hash length - return fmt.Errorf("%w: hash too short (%d bytes)", ErrPHCInvalidHash, len(hash)) - } - if len(hash) > MaxArgonKeyLen { - return fmt.Errorf("%w: hash too long (%d bytes)", ErrPHCInvalidHash, len(hash)) - } - - return nil + _, err := parsePHC(phcHash) + return err } // parsed + verified PHC material, reused to avoid a second KDF pass type phcResult struct { - derived []byte // argon2.IDKey output; == SCRAM salted password when len == DefaultArgonKeyLen - salt []byte - time uint32 - memory uint32 - threads uint8 + derived []byte + expectedHash []byte + salt []byte + time uint32 + memory uint32 + threads uint8 +} + +// verifyPHC validates format, bounds the password, runs the KDF once, +// and constant-time compares against the encoded digest. +func parsePHC(phcHash string) (*phcResult, error) { + if len(phcHash) > MaxPHCHashLen { + return nil, fmt.Errorf("%w: encoded hash exceeds %d bytes", ErrPHCInvalidFormat, MaxPHCHashLen) + } + + parts := strings.Split(phcHash, "$") + if len(parts) != 6 { + return nil, fmt.Errorf("%w: expected 6 parts, got %d", ErrPHCInvalidFormat, len(parts)) + } + if parts[0] != "" { + return nil, fmt.Errorf("%w: hash must start with $", ErrPHCInvalidFormat) + } + if parts[1] != "argon2id" { + return nil, fmt.Errorf("%w: unsupported algorithm %q, expected argon2id", ErrPHCInvalidFormat, parts[1]) + } + + var version int + if _, err := fmt.Sscanf(parts[2], "v=%d", &version); err != nil || parts[2] != fmt.Sprintf("v=%d", version) { + return nil, fmt.Errorf("%w: invalid version format", ErrPHCInvalidFormat) + } + if version != argon2.Version { + return nil, fmt.Errorf("%w: unsupported version %d, expected %d", ErrPHCInvalidFormat, version, argon2.Version) + } + + var memory, time uint32 + var threads uint8 + if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads); err != nil || parts[3] != fmt.Sprintf("m=%d,t=%d,p=%d", memory, time, threads) { + return nil, fmt.Errorf("%w: failed to parse parameters", ErrPHCInvalidFormat) + } + + if time == 0 || memory == 0 || threads == 0 { + return nil, fmt.Errorf("%w: parameters must be non-zero", ErrPHCInvalidFormat) + } + if memory > 4*1024*1024 { + return nil, fmt.Errorf("%w: memory parameter exceeds maximum (4GB)", ErrPHCInvalidFormat) + } + if time > 1000 { + return nil, fmt.Errorf("%w: time parameter exceeds maximum (1000)", ErrPHCInvalidFormat) + } + if threads > 255 { + return nil, fmt.Errorf("%w: threads parameter exceeds maximum (255)", ErrPHCInvalidFormat) + } + + salt, err := base64.RawStdEncoding.DecodeString(parts[4]) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrPHCInvalidSalt, err) + } + if len(salt) < 8 { + return nil, fmt.Errorf("%w: salt too short (%d bytes)", ErrPHCInvalidSalt, len(salt)) + } + if len(salt) > MaxArgonSaltLen { + return nil, fmt.Errorf("%w: salt too long (%d bytes)", ErrPHCInvalidSalt, len(salt)) + } + + hash, err := base64.RawStdEncoding.DecodeString(parts[5]) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrPHCInvalidHash, err) + } + if len(hash) < 16 { + return nil, fmt.Errorf("%w: hash too short (%d bytes)", ErrPHCInvalidHash, len(hash)) + } + if len(hash) > MaxArgonKeyLen { + return nil, fmt.Errorf("%w: hash too long (%d bytes)", ErrPHCInvalidHash, len(hash)) + } + + return &phcResult{ + expectedHash: hash, + salt: salt, + time: time, + memory: memory, + threads: threads, + }, nil } // verifyPHC validates format, bounds the password, runs the KDF once, and // constant-time compares against the encoded digest. func verifyPHC(password, phcHash string) (*phcResult, error) { - if err := ValidatePHCHashFormat(phcHash); err != nil { - return nil, err - } if len(password) > MaxPasswordLen { return nil, ErrPasswordTooLong } - parts := strings.Split(phcHash, "$") + r, err := parsePHC(phcHash) + if err != nil { + return nil, err + } - r := &phcResult{} - // Parse is guaranteed well-formed by ValidatePHCHashFormat above. - fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &r.memory, &r.time, &r.threads) + // Bound the KDF before it runs; the record is untrusted input + if err := checkArgonCost(r.memory, r.time, r.threads); err != nil { + return nil, err + } - // Encodings validated above; errors are unreachable. - r.salt, _ = base64.RawStdEncoding.DecodeString(parts[4]) - expected, _ := base64.RawStdEncoding.DecodeString(parts[5]) - - r.derived = argon2.IDKey([]byte(password), r.salt, r.time, r.memory, r.threads, uint32(len(expected))) - if !hmac.Equal(r.derived, expected) { + r.derived = argon2.IDKey([]byte(password), r.salt, r.time, r.memory, r.threads, uint32(len(r.expectedHash))) + if subtle.ConstantTimeCompare(r.derived, r.expectedHash) != 1 { return nil, ErrInvalidCredentials } return r, nil } + +func checkArgonCost(memory, time uint32, threads uint8) error { + switch { + case memory > MaxVerifyArgonMemory: + return fmt.Errorf("%w: memory %d KiB exceeds %d", ErrPHCCostTooHigh, memory, MaxVerifyArgonMemory) + case time > MaxVerifyArgonTime: + return fmt.Errorf("%w: time %d exceeds %d", ErrPHCCostTooHigh, time, MaxVerifyArgonTime) + case threads > MaxVerifyArgonThreads: + return fmt.Errorf("%w: threads %d exceeds %d", ErrPHCCostTooHigh, threads, MaxVerifyArgonThreads) + } + return nil +} diff --git a/argon2_test.go b/argon2_test.go index 3c0c00a..50392f3 100644 --- a/argon2_test.go +++ b/argon2_test.go @@ -1,192 +1,427 @@ package auth import ( + "bytes" + "crypto/sha256" "encoding/base64" + "fmt" "strings" "sync" "testing" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" + "golang.org/x/crypto/argon2" ) -func TestPasswordHashing(t *testing.T) { - password := "testPassword123" +func TestHashPasswordEncoding(t *testing.T) { + const pw = "testPassword123" + hash, err := HashPassword(pw) + noErr(t, err, "HashPassword") - // Test hashing with default parameters - hash, err := HashPassword(password) - require.NoError(t, err, "Failed to hash password") + parts := strings.Split(hash, "$") + eq(t, len(parts), 6, "field count") + eq(t, parts[0], "", "leading field") + eq(t, parts[1], "argon2id", "algorithm") + eq(t, parts[2], fmt.Sprintf("v=%d", argon2.Version), "version") + eq(t, parts[3], fmt.Sprintf("m=%d,t=%d,p=%d", DefaultArgonMemory, DefaultArgonTime, DefaultArgonThreads), "parameters") - // Verify PHC format - assert.True(t, strings.HasPrefix(hash, "$argon2id$"), - "Hash should have argon2id prefix, got: %s", hash) + salt, err := base64.RawStdEncoding.DecodeString(parts[4]) + noErr(t, err, "salt decode") + eq(t, len(salt), DefaultArgonSaltLen, "salt length") - // Test verification with correct password - err = VerifyPassword(password, hash) - assert.NoError(t, err, "Failed to verify correct password") + digest, err := base64.RawStdEncoding.DecodeString(parts[5]) + noErr(t, err, "digest decode") + eq(t, len(digest), DefaultArgonKeyLen, "digest length") - // Test verification with incorrect password - err = VerifyPassword("wrongPassword", hash) - assert.Error(t, err, "Verification should fail for incorrect password") - assert.Equal(t, ErrInvalidCredentials, err) + // the encoded digest must be reproducible from the encoded material + want := argon2.IDKey([]byte(pw), salt, DefaultArgonTime, DefaultArgonMemory, DefaultArgonThreads, DefaultArgonKeyLen) + eqBytes(t, digest, want, "digest") - // Test weak password - _, err = HashPassword("weak") - assert.Equal(t, ErrWeakPassword, err, "Should reject weak password") + isTrue(t, len(hash) <= MaxPHCHashLen, "default encoding within MaxPHCHashLen") + noErr(t, ValidatePHCHashFormat(hash), "self-validation") +} - // Test with custom options - hash, err = HashPassword(password, - WithTime(5), - WithMemory(128*1024), - WithThreads(8)) - require.NoError(t, err) - - err = VerifyPassword(password, hash) - assert.NoError(t, err) - - // Test malformed PHC hash - err = VerifyPassword(password, "$invalid$format") - assert.Error(t, err, "Should reject malformed hash") - - // Test corrupted salt - corruptedHash := strings.Replace(hash, "$argon2id$", "$argon2id$", 1) - parts := strings.Split(corruptedHash, "$") - if len(parts) == 6 { - parts[4] = "invalid!base64" - corruptedHash = strings.Join(parts, "$") - err = VerifyPassword(password, corruptedHash) - assert.Error(t, err, "Should reject corrupted salt") +func TestHashPasswordSaltUniqueness(t *testing.T) { + seen := make(map[string]struct{}, 64) + for range 64 { + h, err := HashPassword("testPassword123", cheapArgon...) + noErr(t, err, "HashPassword") + salt := strings.Split(h, "$")[4] + if _, dup := seen[salt]; dup { + t.Fatalf("duplicate salt: %s", salt) + } + seen[salt] = struct{}{} } } -func TestEmptyPasswordAfterValidation(t *testing.T) { - // Empty password should be rejected by length check +func TestVerifyPassword(t *testing.T) { + const pw = "testPassword123" + hash, err := HashPassword(pw, cheapArgon...) + noErr(t, err, "HashPassword") + + noErr(t, VerifyPassword(pw, hash), "correct password") + noErr(t, VerifyPassword(pw, hash), "repeat verification") + + for _, wrong := range []string{"wrongPassword", "", pw + "\x00", pw + " ", strings.ToUpper(pw)} { + errIs(t, VerifyPassword(wrong, hash), ErrInvalidCredentials, fmt.Sprintf("password %q", wrong)) + } + + // verification has no minimum length: legacy hashes must stay verifiable + weak := phcFor("short", []byte("0123456789abcdef"), DefaultArgonKeyLen) + noErr(t, VerifyPassword("short", weak), "sub-minimum password verifies") +} + +func TestPasswordLengthBounds(t *testing.T) { _, err := HashPassword("") - assert.Equal(t, ErrWeakPassword, err) + errIs(t, err, ErrWeakPassword, "empty") + _, err = HashPassword("1234567") + errIs(t, err, ErrWeakPassword, "seven bytes") - // 8-character password should pass - hash, err := HashPassword("12345678") - require.NoError(t, err) + h8, err := HashPassword("12345678", cheapArgon...) + noErr(t, err, "eight bytes") + noErr(t, VerifyPassword("12345678", h8), "verify eight bytes") - err = VerifyPassword("12345678", hash) - assert.NoError(t, err) + // the minimum is measured in bytes, not runes + const multibyte = "ünïcödé" + eq(t, len([]rune(multibyte)), 7, "rune count") + eq(t, len(multibyte), 11, "byte count") + hm, err := HashPassword(multibyte, cheapArgon...) + noErr(t, err, "multibyte password") + noErr(t, VerifyPassword(multibyte, hm), "verify multibyte") + + maxPw := strings.Repeat("a", MaxPasswordLen) + hMax, err := HashPassword(maxPw, cheapArgon...) + noErr(t, err, "maximum length") + noErr(t, VerifyPassword(maxPw, hMax), "verify maximum length") + + _, err = HashPassword(maxPw+"a", cheapArgon...) + errIs(t, err, ErrPasswordTooLong, "over maximum") + // rejected before the KDF runs + errIs(t, VerifyPassword(maxPw+"a", hMax), ErrPasswordTooLong, "verify over maximum") } -func TestConcurrentPasswordOperations(t *testing.T) { - password := "testPassword123" - hash, err := HashPassword(password) - require.NoError(t, err) +func TestHashPasswordOptions(t *testing.T) { + const pw = "testPassword123" - // Test concurrent verification - var wg sync.WaitGroup - for i := 0; i < 10; i++ { - wg.Add(1) - go func() { - defer wg.Done() - err := VerifyPassword(password, hash) - assert.NoError(t, err) - }() + h, err := HashPassword(pw, WithTime(2), WithMemory(16*1024), WithThreads(2)) + noErr(t, err, "custom parameters") + eq(t, strings.Split(h, "$")[3], "m=16384,t=2,p=2", "encoded parameters") + noErr(t, VerifyPassword(pw, h), "verify custom parameters") + + // zero values are discarded by the option guards + h, err = HashPassword(pw, WithMemory(testArgonMemory), WithTime(0), WithThreads(0)) + noErr(t, err, "zero-valued options") + eq(t, strings.Split(h, "$")[3], + fmt.Sprintf("m=%d,t=%d,p=%d", testArgonMemory, DefaultArgonTime, DefaultArgonThreads), + "defaults retained") + + // the last option wins + h, err = HashPassword(pw, WithMemory(64*1024), WithMemory(testArgonMemory), WithTime(1), WithThreads(1)) + noErr(t, err, "repeated option") + eq(t, strings.Split(h, "$")[3], fmt.Sprintf("m=%d,t=1,p=1", testArgonMemory), "last option applied") +} + +func TestVerifyPasswordNonStandardDigestLength(t *testing.T) { + // Records produced elsewhere may encode digests other than 32 bytes; + // verification must derive at the encoded length. + const pw = "testPassword123" + salt := []byte("0123456789abcdef0123") + + for _, keyLen := range []uint32{16, 20, 32, MaxArgonKeyLen} { + ctx := fmt.Sprintf("keyLen=%d", keyLen) + h := phcFor(pw, salt, keyLen) + noErr(t, ValidatePHCHashFormat(h), "format "+ctx) + noErr(t, VerifyPassword(pw, h), "verify "+ctx) + errIs(t, VerifyPassword("wrongPassword", h), ErrInvalidCredentials, "wrong password "+ctx) } - wg.Wait() -} - -func TestPHCMigration(t *testing.T) { - password := "testPassword123" - username := "migrationUser" - - // Generate PHC hash - phcHash, err := HashPassword(password) - require.NoError(t, err) - - // Migrate to SCRAM credential - cred, err := MigrateFromPHC(username, password, phcHash) - require.NoError(t, err) - assert.Equal(t, username, cred.Username) - assert.NotNil(t, cred.StoredKey) - assert.NotNil(t, cred.ServerKey) - - // Test with wrong password - _, err = MigrateFromPHC(username, "wrongPassword", phcHash) - assert.Equal(t, ErrInvalidCredentials, err) - - // Test with invalid PHC format - _, err = MigrateFromPHC(username, password, "$invalid$format") - assert.Error(t, err) } func TestValidatePHCHashFormat(t *testing.T) { - // Generate valid hash for testing - validHash, err := HashPassword("testPassword123") - require.NoError(t, err) + ver := fmt.Sprintf("v=%d", argon2.Version) + const params = "m=65536,t=3,p=4" + okSalt, okDigest := rawB64(16), rawB64(32) - // Test valid hash - err = ValidatePHCHashFormat(validHash) - assert.NoError(t, err, "Valid hash should pass validation") + build := func(alg, version, prm, salt, digest string) string { + return "$" + alg + "$" + version + "$" + prm + "$" + salt + "$" + digest + } + std := func(prm, salt, digest string) string { return build("argon2id", ver, prm, salt, digest) } - // Test malformed formats - testCases := []struct { - name string - hash string - wantErr error + generated, err := HashPassword("testPassword123", cheapArgon...) + noErr(t, err, "HashPassword") + + // largest structurally valid record: proves MaxPHCHashLen never binds + maxRecord := std("m=4194304,t=1000,p=255", rawB64(MaxArgonSaltLen), rawB64(MaxArgonKeyLen)) + isTrue(t, len(maxRecord) <= MaxPHCHashLen, "maximum record within length cap") + + cases := []struct { + name string + hash string + want error }{ + {"generated", generated, nil}, + {"minimal", std(params, okSalt, okDigest), nil}, + {"maximum record", maxRecord, nil}, + {"empty", "", ErrPHCInvalidFormat}, - {"not PHC format", "plaintext", ErrPHCInvalidFormat}, - {"wrong prefix", "argon2id$v=19$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, - {"wrong algorithm", "$bcrypt$v=19$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, - {"missing version", "$argon2id$$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, - {"wrong version", "$argon2id$v=1$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, - {"missing params", "$argon2id$v=19$$salt$hash", ErrPHCInvalidFormat}, - {"invalid params format", "$argon2id$v=19$invalid$salt$hash", ErrPHCInvalidFormat}, - {"zero time", "$argon2id$v=19$m=65536,t=0,p=4$salt$hash", ErrPHCInvalidFormat}, - {"zero memory", "$argon2id$v=19$m=0,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, - {"zero threads", "$argon2id$v=19$m=65536,t=3,p=0$salt$hash", ErrPHCInvalidFormat}, - {"excessive memory", "$argon2id$v=19$m=5000000,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, - {"excessive time", "$argon2id$v=19$m=65536,t=2000,p=4$salt$hash", ErrPHCInvalidFormat}, - {"invalid salt encoding", "$argon2id$v=19$m=65536,t=3,p=4$!!!invalid!!!$hash", ErrPHCInvalidSalt}, - {"invalid hash encoding", "$argon2id$v=19$m=65536,t=3,p=4$" + - base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$!!!invalid!!!", ErrPHCInvalidHash}, - {"short salt", "$argon2id$v=19$m=65536,t=3,p=4$" + - base64.RawStdEncoding.EncodeToString([]byte("short")) + "$" + - base64.RawStdEncoding.EncodeToString([]byte("hash1234567890123456")), ErrPHCInvalidSalt}, - {"short hash", "$argon2id$v=19$m=65536,t=3,p=4$" + - base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$" + - base64.RawStdEncoding.EncodeToString([]byte("short")), ErrPHCInvalidHash}, - {"too few parts", "$argon2id$v=19$m=65536,t=3,p=4", ErrPHCInvalidFormat}, - {"too many parts", "$argon2id$v=19$m=65536,t=3,p=4$salt$hash$extra", ErrPHCInvalidFormat}, - {"oversized salt", "$argon2id$v=19$m=65536,t=3,p=4$" + - base64.RawStdEncoding.EncodeToString(make([]byte, 128)) + "$" + - base64.RawStdEncoding.EncodeToString([]byte("hash1234567890123456")), ErrPHCInvalidSalt}, - {"oversized hash", "$argon2id$v=19$m=65536,t=3,p=4$" + - base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$" + - base64.RawStdEncoding.EncodeToString(make([]byte, 128)), ErrPHCInvalidHash}, - {"oversized input", "$argon2id$v=19$m=65536,t=3,p=4$" + - strings.Repeat("A", 512) + "$hash", ErrPHCInvalidFormat}, + {"plaintext", "plaintext", ErrPHCInvalidFormat}, + {"missing leading separator", "argon2id$" + ver + "$" + params + "$" + okSalt + "$" + okDigest, ErrPHCInvalidFormat}, + {"leading garbage", "x" + std(params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"leading space", " " + std(params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"too few fields", "$argon2id$" + ver + "$" + params, ErrPHCInvalidFormat}, + {"too many fields", std(params, okSalt, okDigest) + "$extra", ErrPHCInvalidFormat}, + {"over length cap", strings.Repeat("A", MaxPHCHashLen+1), ErrPHCInvalidFormat}, + + {"algorithm bcrypt", build("bcrypt", ver, params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"algorithm argon2i", build("argon2i", ver, params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"algorithm argon2d", build("argon2d", ver, params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"algorithm case", build("ARGON2ID", ver, params, okSalt, okDigest), ErrPHCInvalidFormat}, + + {"version empty", build("argon2id", "", params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"version malformed", build("argon2id", "version=19", params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"version leading zero", build("argon2id", "v=019", params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"version negative", build("argon2id", "v=-19", params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"version too low", build("argon2id", "v=18", params, okSalt, okDigest), ErrPHCInvalidFormat}, + {"version too high", build("argon2id", "v=20", params, okSalt, okDigest), ErrPHCInvalidFormat}, + + {"parameters empty", std("", okSalt, okDigest), ErrPHCInvalidFormat}, + {"parameters reordered", std("t=3,m=65536,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"parameters spaced", std("m=65536, t=3, p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"parameters trailing", std("m=65536,t=3,p=4,x=1", okSalt, okDigest), ErrPHCInvalidFormat}, + {"parameters negative", std("m=-1,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"parameters overflow", std("m=99999999999999999999,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"zero time", std("m=65536,t=0,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"zero memory", std("m=0,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"zero threads", std("m=65536,t=3,p=0", okSalt, okDigest), ErrPHCInvalidFormat}, + {"memory at cap", std("m=4194304,t=3,p=4", okSalt, okDigest), nil}, + {"memory over cap", std("m=4194305,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"time at cap", std("m=65536,t=1000,p=4", okSalt, okDigest), nil}, + {"time over cap", std("m=65536,t=1001,p=4", okSalt, okDigest), ErrPHCInvalidFormat}, + {"threads at cap", std("m=65536,t=3,p=255", okSalt, okDigest), nil}, + {"threads overflow", std("m=65536,t=3,p=256", okSalt, okDigest), ErrPHCInvalidFormat}, + + {"salt not base64", std(params, "!!!invalid!!!", okDigest), ErrPHCInvalidSalt}, + {"salt padded", std(params, base64.StdEncoding.EncodeToString(make([]byte, 10)), okDigest), ErrPHCInvalidSalt}, + {"salt url alphabet", std(params, "abc-def_ghijklmn", okDigest), ErrPHCInvalidSalt}, + {"salt empty", std(params, "", okDigest), ErrPHCInvalidSalt}, + {"salt too short", std(params, rawB64(7), okDigest), ErrPHCInvalidSalt}, + {"salt minimum", std(params, rawB64(8), okDigest), nil}, + {"salt at cap", std(params, rawB64(MaxArgonSaltLen), okDigest), nil}, + {"salt over cap", std(params, rawB64(MaxArgonSaltLen+1), okDigest), ErrPHCInvalidSalt}, + + {"digest not base64", std(params, okSalt, "!!!invalid!!!"), ErrPHCInvalidHash}, + {"digest empty", std(params, okSalt, ""), ErrPHCInvalidHash}, + {"digest too short", std(params, okSalt, rawB64(15)), ErrPHCInvalidHash}, + {"digest minimum", std(params, okSalt, rawB64(16)), nil}, + {"digest at cap", std(params, okSalt, rawB64(MaxArgonKeyLen)), nil}, + {"digest over cap", std(params, okSalt, rawB64(MaxArgonKeyLen+1)), ErrPHCInvalidHash}, + + // encoding/base64 discards CR and LF, so PHC records are not + // byte-canonical: never compare them as strings. + {"digest with newline", std(params, okSalt, okDigest[:20]+"\n"+okDigest[20:]), nil}, } - for _, tc := range testCases { + for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { err := ValidatePHCHashFormat(tc.hash) - assert.ErrorIs(t, err, tc.wantErr, "Test case: %s", tc.name) + if tc.want == nil { + noErr(t, err, tc.name) + return + } + errIs(t, err, tc.want, tc.name) }) } - - // Test that validation doesn't require password - err = ValidatePHCHashFormat(validHash) - assert.NoError(t, err, "Should validate format without password") - - // Verify that a validated hash can still be used for verification - err = ValidatePHCHashFormat(validHash) - require.NoError(t, err) - err = VerifyPassword("testPassword123", validHash) - assert.NoError(t, err, "Validated hash should still work for password verification") } -func TestVerifyPassword_MalformedParamsNoPanic(t *testing.T) { +func TestVerifyPasswordMalformedNoPanic(t *testing.T) { for _, h := range []string{ - "$argon2id$v=19$m=65536,t=0,p=4$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA", - "$argon2id$v=19$m=65536,t=3,p=0$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA", - "$argon2id$v=19$garbage$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA", + "", "$", "$$$$$", "$argon2id$", strings.Repeat("$", 1000), + "$argon2id$v=19$m=65536,t=0,p=4$" + rawB64(16) + "$" + rawB64(32), + "$argon2id$v=19$m=65536,t=3,p=0$" + rawB64(16) + "$" + rawB64(32), + "$argon2id$v=19$garbage$" + rawB64(16) + "$" + rawB64(32), + "$argon2id$v=19$m=65536,t=3,p=4$$", + "$argon2id$v=19$m=99999999999999999999,t=3,p=4$" + rawB64(16) + "$" + rawB64(32), } { - assert.Error(t, VerifyPassword("whatever", h)) + hasErr(t, VerifyPassword("whatever", h), fmt.Sprintf("hash %q", h)) } } + +func TestMigrateFromPHC(t *testing.T) { + const user, pw = "migrationUser", "testPassword123" + phcHash, err := HashPassword(pw, cheapArgon...) + noErr(t, err, "HashPassword") + + cred, err := MigrateFromPHC(user, pw, phcHash) + noErr(t, err, "MigrateFromPHC") + eq(t, cred.Username, user, "username") + eq(t, cred.ArgonTime, uint32(testArgonTime), "time") + eq(t, cred.ArgonMemory, uint32(testArgonMemory), "memory") + eq(t, cred.ArgonThreads, uint8(testArgonThreads), "threads") + eq(t, len(cred.StoredKey), sha256.Size, "stored key length") + eq(t, len(cred.ServerKey), sha256.Size, "server key length") + + salt, err := base64.RawStdEncoding.DecodeString(strings.Split(phcHash, "$")[4]) + noErr(t, err, "salt decode") + eqBytes(t, cred.Salt, salt, "salt carried over") + + // migration must agree with direct derivation over the same material + direct, err := DeriveCredential(user, pw, cred.Salt, cred.ArgonTime, cred.ArgonMemory, cred.ArgonThreads) + noErr(t, err, "DeriveCredential") + eqBytes(t, cred.StoredKey, direct.StoredKey, "stored key") + eqBytes(t, cred.ServerKey, direct.ServerKey, "server key") + + // keys are derived from the salted password, not copies of it + salted, err := base64.RawStdEncoding.DecodeString(strings.Split(phcHash, "$")[5]) + noErr(t, err, "digest decode") + want := sha256.Sum256(computeHMAC(salted, []byte("Client Key"))) + eqBytes(t, cred.StoredKey, want[:], "stored key derivation") + eqBytes(t, cred.ServerKey, computeHMAC(salted, []byte("Server Key")), "server key derivation") + if bytes.Equal(cred.StoredKey, salted) || bytes.Equal(cred.ServerKey, salted) { + t.Fatal("credential exposes the salted password") + } + + _, err = MigrateFromPHC(user, "wrongPassword", phcHash) + errIs(t, err, ErrInvalidCredentials, "wrong password") + _, err = MigrateFromPHC(user, pw, "$invalid$format") + errIs(t, err, ErrPHCInvalidFormat, "malformed record") + _, err = MigrateFromPHC(user, strings.Repeat("a", MaxPasswordLen+1), phcHash) + errIs(t, err, ErrPasswordTooLong, "oversized password") +} + +func TestMigrateFromPHCShortSalt(t *testing.T) { + // parsePHC accepts 8..64 byte salts; SCRAM requires >= 16. The 32-byte + // digest branch skips that check, the fallback branch enforces it. + // See the ‼️ note on MigrateFromPHC. + const pw = "testPassword123" + salt := []byte("12345678") + + cred, err := MigrateFromPHC("u", pw, phcFor(pw, salt, DefaultArgonKeyLen)) + noErr(t, err, "32-byte digest branch") + eq(t, len(cred.Salt), 8, "short salt retained") + + // the credential it produced cannot survive an export/import cycle + _, err = ImportCredential(cred.Export()) + errIs(t, err, ErrSCRAMSaltTooShort, "re-import") + + _, err = MigrateFromPHC("u", pw, phcFor(pw, salt, 20)) + errIs(t, err, ErrSCRAMSaltTooShort, "fallback branch") +} + +func TestConcurrentPasswordOperations(t *testing.T) { + const pw = "testPassword123" + hash, err := HashPassword(pw, cheapArgon...) + noErr(t, err, "HashPassword") + + const n = 16 + errs := make(chan error, 2*n) + var wg sync.WaitGroup + for i := range n { + wg.Add(1) + go func() { + defer wg.Done() + if err := VerifyPassword(pw, hash); err != nil { + errs <- err + } + local := fmt.Sprintf("password-%04d", i) + h, err := HashPassword(local, cheapArgon...) + if err != nil { + errs <- err + return + } + if err := VerifyPassword(local, h); err != nil { + errs <- err + } + }() + } + wg.Wait() + close(errs) + for err := range errs { + t.Errorf("concurrent operation: %v", err) + } +} + +func FuzzParsePHC(f *testing.F) { + h, err := HashPassword("testPassword123", cheapArgon...) + if err != nil { + f.Fatal(err) + } + f.Add(h) + f.Add("") + f.Add("$argon2id$v=19$m=65536,t=3,p=4$" + rawB64(16) + "$" + rawB64(32)) + f.Add("$argon2id$v=19$m=0,t=0,p=0$$") + + f.Fuzz(func(t *testing.T, s string) { + r, err := parsePHC(s) + if err != nil { + if r != nil { + t.Fatalf("result returned alongside error: %v", err) + } + return + } + if r.time == 0 || r.memory == 0 || r.threads == 0 { + t.Fatalf("zero parameter accepted: m=%d t=%d p=%d", r.memory, r.time, r.threads) + } + if r.memory > 4*1024*1024 || r.time > 1000 { + t.Fatalf("cost bound exceeded: m=%d t=%d", r.memory, r.time) + } + if len(r.salt) < 8 || len(r.salt) > MaxArgonSaltLen { + t.Fatalf("salt length %d accepted", len(r.salt)) + } + if len(r.expectedHash) < 16 || len(r.expectedHash) > MaxArgonKeyLen { + t.Fatalf("digest length %d accepted", len(r.expectedHash)) + } + }) +} + +func FuzzVerifyPassword(f *testing.F) { + h, err := HashPassword("testPassword123", cheapArgon...) + if err != nil { + f.Fatal(err) + } + f.Add("testPassword123", h) + f.Add("", "") + f.Add("x", "$argon2id$v=19$m=8192,t=1,p=1$"+rawB64(16)+"$"+rawB64(32)) + + f.Fuzz(func(t *testing.T, password, hash string) { + // parsePHC admits m up to 4 GiB from untrusted input; clamp before + // letting the fuzzer choose the KDF cost. + if r, err := parsePHC(hash); err == nil && + (r.memory > 64*1024 || r.time > 4 || r.threads > 8) { + return + } + _ = VerifyPassword(password, hash) + }) +} + +func BenchmarkHashPasswordDefault(b *testing.B) { + for b.Loop() { + if _, err := HashPassword("testPassword123"); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkVerifyPasswordDefault(b *testing.B) { + hash, err := HashPassword("testPassword123") + if err != nil { + b.Fatal(err) + } + for b.Loop() { + if err := VerifyPassword("testPassword123", hash); err != nil { + b.Fatal(err) + } + } +} + +func TestVerifyPasswordCostCeiling(t *testing.T) { + const pw = "testPassword123" + // well-formed per PHC, refused before the KDF runs + over := fmt.Sprintf("$argon2id$v=%d$m=%d,t=1,p=1$%s$%s", + argon2.Version, MaxVerifyArgonMemory+1, rawB64(16), rawB64(32)) + noErr(t, ValidatePHCHashFormat(over), "format validation is unaffected") + errIs(t, VerifyPassword(pw, over), ErrPHCCostTooHigh, "memory over ceiling") + + overTime := fmt.Sprintf("$argon2id$v=%d$m=8192,t=%d,p=1$%s$%s", + argon2.Version, MaxVerifyArgonTime+1, rawB64(16), rawB64(32)) + errIs(t, VerifyPassword(pw, overTime), ErrPHCCostTooHigh, "time over ceiling") + _, err := MigrateFromPHC("u", pw, over) + errIs(t, err, ErrPHCCostTooHigh, "migration honors the ceiling") +} diff --git a/error.go b/error.go index d416a8f..964aec2 100644 --- a/error.go +++ b/error.go @@ -42,6 +42,7 @@ var ( ErrPHCInvalidFormat = errors.New("phc: invalid format") ErrPHCInvalidSalt = errors.New("phc: invalid salt encoding") ErrPHCInvalidHash = errors.New("phc: invalid hash encoding") + ErrPHCCostTooHigh = errors.New("phc: cost parameters exceed verification limit") ) // SCRAM-specific errors @@ -57,6 +58,7 @@ var ( ErrSCRAMZeroParams = errors.New("scram: invalid Argon2 parameters") ErrSCRAMSaltTooShort = errors.New("scram: salt must be at least 16 bytes") ErrSCRAMTooManyHandshakes = errors.New("scram: handshake capacity exceeded") + ErrSCRAMParamsTooLarge = errors.New("scram: Argon2 parameters exceed limit") ) // Credential import/export errors diff --git a/go.mod b/go.mod index 1b631d1..ec5a876 100644 --- a/go.mod +++ b/go.mod @@ -4,13 +4,7 @@ go 1.26.0 require ( github.com/golang-jwt/jwt/v5 v5.3.1 - github.com/stretchr/testify v1.11.1 golang.org/x/crypto v0.54.0 ) -require ( - github.com/davecgh/go-spew v1.1.1 // indirect - github.com/pmezard/go-difflib v1.0.0 // indirect - golang.org/x/sys v0.47.0 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect -) +require golang.org/x/sys v0.47.0 // indirect diff --git a/go.sum b/go.sum index 87ca3fd..6f26f6f 100644 --- a/go.sum +++ b/go.sum @@ -1,30 +1,6 @@ -github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= -github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo= -github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= -github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= -github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= -github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= -golang.org/x/crypto v0.43.0 h1:dduJYIi3A3KOfdGOHX8AVZ/jGiyPa3IbBozJ5kNuE04= -golang.org/x/crypto v0.43.0/go.mod h1:BFbav4mRNlXJL4wNeejLpWxB7wMbc79PdRGhWKncxR0= -golang.org/x/crypto v0.52.0 h1:RMs7fP2rXdep0CftQlK8Uf+kibLm7qkCcradZWYz988= -golang.org/x/crypto v0.52.0/go.mod h1:1QgfPxDqh0T2M/elOJtp9RvuR95kVjir0e6/BvEmGbc= -golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= -golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= -golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ= -golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= -golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= -golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= -golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= -gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/http.go b/http.go index 474167b..dd1f6af 100644 --- a/http.go +++ b/http.go @@ -7,34 +7,30 @@ import ( // ParseBasicAuth extracts username/password from Basic auth header func ParseBasicAuth(header string) (username, password string, err error) { - const prefix = "Basic " - if !strings.HasPrefix(header, prefix) { + encoded, ok := strings.CutPrefix(header, "Basic ") + if !ok { return "", "", ErrAuthInvalidBasicFormat } - encoded := strings.TrimPrefix(header, prefix) decoded, err := base64.StdEncoding.DecodeString(encoded) if err != nil { return "", "", ErrAuthInvalidBasicEncoding } - credentials := string(decoded) - idx := strings.IndexByte(credentials, ':') - if idx < 0 { + username, password, ok = strings.Cut(string(decoded), ":") + if !ok { return "", "", ErrAuthInvalidBasicCreds } - return credentials[:idx], credentials[idx+1:], nil + return username, password, nil } // ParseBearerToken extracts token from Bearer auth header func ParseBearerToken(header string) (token string, err error) { - const prefix = "Bearer " - if !strings.HasPrefix(header, prefix) { + token, ok := strings.CutPrefix(header, "Bearer ") + if !ok { return "", ErrAuthInvalidBearerFormat } - - token = strings.TrimPrefix(header, prefix) if token == "" { return "", ErrAuthEmptyBearerToken } @@ -44,18 +40,8 @@ func ParseBearerToken(header string) (token string, err error) { // ExtractAuthType returns authentication type from header func ExtractAuthType(header string) string { - if strings.HasPrefix(header, "Basic ") { - return "Basic" + if authType, _, ok := strings.Cut(header, " "); ok { + return authType } - if strings.HasPrefix(header, "Bearer ") { - return "Bearer" - } - - // Extract first word as auth type - idx := strings.IndexByte(header, ' ') - if idx > 0 { - return header[:idx] - } - return "" + return "" // Matches original behavior if no space is found or string is empty } - diff --git a/http_test.go b/http_test.go index 0dbb42e..5e5471c 100644 --- a/http_test.go +++ b/http_test.go @@ -2,53 +2,122 @@ package auth import ( "encoding/base64" + "fmt" + "strings" "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" ) -func TestHTTPAuthParsing(t *testing.T) { - // Test Basic Auth - basicHeader := "Basic " + base64.StdEncoding.EncodeToString([]byte("user:pass")) - username, password, err := ParseBasicAuth(basicHeader) - require.NoError(t, err) - assert.Equal(t, "user", username) - assert.Equal(t, "pass", password) - - // Test Bearer Token - bearerHeader := "Bearer test-token-xyz" - token, err := ParseBearerToken(bearerHeader) - require.NoError(t, err) - assert.Equal(t, "test-token-xyz", token) - - // Test ExtractAuthType - assert.Equal(t, "Basic", ExtractAuthType(basicHeader)) - assert.Equal(t, "Bearer", ExtractAuthType(bearerHeader)) - assert.Equal(t, "Custom", ExtractAuthType("Custom somedata")) - assert.Equal(t, "", ExtractAuthType("InvalidHeader")) - - // Test invalid formats - _, _, err = ParseBasicAuth("Invalid header") - assert.Error(t, err) - assert.Equal(t, ErrAuthInvalidBasicFormat, err) - - _, err = ParseBearerToken("Invalid header") - assert.Error(t, err) - assert.Equal(t, ErrAuthInvalidBearerFormat, err) - - // Test malformed Basic auth - _, _, err = ParseBasicAuth("Basic not-base64!") - assert.Error(t, err) - assert.Equal(t, ErrAuthInvalidBasicEncoding, err) - - _, _, err = ParseBasicAuth("Basic " + base64.StdEncoding.EncodeToString([]byte("no-colon"))) - assert.Error(t, err) - assert.Equal(t, ErrAuthInvalidBasicCreds, err) - - // Test empty Bearer token - _, err = ParseBearerToken("Bearer ") - assert.Error(t, err) - assert.Equal(t, ErrAuthEmptyBearerToken, err) +func basicHeader(payload string) string { + return "Basic " + base64.StdEncoding.EncodeToString([]byte(payload)) } +func TestParseBasicAuth(t *testing.T) { + cases := []struct { + name, header, username, password string + want error + }{ + {name: "standard", header: basicHeader("user:pass"), username: "user", password: "pass"}, + {name: "empty password", header: basicHeader("user:"), username: "user"}, + {name: "empty username", header: basicHeader(":pass"), password: "pass"}, + {name: "colon in password", header: basicHeader("user:pa:ss"), username: "user", password: "pa:ss"}, + {name: "unicode", header: basicHeader("üser:pässwörd"), username: "üser", password: "pässwörd"}, + {name: "separator only", header: basicHeader(":")}, + + {name: "missing scheme", header: "Invalid header", want: ErrAuthInvalidBasicFormat}, + {name: "no space", header: "Basic", want: ErrAuthInvalidBasicFormat}, + {name: "empty header", header: "", want: ErrAuthInvalidBasicFormat}, + // ☢ RFC 7235 auth-scheme is case-insensitive; matching here is not + {name: "lowercase scheme", header: "basic dXNlcjpwYXNz", want: ErrAuthInvalidBasicFormat}, + {name: "bad base64", header: "Basic not-base64!", want: ErrAuthInvalidBasicEncoding}, + {name: "no colon", header: basicHeader("no-colon"), want: ErrAuthInvalidBasicCreds}, + {name: "empty payload", header: "Basic ", want: ErrAuthInvalidBasicCreds}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + username, password, err := ParseBasicAuth(tc.header) + if tc.want != nil { + errIs(t, err, tc.want, tc.name) + eq(t, username, "", "username on failure") + eq(t, password, "", "password on failure") + return + } + noErr(t, err, tc.name) + eq(t, username, tc.username, "username") + eq(t, password, tc.password, "password") + }) + } +} + +func TestParseBearerToken(t *testing.T) { + cases := []struct { + name, header, token string + want error + }{ + {name: "standard", header: "Bearer test-token-xyz", token: "test-token-xyz"}, + {name: "embedded space", header: "Bearer a b", token: "a b"}, + {name: "trailing space", header: "Bearer x ", token: "x "}, + + {name: "empty token", header: "Bearer ", want: ErrAuthEmptyBearerToken}, + {name: "missing scheme", header: "Invalid header", want: ErrAuthInvalidBearerFormat}, + {name: "lowercase scheme", header: "bearer x", want: ErrAuthInvalidBearerFormat}, + {name: "no space", header: "Bearer", want: ErrAuthInvalidBearerFormat}, + {name: "empty header", header: "", want: ErrAuthInvalidBearerFormat}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + token, err := ParseBearerToken(tc.header) + if tc.want != nil { + errIs(t, err, tc.want, tc.name) + eq(t, token, "", "token on failure") + return + } + noErr(t, err, tc.name) + eq(t, token, tc.token, "token") + }) + } +} + +func TestExtractAuthType(t *testing.T) { + cases := map[string]string{ + "Basic dXNlcjpwYXNz": "Basic", + "Bearer token": "Bearer", + "Custom somedata": "Custom", + "Basic ": "Basic", + "InvalidHeader": "", + "": "", + " Bearer token": "", // cut at the first space yields an empty scheme + } + for header, want := range cases { + eq(t, ExtractAuthType(header), want, fmt.Sprintf("header %q", header)) + } +} + +func FuzzParseBasicAuth(f *testing.F) { + f.Add(basicHeader("user:pass")) + f.Add("Basic ") + f.Add("") + f.Add("Basic !!!") + + f.Fuzz(func(t *testing.T, header string) { + username, password, err := ParseBasicAuth(header) + if err != nil { + if username != "" || password != "" { + t.Fatal("values returned alongside error") + } + return + } + encoded, ok := strings.CutPrefix(header, "Basic ") + if !ok { + t.Fatal("accepted a header without the Basic prefix") + } + decoded, decErr := base64.StdEncoding.DecodeString(encoded) + if decErr != nil { + t.Fatal("accepted an undecodable payload") + } + if string(decoded) != username+":"+password { + t.Fatalf("lossy split: %q vs %q", decoded, username+":"+password) + } + }) +} diff --git a/jwt.go b/jwt.go index 6c342c7..dd48357 100644 --- a/jwt.go +++ b/jwt.go @@ -111,6 +111,16 @@ func NewJWTRSA(privateKey *rsa.PrivateKey, opts ...JWTOption) (*JWT, error) { return j, nil } +// NewJWTRSAFromPEM creates a JWT manager for RS256 from raw PEM-encoded private key data. +func NewJWTRSAFromPEM(privateKeyPEM []byte, opts ...JWTOption) (*JWT, error) { + privateKey, err := parseRSAPrivateKey(privateKeyPEM) + if err != nil { + return nil, err + } + // Call the original constructor with the now-parsed key + return NewJWTRSA(privateKey, opts...) +} + // NewJWTVerifier creates JWT manager for verification only (RS256) func NewJWTVerifier(publicKey *rsa.PublicKey, opts ...JWTOption) (*JWT, error) { if publicKey == nil { @@ -132,6 +142,16 @@ func NewJWTVerifier(publicKey *rsa.PublicKey, opts ...JWTOption) (*JWT, error) { return j, nil } +// NewJWTVerifierFromPEM creates a JWT manager for verification from raw PEM-encoded public key data. +func NewJWTVerifierFromPEM(publicKeyPEM []byte, opts ...JWTOption) (*JWT, error) { + publicKey, err := parseRSAPublicKey(publicKeyPEM) + if err != nil { + return nil, err + } + // Call the original constructor with the now-parsed key + return NewJWTVerifier(publicKey, opts...) +} + // GenerateToken creates signed JWT with claims func (j *JWT) GenerateToken(userID string, claims map[string]any) (string, error) { if userID == "" { @@ -191,27 +211,26 @@ func (j *JWT) ValidateToken(tokenString string) (string, map[string]any, error) func mapJWTError(err error) error { switch { case errors.Is(err, jwt.ErrTokenMalformed): - return fmt.Errorf("%w : %w", ErrTokenMalformed, err) + return fmt.Errorf("%w: %w", ErrTokenMalformed, err) case errors.Is(err, jwt.ErrTokenUnverifiable): - return fmt.Errorf("%w : %w", ErrTokenMalformed, err) + return fmt.Errorf("%w: %w", ErrTokenMalformed, err) case errors.Is(err, jwt.ErrTokenSignatureInvalid): - return fmt.Errorf("%w : %w", ErrTokenInvalidSignature, err) + return fmt.Errorf("%w: %w", ErrTokenInvalidSignature, err) case errors.Is(err, jwt.ErrTokenExpired): - return fmt.Errorf("%w : %w", ErrTokenExpired, err) + return fmt.Errorf("%w: %w", ErrTokenExpired, err) case errors.Is(err, jwt.ErrTokenNotValidYet): - return fmt.Errorf("%w : %w", ErrTokenNotYetValid, err) + return fmt.Errorf("%w: %w", ErrTokenNotYetValid, err) case errors.Is(err, jwt.ErrTokenInvalidAudience): - return fmt.Errorf("%w : %w", ErrTokenMissingClaim, err) + return fmt.Errorf("%w: %w", ErrTokenMissingClaim, err) case errors.Is(err, jwt.ErrTokenInvalidIssuer): - return fmt.Errorf("%w : %w", ErrTokenMissingClaim, err) + return fmt.Errorf("%w: %w", ErrTokenMissingClaim, err) + case errors.Is(err, jwt.ErrTokenRequiredClaimMissing): + return fmt.Errorf("%w: %w", ErrTokenMissingClaim, err) default: - // Alg rejection (WithValidMethods) surfaces as ErrTokenSignatureInvalid. - return fmt.Errorf("%w : %w", ErrTokenMalformed, err) + return fmt.Errorf("%w: %w", ErrTokenMalformed, err) } } -// Standalone helper functions for one-off operations - // GenerateHS256Token creates HS256 JWT without manager instance func GenerateHS256Token(secret []byte, userID string, claims map[string]any, lifetime time.Duration) (string, error) { if len(secret) < 32 { @@ -263,28 +282,6 @@ func ValidateHS256Token(secret []byte, tokenString string) (string, map[string]a return claims.Subject, claims.Extra, nil } -// RSA Utilities - -// NewJWTRSAFromPEM creates a JWT manager for RS256 from raw PEM-encoded private key data. -func NewJWTRSAFromPEM(privateKeyPEM []byte, opts ...JWTOption) (*JWT, error) { - privateKey, err := parseRSAPrivateKey(privateKeyPEM) - if err != nil { - return nil, err - } - // Call the original constructor with the now-parsed key - return NewJWTRSA(privateKey, opts...) -} - -// NewJWTVerifierFromPEM creates a JWT manager for verification from raw PEM-encoded public key data. -func NewJWTVerifierFromPEM(publicKeyPEM []byte, opts ...JWTOption) (*JWT, error) { - publicKey, err := parseRSAPublicKey(publicKeyPEM) - if err != nil { - return nil, err - } - // Call the original constructor with the now-parsed key - return NewJWTVerifier(publicKey, opts...) -} - // parseRSAPrivateKey parses a PEM-encoded RSA private key. func parseRSAPrivateKey(pemBytes []byte) (*rsa.PrivateKey, error) { block, _ := pem.Decode(pemBytes) diff --git a/jwt_test.go b/jwt_test.go index 3aebc96..befaf2a 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -1,259 +1,606 @@ package auth import ( + "bytes" + "crypto/ed25519" + "crypto/hmac" "crypto/rand" "crypto/rsa" + "crypto/sha256" "crypto/x509" + "encoding/base64" + "encoding/json" "encoding/pem" + "errors" + "fmt" "strings" + "sync" "testing" "time" "github.com/golang-jwt/jwt/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" ) -func TestJWTHS256(t *testing.T) { - secret := []byte("test-secret-key-must-be-32-bytes") - jwtMgr, err := NewJWT(secret) - require.NoError(t, err) +var testSecret = []byte("test-secret-key-must-be-32-bytes") - userID := "user123" - claims := map[string]any{ +func genRSAKey() *rsa.PrivateKey { + key, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + panic(err) + } + return key +} + +// RSA key generation is the dominant cost in this file; amortize it. +var ( + testRSAKey = sync.OnceValue(genRSAKey) + testRSAKeyAlt = sync.OnceValue(genRSAKey) +) + +func defaultHeader() map[string]any { return map[string]any{"alg": "HS256", "typ": "JWT"} } + +// signHS256 assembles a token from raw maps, bypassing the package so that +// malformed and hostile tokens can be constructed. +func signHS256(t *testing.T, secret []byte, header, claims map[string]any) string { + t.Helper() + h, err := json.Marshal(header) + noErr(t, err, "marshal header") + c, err := json.Marshal(claims) + noErr(t, err, "marshal claims") + + signing := base64.RawURLEncoding.EncodeToString(h) + "." + base64.RawURLEncoding.EncodeToString(c) + mac := hmac.New(sha256.New, secret) + mac.Write([]byte(signing)) + return signing + "." + base64.RawURLEncoding.EncodeToString(mac.Sum(nil)) +} + +func unsignedToken(t *testing.T, header, claims map[string]any) string { + t.Helper() + h, err := json.Marshal(header) + noErr(t, err, "marshal header") + c, err := json.Marshal(claims) + noErr(t, err, "marshal claims") + return base64.RawURLEncoding.EncodeToString(h) + "." + base64.RawURLEncoding.EncodeToString(c) + "." +} + +func decodeSegment(t *testing.T, segment string) map[string]any { + t.Helper() + raw, err := base64.RawURLEncoding.DecodeString(segment) + noErr(t, err, "segment decode") + var m map[string]any + noErr(t, json.Unmarshal(raw, &m), "segment unmarshal") + return m +} + +func jwtParts(t *testing.T, token string) (header, payload map[string]any) { + t.Helper() + parts := strings.Split(token, ".") + eq(t, len(parts), 3, "token segments") + return decodeSegment(t, parts[0]), decodeSegment(t, parts[1]) +} + +func str(t *testing.T, m map[string]any, key string) string { + t.Helper() + v, ok := m[key].(string) + if !ok { + t.Fatalf("claim %q: %v is not a string", key, m[key]) + } + return v +} + +func num(t *testing.T, m map[string]any, key string) float64 { + t.Helper() + v, ok := m[key].(float64) + if !ok { + t.Fatalf("claim %q: %v is not a number", key, m[key]) + } + return v +} + +func TestJWTHS256RoundTrip(t *testing.T) { + manager, err := NewJWT(testSecret) + noErr(t, err, "NewJWT") + + token, err := manager.GenerateToken("user123", map[string]any{ "email": "test@example.com", "role": "admin", + }) + noErr(t, err, "GenerateToken") + + header, payload := jwtParts(t, token) + eq(t, str(t, header, "alg"), "HS256", "alg") + eq(t, str(t, header, "typ"), "JWT", "typ") + eq(t, str(t, payload, "sub"), "user123", "subject") + + exp, iat, nbf := num(t, payload, "exp"), num(t, payload, "iat"), num(t, payload, "nbf") + eq(t, int64(exp-iat), int64(DefaultTokenLifetime/time.Second), "default lifetime") + eq(t, nbf, iat, "nbf equals iat") + if _, present := payload["iss"]; present { + t.Fatal("issuer emitted without configuration") } - // Generate token - token, err := jwtMgr.GenerateToken(userID, claims) - require.NoError(t, err) - assert.NotEmpty(t, token) + userID, claims, err := manager.ValidateToken(token) + noErr(t, err, "ValidateToken") + eq(t, userID, "user123", "user id") + eq(t, len(claims), 2, "claim count") + eq(t, str(t, claims, "email"), "test@example.com", "email claim") + eq(t, str(t, claims, "role"), "admin", "role claim") - // Validate token - extractedUserID, extractedClaims, err := jwtMgr.ValidateToken(token) - require.NoError(t, err) + // nil claims must not emit an extra object + bare, err := manager.GenerateToken("user123", nil) + noErr(t, err, "GenerateToken without claims") + _, barePayload := jwtParts(t, bare) + if _, present := barePayload["extra"]; present { + t.Fatal("empty extra claim emitted") + } + _, claims, err = manager.ValidateToken(bare) + noErr(t, err, "ValidateToken without claims") + eq(t, len(claims), 0, "no extra claims") +} - assert.Equal(t, userID, extractedUserID) - assert.Equal(t, "test@example.com", extractedClaims["email"]) - assert.Equal(t, "admin", extractedClaims["role"]) +func TestJWTSecretLength(t *testing.T) { + _, err := NewJWT(nil) + errIs(t, err, ErrSecretTooShort, "nil secret") + _, err = NewJWT(make([]byte, 31)) + errIs(t, err, ErrSecretTooShort, "31 bytes") + _, err = NewJWT(make([]byte, 32)) + noErr(t, err, "32 bytes") + + _, err = GenerateHS256Token(make([]byte, 31), "u", nil, time.Hour) + errIs(t, err, ErrSecretTooShort, "standalone generate") + _, _, err = ValidateHS256Token(make([]byte, 31), "irrelevant") + errIs(t, err, ErrSecretTooShort, "standalone validate") +} + +func TestJWTEmptyUserID(t *testing.T) { + manager, err := NewJWT(testSecret) + noErr(t, err, "NewJWT") + _, err = manager.GenerateToken("", map[string]any{"role": "admin"}) + errIs(t, err, ErrTokenEmptyUserID, "empty user id") + _, err = GenerateHS256Token(testSecret, "", nil, time.Hour) + errIs(t, err, ErrTokenEmptyUserID, "standalone empty user id") } func TestJWTRS256(t *testing.T) { - // Generate RSA key pair - privateKey, err := rsa.GenerateKey(rand.Reader, 2048) - require.NoError(t, err) + key := testRSAKey() - // Test with private key (can sign and verify) - jwtMgr, err := NewJWTRSA(privateKey) - require.NoError(t, err) + signer, err := NewJWTRSA(key) + noErr(t, err, "NewJWTRSA") + token, err := signer.GenerateToken("user456", map[string]any{"scope": "read:all"}) + noErr(t, err, "GenerateToken") - userID := "user456" - claims := map[string]any{ - "scope": "read:all", - } + header, _ := jwtParts(t, token) + eq(t, str(t, header, "alg"), "RS256", "alg") - // Generate token - token, err := jwtMgr.GenerateToken(userID, claims) - require.NoError(t, err) - assert.NotEmpty(t, token) + userID, claims, err := signer.ValidateToken(token) + noErr(t, err, "self validation") + eq(t, userID, "user456", "user id") + eq(t, str(t, claims, "scope"), "read:all", "scope claim") - // Validate with same manager - extractedUserID, extractedClaims, err := jwtMgr.ValidateToken(token) - require.NoError(t, err) - assert.Equal(t, userID, extractedUserID) - assert.Equal(t, "read:all", extractedClaims["scope"]) + verifier, err := NewJWTVerifier(&key.PublicKey) + noErr(t, err, "NewJWTVerifier") + userID, _, err = verifier.ValidateToken(token) + noErr(t, err, "verifier validation") + eq(t, userID, "user456", "user id from verifier") - // Test with verifier only (public key) - verifier, err := NewJWTVerifier(&privateKey.PublicKey) - require.NoError(t, err) + _, err = verifier.GenerateToken("user456", nil) + errIs(t, err, ErrTokenNoPrivateKey, "verifier must not sign") - // Should validate token - extractedUserID, _, err = verifier.ValidateToken(token) - require.NoError(t, err) - assert.Equal(t, userID, extractedUserID) + // an unrelated key must not verify + foreign, err := NewJWTVerifier(&testRSAKeyAlt().PublicKey) + noErr(t, err, "NewJWTVerifier foreign") + _, _, err = foreign.ValidateToken(token) + errIs(t, err, ErrTokenInvalidSignature, "foreign public key") - // Should not generate token - _, err = verifier.GenerateToken(userID, claims) - assert.Equal(t, ErrTokenNoPrivateKey, err) + _, err = NewJWTRSA(nil) + errIs(t, err, ErrTokenNoPrivateKey, "nil private key") + _, err = NewJWTVerifier(nil) + errIs(t, err, ErrTokenNoPublicKey, "nil public key") } -func TestJWTOptions(t *testing.T) { - secret := []byte("test-secret-key-must-be-32-bytes") +func TestJWTAlgorithmEnforcement(t *testing.T) { + key := testRSAKey() + hs, err := NewJWT(testSecret) + noErr(t, err, "NewJWT") + rs, err := NewJWTRSA(key) + noErr(t, err, "NewJWTRSA") - // Test custom lifetime - jwtMgr, err := NewJWT(secret, - WithTokenLifetime(1*time.Hour), + hsToken, err := hs.GenerateToken("u", nil) + noErr(t, err, "HS256 token") + rsToken, err := rs.GenerateToken("u", nil) + noErr(t, err, "RS256 token") + + _, _, err = hs.ValidateToken(rsToken) + errIs(t, err, ErrTokenInvalidSignature, "RS256 token to HS256 manager") + _, _, err = rs.ValidateToken(hsToken) + errIs(t, err, ErrTokenInvalidSignature, "HS256 token to RS256 manager") + + // ☢ algorithm confusion: HS256 token keyed with the RSA public key + pubBytes, err := x509.MarshalPKIXPublicKey(&key.PublicKey) + noErr(t, err, "marshal public key") + pubPEM := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: pubBytes}) + forged := signHS256(t, pubPEM, defaultHeader(), map[string]any{ + "sub": "attacker", "exp": time.Now().Add(time.Hour).Unix(), + }) + _, _, err = rs.ValidateToken(forged) + errIs(t, err, ErrTokenInvalidSignature, "algorithm confusion") + + // alg: none + none := unsignedToken(t, map[string]any{"alg": "none", "typ": "JWT"}, map[string]any{ + "sub": "attacker", "exp": time.Now().Add(time.Hour).Unix(), + }) + _, _, err = hs.ValidateToken(none) + errIs(t, err, ErrTokenInvalidSignature, "alg none against HS256") + _, _, err = rs.ValidateToken(none) + errIs(t, err, ErrTokenInvalidSignature, "alg none against RS256") + + // an unregistered algorithm is unverifiable, mapped to malformed + unknown := signHS256(t, testSecret, map[string]any{"alg": "HS999", "typ": "JWT"}, map[string]any{ + "sub": "attacker", "exp": time.Now().Add(time.Hour).Unix(), + }) + _, _, err = hs.ValidateToken(unknown) + errIs(t, err, ErrTokenMalformed, "unknown algorithm") +} + +func TestJWTTampering(t *testing.T) { + manager, err := NewJWT(testSecret) + noErr(t, err, "NewJWT") + token, err := manager.GenerateToken("user1", map[string]any{"role": "user"}) + noErr(t, err, "GenerateToken") + parts := strings.Split(token, ".") + + escalated, err := json.Marshal(map[string]any{ + "sub": "user1", "exp": time.Now().Add(time.Hour).Unix(), + "extra": map[string]any{"role": "admin"}, + }) + noErr(t, err, "marshal forged claims") + + // length-preserving corruption, so the segment still decodes and + // the failure is attributable to the MAC rather than the encoding + corrupt := []byte(parts[2]) + if corrupt[0] == 'A' { + corrupt[0] = 'B' + } else { + corrupt[0] = 'A' + } + + cases := []struct { + name string + token string + want error + }{ + {"payload rewrite", parts[0] + "." + base64.RawURLEncoding.EncodeToString(escalated) + "." + parts[2], ErrTokenInvalidSignature}, + {"corrupt signature", parts[0] + "." + parts[1] + "." + string(corrupt), ErrTokenInvalidSignature}, + {"truncated signature", parts[0] + "." + parts[1] + "." + parts[2][:len(parts[2])-2], ErrTokenMalformed}, + {"replaced signature", parts[0] + "." + parts[1] + ".invalidsignature", ErrTokenInvalidSignature}, + {"empty", "", ErrTokenMalformed}, + {"two segments", parts[0] + "." + parts[1], ErrTokenMalformed}, + {"four segments", token + ".extra", ErrTokenMalformed}, + {"separators only", "..", ErrTokenMalformed}, + {"payload not base64", parts[0] + ".!!!." + parts[2], ErrTokenMalformed}, + {"header not base64", "!!!." + parts[1] + "." + parts[2], ErrTokenMalformed}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + userID, claims, err := manager.ValidateToken(tc.token) + errIs(t, err, tc.want, tc.name) + eq(t, userID, "", "user id on failure") + if claims != nil { + t.Fatalf("%s: claims returned on failure", tc.name) + } + }) + } +} + +func TestJWTExpiryAndLeeway(t *testing.T) { + now := time.Now() + strict, err := NewJWT(testSecret, WithLeeway(0)) + noErr(t, err, "NewJWT strict") + lenient, err := NewJWT(testSecret, WithLeeway(5*time.Minute)) + noErr(t, err, "NewJWT lenient") + + expired := signHS256(t, testSecret, defaultHeader(), map[string]any{ + "sub": "u", "iat": now.Add(-2 * time.Hour).Unix(), "exp": now.Add(-time.Minute).Unix(), + }) + _, _, err = strict.ValidateToken(expired) + errIs(t, err, ErrTokenExpired, "expired token") + _, _, err = lenient.ValidateToken(expired) + noErr(t, err, "expiry inside leeway") + + notYet := signHS256(t, testSecret, defaultHeader(), map[string]any{ + "sub": "u", "nbf": now.Add(2 * time.Second).Unix(), "exp": now.Add(time.Hour).Unix(), + }) + _, _, err = strict.ValidateToken(notYet) + errIs(t, err, ErrTokenNotYetValid, "nbf in the future") + _, _, err = lenient.ValidateToken(notYet) + noErr(t, err, "nbf inside leeway") + + // exp is mandatory. ‼️ mapJWTError has no case for + // jwt.ErrTokenRequiredClaimMissing, so this surfaces as malformed. + noExp := signHS256(t, testSecret, defaultHeader(), map[string]any{"sub": "u", "iat": now.Unix()}) + _, _, err = strict.ValidateToken(noExp) + errIs(t, err, ErrTokenMissingClaim, "missing exp") + + // generated lifetimes are honored without waiting for them + short, err := NewJWT(testSecret, WithTokenLifetime(time.Second)) + noErr(t, err, "NewJWT short lifetime") + token, err := short.GenerateToken("u", nil) + noErr(t, err, "GenerateToken") + _, payload := jwtParts(t, token) + eq(t, int64(num(t, payload, "exp")-num(t, payload, "iat")), int64(1), "encoded lifetime") + _, _, err = short.ValidateToken(token) + noErr(t, err, "valid immediately") +} + +func TestJWTIssuerAudience(t *testing.T) { + manager, err := NewJWT(testSecret, + WithTokenLifetime(time.Hour), WithIssuer("test-issuer"), WithAudience([]string{"api.example.com"}), ) - require.NoError(t, err) + noErr(t, err, "NewJWT") - token, err := jwtMgr.GenerateToken("user1", nil) - require.NoError(t, err) + token, err := manager.GenerateToken("user1", nil) + noErr(t, err, "GenerateToken") - // Parse token to check claims - parsed, _ := jwt.Parse(token, func(token *jwt.Token) (any, error) { - return secret, nil + _, payload := jwtParts(t, token) + eq(t, str(t, payload, "iss"), "test-issuer", "issuer") + audience, ok := payload["aud"].([]any) + isTrue(t, ok, "audience encoded as an array") + eq(t, len(audience), 1, "audience length") + eq(t, audience[0], any("api.example.com"), "audience value") + eq(t, int64(num(t, payload, "exp")-num(t, payload, "iat")), int64(3600), "custom lifetime") + + _, _, err = manager.ValidateToken(token) + noErr(t, err, "self-issued token") + + other, err := NewJWT(testSecret, WithTokenLifetime(time.Hour), WithIssuer("other-issuer")) + noErr(t, err, "NewJWT other issuer") + otherToken, err := other.GenerateToken("user1", nil) + noErr(t, err, "GenerateToken other issuer") + _, _, err = manager.ValidateToken(otherToken) + errIs(t, err, ErrTokenMissingClaim, "issuer mismatch") + + missingAudience := signHS256(t, testSecret, defaultHeader(), map[string]any{ + "sub": "u", "iss": "test-issuer", "exp": time.Now().Add(time.Hour).Unix(), }) + _, _, err = manager.ValidateToken(missingAudience) + errIs(t, err, ErrTokenMissingClaim, "absent audience is rejected when expected") - claims := parsed.Claims.(jwt.MapClaims) - - // Check issuer - assert.Equal(t, "test-issuer", claims["iss"]) - - // Check audience - aud := claims["aud"].([]any) - assert.Contains(t, aud, "api.example.com") - - // Check expiration is ~1 hour - exp := int64(claims["exp"].(float64)) - iat := int64(claims["iat"].(float64)) - assert.InDelta(t, 3600, exp-iat, 10) -} - -func TestJWTErrors(t *testing.T) { - secret := []byte("test-secret-key-must-be-32-bytes") - jwtMgr, err := NewJWT(secret) - require.NoError(t, err) - - // Empty user ID - _, err = jwtMgr.GenerateToken("", nil) - assert.Equal(t, ErrTokenEmptyUserID, err) - - // Invalid token format - _, _, err = jwtMgr.ValidateToken("invalid.token") - assert.ErrorIs(t, err, ErrTokenMalformed) - - // Tampered signature - token, _ := jwtMgr.GenerateToken("user1", nil) - parts := strings.Split(token, ".") - tampered := parts[0] + "." + parts[1] + ".invalidsignature" - _, _, err = jwtMgr.ValidateToken(tampered) - assert.ErrorIs(t, err, ErrTokenInvalidSignature) - - // Wrong algorithm - rsaKey, _ := rsa.GenerateKey(rand.Reader, 2048) - rsaMgr, _ := NewJWTRSA(rsaKey) - rsaToken, _ := rsaMgr.GenerateToken("user1", nil) - - _, _, err = jwtMgr.ValidateToken(rsaToken) - assert.ErrorIs(t, err, ErrTokenInvalidSignature) -} - -func TestJWTExpiration(t *testing.T) { - secret := []byte("test-secret-key-must-be-32-bytes") - - // Create token with 1 second lifetime - jwtMgr, err := NewJWT(secret, WithTokenLifetime(1*time.Second), WithLeeway(0)) - require.NoError(t, err) - - token, err := jwtMgr.GenerateToken("user1", nil) - require.NoError(t, err) - - // Should be valid immediately - _, _, err = jwtMgr.ValidateToken(token) - assert.NoError(t, err) - - // Wait for expiration - time.Sleep(2 * time.Second) - - // Should be expired - _, _, err = jwtMgr.ValidateToken(token) - assert.ErrorIs(t, err, ErrTokenExpired) -} - -func TestLeeway(t *testing.T) { - secret := []byte("test-secret-key-must-be-32-bytes") - - // Create manager with no leeway - jwtMgr, err := NewJWT(secret, WithLeeway(0)) - require.NoError(t, err) - - // Manually create a token with NotBefore in future - now := time.Now() - token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ - "sub": "user1", - "nbf": now.Add(2 * time.Second).Unix(), - "exp": now.Add(1 * time.Hour).Unix(), + superset := signHS256(t, testSecret, defaultHeader(), map[string]any{ + "sub": "u", "iss": "test-issuer", "exp": time.Now().Add(time.Hour).Unix(), + "aud": []string{"other.example.com", "api.example.com"}, }) - tokenString, err := token.SignedString(secret) - require.NoError(t, err) + _, _, err = manager.ValidateToken(superset) + noErr(t, err, "expected audience among others") - // Should fail immediately (not valid yet) - _, _, err = jwtMgr.ValidateToken(tokenString) - assert.ErrorIs(t, err, ErrTokenNotYetValid) - - // Create manager with leeway - jwtMgrWithLeeway, err := NewJWT(secret, WithLeeway(5*time.Second)) - require.NoError(t, err) - - // Should pass with leeway - _, _, err = jwtMgrWithLeeway.ValidateToken(tokenString) - assert.NoError(t, err) + // an unconstrained manager imposes neither claim + plain, err := NewJWT(testSecret) + noErr(t, err, "NewJWT plain") + _, _, err = plain.ValidateToken(token) + noErr(t, err, "unconstrained validation") } -func TestStandaloneFunctions(t *testing.T) { - secret := []byte("test-secret-key-must-be-32-bytes") - userID := "standalone-user" - claims := map[string]any{"test": "value"} +func TestJWTUnenforcedClaims(t *testing.T) { + // Documented gaps: sub is not required and iat is not verified. + // Callers must reject an empty user id themselves. + manager, err := NewJWT(testSecret) + noErr(t, err, "NewJWT") - // Generate token - token, err := GenerateHS256Token(secret, userID, claims, 1*time.Hour) - require.NoError(t, err) - - // Validate token - extractedUserID, extractedClaims, err := ValidateHS256Token(secret, token) - require.NoError(t, err) - - assert.Equal(t, userID, extractedUserID) - assert.Equal(t, "value", extractedClaims["test"]) - - // Test with short secret - _, err = GenerateHS256Token([]byte("short"), userID, claims, 1*time.Hour) - assert.Equal(t, ErrSecretTooShort, err) + token := signHS256(t, testSecret, defaultHeader(), map[string]any{ + "exp": time.Now().Add(time.Hour).Unix(), + "iat": time.Now().Add(24 * time.Hour).Unix(), + }) + userID, claims, err := manager.ValidateToken(token) + noErr(t, err, "token without subject") + eq(t, userID, "", "empty subject accepted") + eq(t, len(claims), 0, "no extra claims") } -func TestJWTRSAFromPEM(t *testing.T) { - // 1. Generate a new RSA key pair for this test - privateKey, err := rsa.GenerateKey(rand.Reader, 2048) - require.NoError(t, err) +func TestJWTOptionGuards(t *testing.T) { + manager, err := NewJWT(testSecret, + WithTokenLifetime(0), WithTokenLifetime(-time.Hour), WithLeeway(-time.Second)) + noErr(t, err, "NewJWT") + eq(t, manager.tokenLifetime, DefaultTokenLifetime, "lifetime unchanged") + eq(t, manager.leeway, DefaultLeeway, "leeway unchanged") - // 2. Encode the private key to PEM format - privateKeyPEM := pem.EncodeToMemory(&pem.Block{ - Type: "RSA PRIVATE KEY", - Bytes: x509.MarshalPKCS1PrivateKey(privateKey), - }) + manager, err = NewJWT(testSecret, WithLeeway(0)) + noErr(t, err, "NewJWT zero leeway") + eq(t, manager.leeway, time.Duration(0), "zero leeway applied") - // 3. Encode the public key to PEM format - publicKeyBytes, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey) - require.NoError(t, err) - publicKeyPEM := pem.EncodeToMemory(&pem.Block{ - Type: "PUBLIC KEY", - Bytes: publicKeyBytes, - }) + // options apply to every constructor + verifier, err := NewJWTVerifier(&testRSAKey().PublicKey, WithIssuer("iss"), WithLeeway(time.Minute)) + noErr(t, err, "NewJWTVerifier") + eq(t, verifier.issuer, "iss", "issuer") + eq(t, verifier.leeway, time.Minute, "leeway") +} - // 4. Test the PEM constructor for the signer - jwtMgr, err := NewJWTRSAFromPEM(privateKeyPEM) - require.NoError(t, err) +func TestJWTStandaloneFunctions(t *testing.T) { + token, err := GenerateHS256Token(testSecret, "standalone-user", + map[string]any{"test": "value", "count": 42}, time.Hour) + noErr(t, err, "GenerateHS256Token") - token, err := jwtMgr.GenerateToken("user-from-pem", nil) - require.NoError(t, err) - assert.NotEmpty(t, token) + userID, claims, err := ValidateHS256Token(testSecret, token) + noErr(t, err, "ValidateHS256Token") + eq(t, userID, "standalone-user", "user id") + eq(t, str(t, claims, "test"), "value", "string claim") + eq(t, claims["count"], any(float64(42)), "numeric claim after JSON round trip") - // 5. Test the PEM constructor for the verifier - verifier, err := NewJWTVerifierFromPEM(publicKeyPEM) - require.NoError(t, err) + _, _, err = ValidateHS256Token(bytes.Repeat([]byte("x"), 32), token) + errIs(t, err, ErrTokenInvalidSignature, "wrong secret") - userID, _, err := verifier.ValidateToken(token) - require.NoError(t, err) - assert.Equal(t, "user-from-pem", userID) + expired, err := GenerateHS256Token(testSecret, "u", nil, -time.Hour) + noErr(t, err, "expired token") + _, _, err = ValidateHS256Token(testSecret, expired) + errIs(t, err, ErrTokenExpired, "expired beyond default leeway") + + // standalone validation checks neither issuer nor audience + scoped, err := NewJWT(testSecret, WithIssuer("iss"), WithAudience([]string{"aud"})) + noErr(t, err, "NewJWT scoped") + scopedToken, err := scoped.GenerateToken("u", nil) + noErr(t, err, "GenerateToken scoped") + _, _, err = ValidateHS256Token(testSecret, scopedToken) + noErr(t, err, "issuer and audience are not enforced standalone") + + // tokens are interchangeable with the manager form + managed, err := NewJWT(testSecret) + noErr(t, err, "NewJWT") + _, _, err = managed.ValidateToken(token) + noErr(t, err, "standalone token accepted by manager") +} + +func TestJWTPEM(t *testing.T) { + key := testRSAKey() + + pkcs1 := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)}) + pkcs8Bytes, err := x509.MarshalPKCS8PrivateKey(key) + noErr(t, err, "marshal pkcs8") + pkcs8 := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: pkcs8Bytes}) + pkixBytes, err := x509.MarshalPKIXPublicKey(&key.PublicKey) + noErr(t, err, "marshal pkix") + pkix := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: pkixBytes}) + + for name, blob := range map[string][]byte{"pkcs1": pkcs1, "pkcs8": pkcs8} { + t.Run(name, func(t *testing.T) { + signer, err := NewJWTRSAFromPEM(blob, WithTokenLifetime(time.Hour)) + noErr(t, err, "NewJWTRSAFromPEM") + eq(t, signer.tokenLifetime, time.Hour, "options forwarded") + + token, err := signer.GenerateToken("user-from-pem", nil) + noErr(t, err, "GenerateToken") + + verifier, err := NewJWTVerifierFromPEM(pkix) + noErr(t, err, "NewJWTVerifierFromPEM") + userID, _, err := verifier.ValidateToken(token) + noErr(t, err, "ValidateToken") + eq(t, userID, "user-from-pem", "user id") + }) + } - // 6. Test failure cases with invalid data _, err = NewJWTRSAFromPEM([]byte("invalid pem data")) - assert.ErrorIs(t, err, ErrRSAInvalidPEM) - + errIs(t, err, ErrRSAInvalidPEM, "private: not pem") + _, err = NewJWTRSAFromPEM(nil) + errIs(t, err, ErrRSAInvalidPEM, "private: empty") _, err = NewJWTVerifierFromPEM([]byte("invalid pem data")) - assert.ErrorIs(t, err, ErrRSAInvalidPEM) + errIs(t, err, ErrRSAInvalidPEM, "public: not pem") + + _, err = NewJWTRSAFromPEM(pkix) + errIs(t, err, ErrRSAInvalidPrivateKey, "public key supplied as private") + _, err = NewJWTVerifierFromPEM(pkcs8) + errIs(t, err, ErrRSAInvalidPublicKey, "private key supplied as public") + + // well-formed keys of the wrong algorithm + edPub, edPriv, err := ed25519.GenerateKey(rand.Reader) + noErr(t, err, "ed25519 keygen") + edPubBytes, err := x509.MarshalPKIXPublicKey(edPub) + noErr(t, err, "marshal ed25519 public") + _, err = NewJWTVerifierFromPEM(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: edPubBytes})) + errIs(t, err, ErrRSANotPublicKey, "non-rsa public key") + + edPrivBytes, err := x509.MarshalPKCS8PrivateKey(edPriv) + noErr(t, err, "marshal ed25519 private") + _, err = NewJWTRSAFromPEM(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: edPrivBytes})) + errIs(t, err, ErrRSAInvalidPrivateKey, "non-rsa private key") } +func TestMapJWTError(t *testing.T) { + cases := []struct { + in error + want error + }{ + {jwt.ErrTokenMalformed, ErrTokenMalformed}, + {jwt.ErrTokenUnverifiable, ErrTokenMalformed}, + {jwt.ErrTokenSignatureInvalid, ErrTokenInvalidSignature}, + {jwt.ErrTokenExpired, ErrTokenExpired}, + {jwt.ErrTokenNotValidYet, ErrTokenNotYetValid}, + {jwt.ErrTokenInvalidAudience, ErrTokenMissingClaim}, + {jwt.ErrTokenInvalidIssuer, ErrTokenMissingClaim}, + {jwt.ErrTokenRequiredClaimMissing, ErrTokenMissingClaim}, + {errors.New("unclassified"), ErrTokenMalformed}, + } + for _, tc := range cases { + got := mapJWTError(tc.in) + errIs(t, got, tc.want, "mapped sentinel") + errIs(t, got, tc.in, "original error preserved") + } +} + +func TestJWTConcurrency(t *testing.T) { + manager, err := NewJWT(testSecret, WithTokenLifetime(time.Hour), WithIssuer("iss")) + noErr(t, err, "NewJWT") + shared, err := manager.GenerateToken("user1", map[string]any{"role": "admin"}) + noErr(t, err, "GenerateToken") + + const n = 32 + errs := make(chan error, n) + var wg sync.WaitGroup + for i := range n { + wg.Add(1) + go func() { + defer wg.Done() + if _, _, err := manager.ValidateToken(shared); err != nil { + errs <- err + return + } + token, err := manager.GenerateToken(fmt.Sprintf("user-%d", i), map[string]any{"n": i}) + if err != nil { + errs <- err + return + } + if _, _, err := manager.ValidateToken(token); err != nil { + errs <- err + } + }() + } + wg.Wait() + close(errs) + for err := range errs { + t.Errorf("concurrent operation: %v", err) + } +} + +func FuzzValidateHS256Token(f *testing.F) { + token, err := GenerateHS256Token(testSecret, "seed", map[string]any{"a": 1}, time.Hour) + if err != nil { + f.Fatal(err) + } + f.Add(token) + f.Add("") + f.Add("a.b.c") + f.Add(strings.Repeat(".", 16)) + + f.Fuzz(func(t *testing.T, s string) { + _, _, err := ValidateHS256Token(testSecret, s) + if err != nil { + return + } + // acceptance implies a well-formed, correctly signed token + parts := strings.Split(s, ".") + if len(parts) != 3 { + t.Fatalf("accepted token with %d segments", len(parts)) + } + mac := hmac.New(sha256.New, testSecret) + mac.Write([]byte(parts[0] + "." + parts[1])) + sig, decErr := base64.RawURLEncoding.DecodeString(parts[2]) + if decErr != nil || !hmac.Equal(sig, mac.Sum(nil)) { + t.Fatal("accepted token with an invalid signature") + } + }) +} + +func BenchmarkJWTHS256(b *testing.B) { + manager, err := NewJWT(testSecret) + if err != nil { + b.Fatal(err) + } + claims := map[string]any{"role": "admin"} + for b.Loop() { + token, err := manager.GenerateToken("user1", claims) + if err != nil { + b.Fatal(err) + } + if _, _, err := manager.ValidateToken(token); err != nil { + b.Fatal(err) + } + } +} diff --git a/scram.go b/scram.go index 1303ab1..7e57905 100644 --- a/scram.go +++ b/scram.go @@ -154,12 +154,18 @@ func ImportCredential(data map[string]any) (*Credential, error) { if argonTime == 0 || argonMemory == 0 || argonThreads == 0 { return nil, ErrSCRAMZeroParams } + if err := checkArgonCost(argonMemory, argonTime, argonThreads); err != nil { + return nil, ErrSCRAMParamsTooLarge + } if len(salt) < 16 { return nil, ErrSCRAMSaltTooShort } - if len(storedKey) != sha256.Size || len(serverKey) != sha256.Size { + if len(storedKey) != sha256.Size { return nil, ErrCredInvalidStoredKey } + if len(serverKey) != sha256.Size { + return nil, ErrCredInvalidServerKey + } return &Credential{ Username: username, @@ -405,7 +411,8 @@ func (s *ScramServer) ProcessClientFinalMessage(fullNonce, clientProof string) ( if len(clientProofBytes) != len(clientSignature) { return ServerFinalMessage{}, ErrSCRAMInvalidProofLen } - clientKey := xorBytes(clientProofBytes, clientSignature) + clientKey := make([]byte, len(clientProofBytes)) + subtle.XORBytes(clientKey, clientProofBytes, clientSignature) // Verify by computing StoredKey computedStoredKey := sha256.Sum256(clientKey) @@ -491,6 +498,11 @@ func (c *ScramClient) ProcessServerFirstMessage(msg ServerFirstMessage) (ClientF if msg.ArgonTime == 0 || msg.ArgonMemory == 0 || msg.ArgonThreads == 0 { return ClientFinalRequest{}, ErrSCRAMZeroParams } + // The peer chooses these values. Unbounded, they are a remote OOM + // against every client that talks to a hostile or compromised server. + if err := checkArgonCost(msg.ArgonMemory, msg.ArgonTime, msg.ArgonThreads); err != nil { + return ClientFinalRequest{}, ErrSCRAMParamsTooLarge + } // Derive keys using Argon2id saltedPassword := argon2.IDKey([]byte(c.Password), salt, msg.ArgonTime, msg.ArgonMemory, msg.ArgonThreads, 32) @@ -506,7 +518,8 @@ func (c *ScramClient) ProcessServerFirstMessage(msg ServerFirstMessage) (ClientF // Compute client proof clientSignature := computeHMAC(storedKey[:], []byte(c.authMessage)) - clientProof := xorBytes(clientKey, clientSignature) + clientProof := make([]byte, len(clientKey)) + subtle.XORBytes(clientProof, clientKey, clientSignature) // Store server key for verification c.serverKey = serverKey @@ -589,14 +602,3 @@ func computeHMAC(key, message []byte) []byte { mac.Write(message) return mac.Sum(nil) } - -func xorBytes(a, b []byte) []byte { - if len(a) != len(b) { - panic("xor length mismatch") - } - result := make([]byte, len(a)) - for i := range a { - result[i] = a[i] ^ b[i] - } - return result -} diff --git a/scram_test.go b/scram_test.go index e379e23..01a5721 100644 --- a/scram_test.go +++ b/scram_test.go @@ -1,334 +1,700 @@ package auth import ( + "bytes" + "crypto/sha256" + "encoding/base64" + "encoding/json" "fmt" + "math" + "strings" "sync" "testing" "time" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" + "golang.org/x/crypto/argon2" ) -// setupScramTest is a helper to initialize a server and user credential for testing. -// It performs the full Argon2 -> SCRAM migration workflow. -func setupScramTest(t *testing.T) (server *ScramServer, username, password string, cred *Credential) { - username = "testuser" - password = "SecurePassword123" - - // 1. Start with an Argon2 PHC hash, as a real application would. - phcHash, err := HashPassword(password) - require.NoError(t, err, "Setup failed: could not hash password") - - // 2. Migrate the PHC hash to a SCRAM credential. - cred, err = MigrateFromPHC(username, password, phcHash) - require.NoError(t, err, "Setup failed: could not migrate from PHC hash") - - // 3. Create a server and add the new credential. - server = NewScramServer() - t.Cleanup(server.Stop) - server.AddCredential(cred) - - return server, username, password, cred +func newTestServer(t *testing.T) *ScramServer { + t.Helper() + s := NewScramServer() + t.Cleanup(s.Stop) + return s } -// TestScram_FullRoundtrip_Success simulates a complete, successful authentication handshake. -func TestScram_FullRoundtrip_Success(t *testing.T) { - server, username, password, _ := setupScramTest(t) - client := NewScramClient(username, password) - - // --- Step 1: Client sends its first message --- - clientFirst, err := client.StartAuthentication() - require.NoError(t, err) - - // --- Step 2: Server receives client's message and responds --- - serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) - require.NoError(t, err, "Server failed to process client's first message") - - // --- Step 3: Client receives server's message, computes proof --- - clientFinal, err := client.ProcessServerFirstMessage(serverFirst) - require.NoError(t, err, "Client failed to process server's first message") - - // --- Step 4: Server receives client's proof and verifies it --- - serverFinal, err := server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof) - require.NoError(t, err, "Server failed to verify client's final proof") - assert.NotEmpty(t, serverFinal.ServerSignature, "Server signature should not be empty") - - // --- Step 5: Client verifies server's signature (mutual authentication) --- - err = client.VerifyServerFinalMessage(serverFinal) - assert.NoError(t, err, "Client failed to verify server's final signature") - - t.Log("SCRAM full roundtrip successful") +func testCredential(t *testing.T, username, password string) *Credential { + t.Helper() + phcHash, err := HashPassword(password, cheapArgon...) + noErr(t, err, "HashPassword") + cred, err := MigrateFromPHC(username, password, phcHash) + noErr(t, err, "MigrateFromPHC") + return cred } -// TestScram_FullRoundtrip_WrongPassword ensures authentication fails with an incorrect password. -func TestScram_FullRoundtrip_WrongPassword(t *testing.T) { - server, username, _, _ := setupScramTest(t) - defer server.Stop() - // Create a client with the WRONG password - client := NewScramClient(username, "WrongPassword!!!") - - // Steps 1-3 will appear to succeed, as the client doesn't know the password is wrong yet. - clientFirst, err := client.StartAuthentication() - require.NoError(t, err) - - serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) - require.NoError(t, err) - - clientFinal, err := client.ProcessServerFirstMessage(serverFirst) - require.NoError(t, err) - - // --- Step 4: Server verification should fail here --- - _, err = server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof) - assert.ErrorIs(t, err, ErrInvalidCredentials, "Server should reject proof from wrong password") - - t.Log("SCRAM correctly failed for wrong password") +func setupScram(t *testing.T) (*ScramServer, string, string, *Credential) { + t.Helper() + const username, password = "testuser", "SecurePassword123" + cred := testCredential(t, username, password) + s := newTestServer(t) + s.AddCredential(cred) + return s, username, password, cred } -// TestScram_FullRoundtrip_UserNotFound tests for user enumeration protection. -// The server must be indistinguishable from the wrong-password path: no error at -// first message, stable decoy salt across probes, ErrInvalidCredentials at proof. -func TestScram_FullRoundtrip_UserNotFound(t *testing.T) { - server, _, _, _ := setupScramTest(t) - defer server.Stop() - client := NewScramClient("unknown_user", "any_password") - - clientFirst, err := client.StartAuthentication() - require.NoError(t, err) - - // unknown user must not error here - serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) - require.NoError(t, err, "unknown user must not be signalled at first message") - assert.NotEmpty(t, serverFirst.FullNonce) - assert.NotEmpty(t, serverFirst.Salt) - - // decoy salt must be stable across repeated probes - second, err := server.ProcessClientFirstMessage(clientFirst.Username, "probe-nonce-2") - require.NoError(t, err) - assert.Equal(t, serverFirst.Salt, second.Salt, "decoy salt must be deterministic") - assert.Equal(t, serverFirst.ArgonTime, second.ArgonTime) - assert.Equal(t, serverFirst.ArgonMemory, second.ArgonMemory) - assert.Equal(t, serverFirst.ArgonThreads, second.ArgonThreads) - - clientFinal, err := client.ProcessServerFirstMessage(serverFirst) - require.NoError(t, err) - - // failure surfaces only here, identical to wrong-password path - _, err = server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof) - assert.ErrorIs(t, err, ErrInvalidCredentials, "unknown user must fail like wrong password") +func handshakeCount(s *ScramServer) int { + s.mu.RLock() + defer s.mu.RUnlock() + return len(s.handshakes) } -// TestScram_InvalidNonce simulates a replay attack or message mismatch. -func TestScram_InvalidNonce(t *testing.T) { - server, username, password, _ := setupScramTest(t) - client := NewScramClient(username, password) - - // Perform the first part of the handshake - clientFirst, _ := client.StartAuthentication() - serverFirst, _ := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) - clientFinal, _ := client.ProcessServerFirstMessage(serverFirst) - - // Attempt to finalize with a completely different nonce - _, err := server.ProcessClientFinalMessage("this-is-a-bad-nonce", clientFinal.ClientProof) - assert.ErrorIs(t, err, ErrSCRAMInvalidNonce, "Server should reject a final message with an unknown nonce") +// startHandshake drives a fresh client to the point where a proof is pending. +func startHandshake(t *testing.T, s *ScramServer, username, password string) ClientFinalRequest { + t.Helper() + c := NewScramClient(username, password) + first, err := c.StartAuthentication() + noErr(t, err, "StartAuthentication") + serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce) + noErr(t, err, "ProcessClientFirstMessage") + final, err := c.ProcessServerFirstMessage(serverFirst) + noErr(t, err, "ProcessServerFirstMessage") + return final } -// TestScram_CredentialImportExport verifies that credentials can be serialized and deserialized correctly. -func TestScram_CredentialImportExport(t *testing.T) { - _, _, _, originalCred := setupScramTest(t) - - // Export the credential to a map - exportedData := originalCred.Export() - require.NotNil(t, exportedData) - - // Assert that required fields exist and are strings (as they are base64 encoded) - assert.IsType(t, "", exportedData["salt"]) - assert.IsType(t, "", exportedData["stored_key"]) - assert.IsType(t, "", exportedData["server_key"]) - - // Import the credential back from the map - importedCred, err := ImportCredential(exportedData) - require.NoError(t, err) - require.NotNil(t, importedCred) - - // Verify that the imported credential is identical to the original - assert.Equal(t, originalCred.Username, importedCred.Username) - assert.Equal(t, originalCred.Salt, importedCred.Salt) - assert.Equal(t, originalCred.ArgonTime, importedCred.ArgonTime) - assert.Equal(t, originalCred.ArgonMemory, importedCred.ArgonMemory) - assert.Equal(t, originalCred.ArgonThreads, importedCred.ArgonThreads) - assert.Equal(t, originalCred.StoredKey, importedCred.StoredKey) - assert.Equal(t, originalCred.ServerKey, importedCred.ServerKey) - - t.Log("SCRAM credential import/export successful") -} - -// TestScramServerCleanup verifies automatic cleanup of expired handshakes -func TestScramServerCleanup(t *testing.T) { - // Create server with short cleanup interval for testing - server := NewScramServer() - defer server.Stop() - - // Add a test credential - cred := &Credential{ - Username: "testuser", - Salt: []byte("salt1234567890123456"), - ArgonTime: 1, - ArgonMemory: 64, - ArgonThreads: 1, - StoredKey: []byte("stored_key_placeholder"), - ServerKey: []byte("server_key_placeholder"), +func runHandshake(s *ScramServer, c *ScramClient) error { + first, err := c.StartAuthentication() + if err != nil { + return err } - server.AddCredential(cred) + serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce) + if err != nil { + return err + } + final, err := c.ProcessServerFirstMessage(serverFirst) + if err != nil { + return err + } + serverFinal, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + if err != nil { + return err + } + return c.VerifyServerFinalMessage(serverFinal) +} - // Start multiple handshakes - var nonces []string - for i := 0; i < 5; i++ { - clientNonce := fmt.Sprintf("client-nonce-%d", i) - msg, err := server.ProcessClientFirstMessage("testuser", clientNonce) - require.NoError(t, err) - nonces = append(nonces, msg.FullNonce) +func TestScramRoundtrip(t *testing.T) { + s, user, pw, _ := setupScram(t) + noErr(t, runHandshake(s, NewScramClient(user, pw)), "handshake") + eq(t, handshakeCount(s), 0, "handshake retained after success") +} + +func TestScramWrongPassword(t *testing.T) { + s, user, _, _ := setupScram(t) + errIs(t, runHandshake(s, NewScramClient(user, "WrongPassword!!!")), ErrInvalidCredentials, "wrong password") + eq(t, handshakeCount(s), 0, "handshake retained after failure") +} + +func TestScramUnknownUser(t *testing.T) { + s, _, _, cred := setupScram(t) + c := NewScramClient("unknown_user", "any_password") + + first, err := c.StartAuthentication() + noErr(t, err, "StartAuthentication") + + serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce) + noErr(t, err, "unknown user must not be signalled at the first message") + + // the decoy must mirror the registered parameter shape + eq(t, serverFirst.ArgonTime, cred.ArgonTime, "decoy time") + eq(t, serverFirst.ArgonMemory, cred.ArgonMemory, "decoy memory") + eq(t, serverFirst.ArgonThreads, cred.ArgonThreads, "decoy threads") + decoySalt, err := base64.StdEncoding.DecodeString(serverFirst.Salt) + noErr(t, err, "decoy salt decode") + eq(t, len(decoySalt), len(cred.Salt), "decoy salt length") + + // stable across probes, distinct per username + second, err := s.ProcessClientFirstMessage("unknown_user", "probe-2") + noErr(t, err, "second probe") + eq(t, second.Salt, serverFirst.Salt, "decoy salt must be deterministic") + third, err := s.ProcessClientFirstMessage("other_unknown", "probe-3") + noErr(t, err, "third probe") + if third.Salt == serverFirst.Salt { + t.Fatal("decoy salt is not username-bound") } - // Verify all handshakes exist - server.mu.RLock() - assert.Len(t, server.handshakes, 5) - server.mu.RUnlock() + // failure surfaces only at the proof step, as for a wrong password + final, err := c.ProcessServerFirstMessage(serverFirst) + noErr(t, err, "ProcessServerFirstMessage") + _, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + errIs(t, err, ErrInvalidCredentials, "unknown user") +} - // Manually set old timestamp for first 3 handshakes - server.mu.Lock() - oldTime := time.Now().Add(-2 * ScramHandshakeTimeout) - count := 0 - for nonce := range server.handshakes { - if count < 3 { - server.handshakes[nonce].CreatedAt = oldTime - count++ +func TestScramDecoySaltIsolation(t *testing.T) { + // decoyKey is per-instance: a shared decoy salt would be a global oracle + // for account existence across a cluster. + a, b := newTestServer(t), newTestServer(t) + first, err := a.ProcessClientFirstMessage("ghost", "n1") + noErr(t, err, "server a") + second, err := b.ProcessClientFirstMessage("ghost", "n2") + noErr(t, err, "server b") + if first.Salt == second.Salt { + t.Fatal("decoy salt is identical across server instances") + } + + // with no credential registered the template is empty and defaults apply + raw, err := base64.StdEncoding.DecodeString(first.Salt) + noErr(t, err, "decode") + eq(t, len(raw), DefaultArgonSaltLen, "fallback salt length") + eq(t, first.ArgonTime, uint32(DefaultArgonTime), "fallback time") + eq(t, first.ArgonMemory, uint32(DefaultArgonMemory), "fallback memory") + eq(t, first.ArgonThreads, uint8(DefaultArgonThreads), "fallback threads") +} + +func TestScramDecoySaltMultiBlock(t *testing.T) { + s := newTestServer(t) + cred := testCredential(t, "u", "SecurePassword123") + cred.Salt = make([]byte, 48) // exceeds one HMAC-SHA256 block + s.AddCredential(cred) + + msg, err := s.ProcessClientFirstMessage("ghost", "n") + noErr(t, err, "first message") + raw, err := base64.StdEncoding.DecodeString(msg.Salt) + noErr(t, err, "decode") + eq(t, len(raw), 48, "decoy salt length") +} + +func TestScramReplayAndUnknownNonce(t *testing.T) { + s, user, pw, _ := setupScram(t) + final := startHandshake(t, s, user, pw) + + _, err := s.ProcessClientFinalMessage("this-is-a-bad-nonce", final.ClientProof) + errIs(t, err, ErrSCRAMInvalidNonce, "unknown nonce") + + _, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + noErr(t, err, "first proof") + + _, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + errIs(t, err, ErrSCRAMInvalidNonce, "replayed proof") +} + +func TestScramProofBinding(t *testing.T) { + // A proof commits to its own auth message; moving it to another live + // handshake for the same user must fail. + s, user, pw, _ := setupScram(t) + a := startHandshake(t, s, user, pw) + b := startHandshake(t, s, user, pw) + + _, err := s.ProcessClientFinalMessage(a.FullNonce, b.ClientProof) + errIs(t, err, ErrInvalidCredentials, "cross-handshake proof") + + // the rejected attempt consumed handshake a but left b intact + _, err = s.ProcessClientFinalMessage(a.FullNonce, a.ClientProof) + errIs(t, err, ErrSCRAMInvalidNonce, "handshake a consumed") + _, err = s.ProcessClientFinalMessage(b.FullNonce, b.ClientProof) + noErr(t, err, "handshake b unaffected") +} + +func TestScramProofEncoding(t *testing.T) { + s, user, pw, _ := setupScram(t) + + cases := []struct { + name string + proof string + want error + }{ + {"not base64", "!!!not base64!!!", ErrSCRAMInvalidProof}, + {"empty", "", ErrSCRAMInvalidProofLen}, + {"short", base64.StdEncoding.EncodeToString(make([]byte, 16)), ErrSCRAMInvalidProofLen}, + {"long", base64.StdEncoding.EncodeToString(make([]byte, sha256.Size+1)), ErrSCRAMInvalidProofLen}, + {"zeroed", base64.StdEncoding.EncodeToString(make([]byte, sha256.Size)), ErrInvalidCredentials}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + final := startHandshake(t, s, user, pw) + _, err := s.ProcessClientFinalMessage(final.FullNonce, tc.proof) + errIs(t, err, tc.want, tc.name) + }) + } +} + +func TestScramVerifyInProgress(t *testing.T) { + s, user, pw, _ := setupScram(t) + final := startHandshake(t, s, user, pw) + + s.mu.RLock() + state := s.handshakes[final.FullNonce] + s.mu.RUnlock() + if state == nil { + t.Fatal("handshake not registered") + } + + state.verifying.Store(1) + _, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + errIs(t, err, ErrSCRAMVerifyInProgress, "concurrent verification") + + // the rejected attempt must not consume the handshake + state.verifying.Store(0) + _, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + noErr(t, err, "retry after release") +} + +func TestScramTimeouts(t *testing.T) { + s, user, pw, _ := setupScram(t) + final := startHandshake(t, s, user, pw) + + s.mu.Lock() + s.handshakes[final.FullNonce].CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout) + s.mu.Unlock() + + _, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + errIs(t, err, ErrSCRAMTimeout, "server-side timeout") + _, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + errIs(t, err, ErrSCRAMInvalidNonce, "expired handshake consumed") + + // client-side clock + c := NewScramClient(user, pw) + _, err = c.StartAuthentication() + noErr(t, err, "StartAuthentication") + c.startTime = time.Now().Add(-2 * ScramHandshakeTimeout) + + _, err = c.ProcessServerFirstMessage(ServerFirstMessage{ + FullNonce: "n", + Salt: base64.StdEncoding.EncodeToString(make([]byte, 16)), + ArgonTime: testArgonTime, + ArgonMemory: testArgonMemory, ArgonThreads: testArgonThreads, + }) + errIs(t, err, ErrSCRAMTimeout, "client timeout on server-first") + + c.authMessage = "seeded" + c.serverKey = make([]byte, sha256.Size) + errIs(t, c.VerifyServerFinalMessage(ServerFinalMessage{}), ErrSCRAMTimeout, "client timeout on server-final") +} + +func TestScramCleanup(t *testing.T) { + s, user, _, _ := setupScram(t) + for i := range 5 { + _, err := s.ProcessClientFirstMessage(user, fmt.Sprintf("client-nonce-%d", i)) + noErr(t, err, "first message") + } + eq(t, handshakeCount(s), 5, "registered handshakes") + + s.mu.Lock() + aged := 0 + for _, state := range s.handshakes { + if aged == 3 { + break + } + state.CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout) + aged++ + } + s.mu.Unlock() + + s.cleanupExpiredHandshakes() + eq(t, handshakeCount(s), 2, "after sweep") + + // a handshake under verification survives the sweep regardless of age + s.mu.Lock() + for _, state := range s.handshakes { + state.CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout) + state.verifying.Store(1) + } + s.mu.Unlock() + + s.cleanupExpiredHandshakes() + eq(t, handshakeCount(s), 2, "verifying handshakes must not be evicted") +} + +func TestScramHandshakeCap(t *testing.T) { + s, user, _, _ := setupScram(t) + for i := range ScramMaxHandshakes { + _, err := s.ProcessClientFirstMessage(user, fmt.Sprintf("n-%d", i)) + noErr(t, err, "first message") + } + eq(t, handshakeCount(s), ScramMaxHandshakes, "at capacity") + + _, err := s.ProcessClientFirstMessage(user, "overflow") + errIs(t, err, ErrSCRAMTooManyHandshakes, "known user at capacity") + + // the cap precedes credential lookup, so it is not an enumeration oracle + _, err = s.ProcessClientFirstMessage("ghost", "overflow") + errIs(t, err, ErrSCRAMTooManyHandshakes, "unknown user at capacity") + + s.mu.Lock() + for _, state := range s.handshakes { + state.CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout) + } + s.mu.Unlock() + + _, err = s.ProcessClientFirstMessage(user, "after-sweep") + noErr(t, err, "capacity reclaimed by the opportunistic sweep") + eq(t, handshakeCount(s), 1, "all expired slots reclaimed") +} + +func TestScramNonceUniqueness(t *testing.T) { + s, user, _, _ := setupScram(t) + seen := make(map[string]struct{}, 256) + for range 256 { + // a fixed client nonce must not produce a fixed full nonce + msg, err := s.ProcessClientFirstMessage(user, "fixed-client-nonce") + noErr(t, err, "first message") + if _, dup := seen[msg.FullNonce]; dup { + t.Fatalf("duplicate full nonce: %s", msg.FullNonce) + } + seen[msg.FullNonce] = struct{}{} + } +} + +func TestScramStopIdempotent(t *testing.T) { + s := NewScramServer() + s.Stop() + s.Stop() // stopOnce must absorb the second close +} + +func TestScramClientState(t *testing.T) { + s, user, pw, _ := setupScram(t) + c := NewScramClient(user, pw) + + errIs(t, c.VerifyServerFinalMessage(ServerFinalMessage{}), ErrSCRAMInvalidState, "unstarted client") + + first, err := c.StartAuthentication() + noErr(t, err, "StartAuthentication") + serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce) + noErr(t, err, "ProcessClientFirstMessage") + final, err := c.ProcessServerFirstMessage(serverFirst) + noErr(t, err, "ProcessServerFirstMessage") + serverFinal, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof) + noErr(t, err, "ProcessClientFinalMessage") + + tampered := serverFinal + tampered.ServerSignature = base64.StdEncoding.EncodeToString(make([]byte, sha256.Size)) + errIs(t, c.VerifyServerFinalMessage(tampered), ErrSCRAMServerAuthFailed, "forged signature") + tampered.ServerSignature = "!!!" + errIs(t, c.VerifyServerFinalMessage(tampered), ErrSCRAMServerAuthFailed, "malformed signature") + noErr(t, c.VerifyServerFinalMessage(serverFinal), "valid signature") + + c.Reset() + errIs(t, c.VerifyServerFinalMessage(serverFinal), ErrSCRAMInvalidState, "after reset") + next, err := c.StartAuthentication() + noErr(t, err, "restart") + if next.ClientNonce == first.ClientNonce { + t.Fatal("client nonce reused after Reset") + } +} + +func TestScramClientRejectsBadServerFirst(t *testing.T) { + c := NewScramClient("u", "SecurePassword123") + _, err := c.StartAuthentication() + noErr(t, err, "StartAuthentication") + + _, err = c.ProcessServerFirstMessage(ServerFirstMessage{ + FullNonce: "n", Salt: "!!!", + ArgonTime: testArgonTime, ArgonMemory: testArgonMemory, ArgonThreads: testArgonThreads, + }) + errIs(t, err, ErrSCRAMInvalidSalt, "salt encoding") + + // ☢ no upper bound is applied to server-supplied cost parameters + good := base64.StdEncoding.EncodeToString(make([]byte, 16)) + for _, msg := range []ServerFirstMessage{ + {FullNonce: "n", Salt: good, ArgonTime: 0, ArgonMemory: testArgonMemory, ArgonThreads: 1}, + {FullNonce: "n", Salt: good, ArgonTime: 1, ArgonMemory: 0, ArgonThreads: 1}, + {FullNonce: "n", Salt: good, ArgonTime: 1, ArgonMemory: testArgonMemory, ArgonThreads: 0}, + } { + _, err = c.ProcessServerFirstMessage(msg) + errIs(t, err, ErrSCRAMZeroParams, "zero parameter") + } + + // A hostile server cannot dictate an unbounded KDF + _, err = c.ProcessServerFirstMessage(ServerFirstMessage{ + FullNonce: "n", Salt: good, + ArgonTime: 1, ArgonMemory: MaxVerifyArgonMemory + 1, ArgonThreads: 1, + }) + errIs(t, err, ErrSCRAMParamsTooLarge, "memory over ceiling") +} + +func TestScramClientOversizedPassword(t *testing.T) { + c := NewScramClient("u", strings.Repeat("a", MaxPasswordLen+1)) + _, err := c.StartAuthentication() + errIs(t, err, ErrPasswordTooLong, "oversized password rejected before the KDF") +} + +func TestServerFirstMessageMarshal(t *testing.T) { + // the auth message binds this exact encoding; changes break every client + msg := ServerFirstMessage{ + FullNonce: "abc", Salt: "c2FsdA==", + ArgonTime: 3, ArgonMemory: 65536, ArgonThreads: 4, + } + eq(t, msg.Marshal(), "r=abc,s=c2FsdA==,t=3,m=65536,p=4", "marshal") +} + +func TestScramMigratedNonStandardDigest(t *testing.T) { + // MigrateFromPHC falls back to DeriveCredential for digests other than 32 + // bytes; the resulting credential must still complete a handshake. + const user, pw = "legacy", "SecurePassword123" + cred, err := MigrateFromPHC(user, pw, phcFor(pw, []byte("0123456789abcdef"), 20)) + noErr(t, err, "MigrateFromPHC") + eq(t, len(cred.StoredKey), sha256.Size, "stored key length") + + s := newTestServer(t) + s.AddCredential(cred) + noErr(t, runHandshake(s, NewScramClient(user, pw)), "handshake with migrated credential") +} + +func TestDeriveCredential(t *testing.T) { + const pw = "SecurePassword123" + salt := make([]byte, 16) + for i := range salt { + salt[i] = byte(i) + } + + first, err := DeriveCredential("u", pw, salt, testArgonTime, testArgonMemory, testArgonThreads) + noErr(t, err, "DeriveCredential") + second, err := DeriveCredential("u", pw, salt, testArgonTime, testArgonMemory, testArgonThreads) + noErr(t, err, "DeriveCredential repeat") + eqBytes(t, first.StoredKey, second.StoredKey, "deterministic stored key") + eqBytes(t, first.ServerKey, second.ServerKey, "deterministic server key") + + salted := argon2.IDKey([]byte(pw), salt, testArgonTime, testArgonMemory, testArgonThreads, DefaultArgonKeyLen) + want := sha256.Sum256(computeHMAC(salted, []byte("Client Key"))) + eqBytes(t, first.StoredKey, want[:], "stored key derivation") + eqBytes(t, first.ServerKey, computeHMAC(salted, []byte("Server Key")), "server key derivation") + if bytes.Equal(first.StoredKey, salted) || bytes.Equal(first.ServerKey, salted) { + t.Fatal("credential exposes the salted password") + } + + // a different password must not collide + other, err := DeriveCredential("u", pw+"x", salt, testArgonTime, testArgonMemory, testArgonThreads) + noErr(t, err, "DeriveCredential other password") + if bytes.Equal(first.StoredKey, other.StoredKey) { + t.Fatal("stored key is independent of the password") + } + + _, err = DeriveCredential("u", pw, make([]byte, 15), testArgonTime, testArgonMemory, testArgonThreads) + errIs(t, err, ErrSCRAMSaltTooShort, "short salt") + + for _, p := range []struct { + time, memory uint32 + threads uint8 + }{{0, testArgonMemory, 1}, {1, 0, 1}, {1, testArgonMemory, 0}} { + _, err = DeriveCredential("u", pw, salt, p.time, p.memory, p.threads) + errIs(t, err, ErrSCRAMZeroParams, "zero parameter") + } + + _, err = DeriveCredential("u", strings.Repeat("a", MaxPasswordLen+1), salt, + testArgonTime, testArgonMemory, testArgonThreads) + errIs(t, err, ErrPasswordTooLong, "oversized password") +} + +func TestCredentialExportImportRoundTrip(t *testing.T) { + cred := testCredential(t, "roundtrip", "SecurePassword123") + + imported, err := ImportCredential(cred.Export()) + noErr(t, err, "ImportCredential") + eq(t, imported.Username, cred.Username, "username") + eqBytes(t, imported.Salt, cred.Salt, "salt") + eq(t, imported.ArgonTime, cred.ArgonTime, "time") + eq(t, imported.ArgonMemory, cred.ArgonMemory, "memory") + eq(t, imported.ArgonThreads, cred.ArgonThreads, "threads") + eqBytes(t, imported.StoredKey, cred.StoredKey, "stored key") + eqBytes(t, imported.ServerKey, cred.ServerKey, "server key") + + // JSON transport converts every number to float64 + raw, err := json.Marshal(cred.Export()) + noErr(t, err, "marshal") + var decoded map[string]any + noErr(t, json.Unmarshal(raw, &decoded), "unmarshal") + viaJSON, err := ImportCredential(decoded) + noErr(t, err, "import via JSON") + eq(t, viaJSON.ArgonMemory, cred.ArgonMemory, "memory via JSON") + eqBytes(t, viaJSON.StoredKey, cred.StoredKey, "stored key via JSON") + + // int-typed input, as produced by YAML and TOML decoders + m := cred.Export() + m["argon_time"] = int(cred.ArgonTime) + m["argon_memory"] = int(cred.ArgonMemory) + m["argon_threads"] = int(cred.ArgonThreads) + viaInt, err := ImportCredential(m) + noErr(t, err, "import from int-typed map") + eq(t, viaInt.ArgonThreads, cred.ArgonThreads, "threads via int") + + // an imported credential must still authenticate + s := newTestServer(t) + s.AddCredential(imported) + noErr(t, runHandshake(s, NewScramClient("roundtrip", "SecurePassword123")), "handshake after import") +} + +func TestImportCredentialErrors(t *testing.T) { + cred := testCredential(t, "u", "SecurePassword123") + with := func(mutate func(map[string]any)) map[string]any { + m := cred.Export() + mutate(m) + return m + } + b64 := func(n int) string { return base64.StdEncoding.EncodeToString(make([]byte, n)) } + + cases := []struct { + name string + data map[string]any + want error + }{ + {"empty map", map[string]any{}, ErrCredMissingUsername}, + {"username wrong type", with(func(m map[string]any) { m["username"] = 42 }), ErrCredMissingUsername}, + + {"salt missing", with(func(m map[string]any) { delete(m, "salt") }), ErrCredMissingSalt}, + {"salt not base64", with(func(m map[string]any) { m["salt"] = "!!!" }), ErrCredInvalidSalt}, + {"salt too short", with(func(m map[string]any) { m["salt"] = b64(15) }), ErrSCRAMSaltTooShort}, + + {"time missing", with(func(m map[string]any) { delete(m, "argon_time") }), ErrCredMissingTime}, + {"time wrong type", with(func(m map[string]any) { m["argon_time"] = "3" }), ErrCredInvalidType}, + {"time fractional", with(func(m map[string]any) { m["argon_time"] = 3.5 }), ErrCredInvalidType}, + {"time negative float", with(func(m map[string]any) { m["argon_time"] = float64(-1) }), ErrCredInvalidType}, + {"time float overflow", with(func(m map[string]any) { m["argon_time"] = float64(math.MaxUint32 + 1) }), ErrCredInvalidType}, + {"time negative int", with(func(m map[string]any) { m["argon_time"] = -1 }), ErrCredInvalidType}, + {"time zero", with(func(m map[string]any) { m["argon_time"] = uint32(0) }), ErrSCRAMZeroParams}, + + {"memory missing", with(func(m map[string]any) { delete(m, "argon_memory") }), ErrCredMissingMemory}, + {"memory zero", with(func(m map[string]any) { m["argon_memory"] = uint32(0) }), ErrSCRAMZeroParams}, + + {"threads missing", with(func(m map[string]any) { delete(m, "argon_threads") }), ErrCredMissingThreads}, + {"threads wrong type", with(func(m map[string]any) { m["argon_threads"] = "4" }), ErrCredInvalidType}, + {"threads overflow", with(func(m map[string]any) { m["argon_threads"] = float64(256) }), ErrCredInvalidType}, + {"threads negative", with(func(m map[string]any) { m["argon_threads"] = -1 }), ErrCredInvalidType}, + {"threads zero", with(func(m map[string]any) { m["argon_threads"] = uint8(0) }), ErrSCRAMZeroParams}, + + {"stored key missing", with(func(m map[string]any) { delete(m, "stored_key") }), ErrCredMissingStoredKey}, + {"stored key not base64", with(func(m map[string]any) { m["stored_key"] = "!!!" }), ErrCredInvalidStoredKey}, + {"stored key short", with(func(m map[string]any) { m["stored_key"] = b64(sha256.Size - 1) }), ErrCredInvalidStoredKey}, + + {"server key missing", with(func(m map[string]any) { delete(m, "server_key") }), ErrCredMissingServerKey}, + {"server key not base64", with(func(m map[string]any) { m["server_key"] = "!!!" }), ErrCredInvalidServerKey}, + {"server key short", with(func(m map[string]any) { m["server_key"] = b64(sha256.Size - 1) }), ErrCredInvalidServerKey}, + + {"stored key long", with(func(m map[string]any) { m["stored_key"] = b64(sha256.Size + 1) }), ErrCredInvalidStoredKey}, + {"server key long", with(func(m map[string]any) { m["server_key"] = b64(sha256.Size + 1) }), ErrCredInvalidServerKey}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ImportCredential(tc.data) + errIs(t, err, tc.want, tc.name) + if got != nil { + t.Fatalf("%s: credential returned alongside error", tc.name) + } + }) + } +} + +func TestScramConcurrentSameUser(t *testing.T) { + s, user, pw, _ := setupScram(t) + + const n = 12 + errs := make(chan error, n) + var wg sync.WaitGroup + for range n { + wg.Add(1) + go func() { + defer wg.Done() + errs <- runHandshake(s, NewScramClient(user, pw)) + }() + } + wg.Wait() + close(errs) + + for err := range errs { + if err != nil { + t.Errorf("concurrent handshake: %v", err) } } - server.mu.Unlock() - - // Trigger cleanup manually - server.cleanupExpiredHandshakes() - - // Verify only 2 handshakes remain - server.mu.RLock() - assert.Len(t, server.handshakes, 2, "Expired handshakes should be cleaned up") - server.mu.RUnlock() + eq(t, handshakeCount(s), 0, "handshakes leaked") } -// TestScramConcurrentSameUser verifies multiple concurrent authentications for same user -func TestScramConcurrentSameUser(t *testing.T) { - server, username, password, _ := setupScramTest(t) - defer server.Stop() - - // Number of concurrent authentication attempts - numAttempts := 10 - results := make(chan error, numAttempts) +func TestScramConcurrentMixedTraffic(t *testing.T) { + s := newTestServer(t) + creds := make([]*Credential, 4) + for i := range creds { + creds[i] = testCredential(t, fmt.Sprintf("user-%d", i), "SecurePassword123") + } var wg sync.WaitGroup - for i := 0; i < numAttempts; i++ { + for _, cred := range creds { wg.Add(1) - go func(attempt int) { + go func() { defer wg.Done() - - // Each goroutine performs full authentication - client := NewScramClient(username, password) - - // Step 1: Client first - clientFirst, err := client.StartAuthentication() - if err != nil { - results <- err - return - } - - // Step 2: Server first - serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) - if err != nil { - results <- err - return - } - - // Step 3: Client final - clientFinal, err := client.ProcessServerFirstMessage(serverFirst) - if err != nil { - results <- err - return - } - - // Step 4: Server final - serverFinal, err := server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof) - if err != nil { - results <- err - return - } - - // Step 5: Client verify - err = client.VerifyServerFinalMessage(serverFinal) - results <- err - }(i) + s.AddCredential(cred) + }() + } + for i := range 16 { + wg.Add(1) + go func() { + defer wg.Done() + // registration races against lookup; both paths take s.mu + _, _ = s.ProcessClientFirstMessage(fmt.Sprintf("user-%d", i%4), fmt.Sprintf("n-%d", i)) + _, _ = s.ProcessClientFirstMessage(fmt.Sprintf("ghost-%d", i), fmt.Sprintf("g-%d", i)) + }() } - wg.Wait() - close(results) +} - // Verify all attempts succeeded - successCount := 0 - for err := range results { - if err == nil { - successCount++ - } else { - t.Logf("Auth attempt failed: %v", err) +func FuzzImportCredential(f *testing.F) { + cred, err := DeriveCredential("u", "SecurePassword123", make([]byte, 16), + testArgonTime, testArgonMemory, testArgonThreads) + if err != nil { + f.Fatal(err) + } + seed, err := json.Marshal(cred.Export()) + if err != nil { + f.Fatal(err) + } + f.Add(seed) + f.Add([]byte(`{}`)) + f.Add([]byte(`{"username":"u","salt":"","argon_time":1e309}`)) + + f.Fuzz(func(t *testing.T, data []byte) { + var m map[string]any + if err := json.Unmarshal(data, &m); err != nil || m == nil { + return + } + got, err := ImportCredential(m) + if err != nil { + if got != nil { + t.Fatal("credential returned alongside error") + } + return + } + if len(got.Salt) < 16 { + t.Fatalf("accepted salt of %d bytes", len(got.Salt)) + } + if len(got.StoredKey) != sha256.Size || len(got.ServerKey) != sha256.Size { + t.Fatalf("accepted keys of %d/%d bytes", len(got.StoredKey), len(got.ServerKey)) + } + if got.ArgonTime == 0 || got.ArgonMemory == 0 || got.ArgonThreads == 0 { + t.Fatalf("accepted zero parameters: %+v", got) + } + }) +} + +func BenchmarkScramHandshake(b *testing.B) { + const user, pw = "bench", "SecurePassword123" + cred, err := DeriveCredential(user, pw, make([]byte, 16), testArgonTime, testArgonMemory, testArgonThreads) + if err != nil { + b.Fatal(err) + } + s := NewScramServer() + defer s.Stop() + s.AddCredential(cred) + + for b.Loop() { + c := NewScramClient(user, pw) + first, err := c.StartAuthentication() + if err != nil { + b.Fatal(err) + } + serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce) + if err != nil { + b.Fatal(err) + } + final, err := c.ProcessServerFirstMessage(serverFirst) + if err != nil { + b.Fatal(err) + } + if _, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof); err != nil { + b.Fatal(err) } } - - assert.Equal(t, numAttempts, successCount, - "All concurrent authentication attempts should succeed") - - // Verify no handshakes remain after completion - server.mu.RLock() - assert.Empty(t, server.handshakes, "All handshakes should be cleaned up after completion") - server.mu.RUnlock() -} - -// TestScramExplicitTimeout verifies timeout enforcement -func TestScramExplicitTimeout(t *testing.T) { - // Save original timeout and set shorter one for testing - originalTimeout := ScramHandshakeTimeout - // Note: Can't modify const at runtime, so we test with delay instead - - server, username, password, _ := setupScramTest(t) - defer server.Stop() - - client := NewScramClient(username, password) - - // Start authentication - clientFirst, err := client.StartAuthentication() - require.NoError(t, err) - - serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) - require.NoError(t, err) - - // Manually expire the handshake - server.mu.Lock() - for nonce := range server.handshakes { - server.handshakes[nonce].CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout) - } - server.mu.Unlock() - - // Client processes server message (should work, client tracks own timeout) - clientFinal, err := client.ProcessServerFirstMessage(serverFirst) - require.NoError(t, err) - - // Server should reject due to timeout - _, err = server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof) - assert.ErrorIs(t, err, ErrSCRAMTimeout, "Server should reject expired handshake") - - // Test client-side timeout - client2 := NewScramClient(username, password) - client2.startTime = time.Now().Add(-2 * ScramHandshakeTimeout) - - _, err = client2.ProcessServerFirstMessage(serverFirst) - assert.ErrorIs(t, err, ErrSCRAMTimeout, "Client should reject after timeout") - - _ = originalTimeout // Suppress unused variable warning } diff --git a/testhelpers_test.go b/testhelpers_test.go new file mode 100644 index 0000000..ad1162f --- /dev/null +++ b/testhelpers_test.go @@ -0,0 +1,85 @@ +package auth + +import ( + "bytes" + "encoding/base64" + "errors" + "fmt" + "testing" + + "golang.org/x/crypto/argon2" +) + +// Assertion helpers. Failures are fatal: most tests below are sequential +// protocol exchanges where continuing past a failure only produces noise. + +func noErr(t *testing.T, err error, ctx string) { + t.Helper() + if err != nil { + t.Fatalf("%s: unexpected error: %v", ctx, err) + } +} + +func hasErr(t *testing.T, err error, ctx string) { + t.Helper() + if err == nil { + t.Fatalf("%s: expected error, got nil", ctx) + } +} + +func errIs(t *testing.T, err, target error, ctx string) { + t.Helper() + if !errors.Is(err, target) { + t.Fatalf("%s: got %v, want %v", ctx, err, target) + } +} + +func eq[T comparable](t *testing.T, got, want T, ctx string) { + t.Helper() + if got != want { + t.Fatalf("%s: got %v, want %v", ctx, got, want) + } +} + +func eqBytes(t *testing.T, got, want []byte, ctx string) { + t.Helper() + if !bytes.Equal(got, want) { + t.Fatalf("%s: got %x, want %x", ctx, got, want) + } +} + +func isTrue(t *testing.T, v bool, ctx string) { + t.Helper() + if !v { + t.Fatalf("%s: got false, want true", ctx) + } +} + +// Test KDF cost. Parameter handling is verified independently of parameter +// magnitude, so every test that does not measure defaults uses these. +const ( + testArgonTime = 1 + testArgonMemory = 8 * 1024 + testArgonThreads = 1 +) + +var cheapArgon = []Option{ + WithTime(testArgonTime), + WithMemory(testArgonMemory), + WithThreads(testArgonThreads), +} + +// rawB64 encodes n zero bytes in the PHC alphabet. +func rawB64(n int) string { + return base64.RawStdEncoding.EncodeToString(make([]byte, n)) +} + +// phcFor builds a PHC record with an arbitrary salt and digest length, which +// HashPassword cannot produce. +func phcFor(password string, salt []byte, keyLen uint32) string { + digest := argon2.IDKey([]byte(password), salt, testArgonTime, testArgonMemory, testArgonThreads, keyLen) + return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", + argon2.Version, testArgonMemory, testArgonTime, testArgonThreads, + base64.RawStdEncoding.EncodeToString(salt), + base64.RawStdEncoding.EncodeToString(digest)) +} diff --git a/token_test.go b/token_test.go index 2ace02c..63696f1 100644 --- a/token_test.go +++ b/token_test.go @@ -2,80 +2,97 @@ package auth import ( "fmt" + "strings" "sync" "testing" - - "github.com/stretchr/testify/assert" ) func TestSimpleTokenValidator(t *testing.T) { - validator := NewSimpleTokenValidator() + v := NewSimpleTokenValidator() + const first, second = "test-token-123", "test-token-456" - token1 := "test-token-123" - token2 := "test-token-456" + isTrue(t, !v.ValidateToken(first), "empty validator rejects") - // Add tokens - validator.AddToken(token1) - validator.AddToken(token2) + v.AddToken(first) + v.AddToken(second) + isTrue(t, v.ValidateToken(first), "first token") + isTrue(t, v.ValidateToken(second), "second token") + isTrue(t, !v.ValidateToken("invalid-token"), "unknown token") - // Validate existing tokens - assert.True(t, validator.ValidateToken(token1)) - assert.True(t, validator.ValidateToken(token2)) + // matching is exact + isTrue(t, !v.ValidateToken(first+"x"), "suffix") + isTrue(t, !v.ValidateToken(first[:len(first)-1]), "prefix") + isTrue(t, !v.ValidateToken(strings.ToUpper(first)), "case") + isTrue(t, !v.ValidateToken(" "+first), "leading space") - // Invalid token - assert.False(t, validator.ValidateToken("invalid-token")) + v.RemoveToken(first) + isTrue(t, !v.ValidateToken(first), "removed token") + isTrue(t, v.ValidateToken(second), "surviving token") - // Remove token - validator.RemoveToken(token1) - assert.False(t, validator.ValidateToken(token1)) - assert.True(t, validator.ValidateToken(token2)) + // removing an absent token is a no-op + v.RemoveToken("never-added") + eq(t, len(v.tokens), 1, "entry count after no-op removal") + + // repeated adds are idempotent + v.AddToken(second) + v.AddToken(second) + eq(t, len(v.tokens), 1, "entry count after duplicate adds") + + // the empty token is storable and matches only itself + v.AddToken("") + isTrue(t, v.ValidateToken(""), "empty token accepted once added") + v.RemoveToken("") + isTrue(t, !v.ValidateToken(""), "empty token removed") } -func TestConcurrentTokenValidator(t *testing.T) { - validator := NewSimpleTokenValidator() +func TestSimpleTokenValidatorKeying(t *testing.T) { + v := NewSimpleTokenValidator() + tokens := []string{ + "", "a", "a\x00b", "a\x00c", "🔑", strings.Repeat("a", 1<<16), + } + for _, tok := range tokens { + v.AddToken(tok) + } + eq(t, len(v.tokens), len(tokens), "distinct entries") - // Add tokens concurrently + for i, tok := range tokens { + isTrue(t, v.ValidateToken(tok), fmt.Sprintf("token %d", i)) + } + for i, tok := range tokens { + v.RemoveToken(tok) + isTrue(t, !v.ValidateToken(tok), fmt.Sprintf("token %d removed", i)) + } + eq(t, len(v.tokens), 0, "empty after removal") +} + +func TestSimpleTokenValidatorConcurrent(t *testing.T) { + v := NewSimpleTokenValidator() + const n = 256 + + // pre-populate half the space so readers see hits and misses + for i := n / 2; i < n; i++ { + v.AddToken(fmt.Sprintf("token-%d", i)) + } + + // each goroutine owns one token, so the final state is deterministic var wg sync.WaitGroup - for i := 0; i < 100; i++ { + for i := range n { wg.Add(1) - go func(idx int) { + go func() { defer wg.Done() - token := fmt.Sprintf("token-%d", idx) - validator.AddToken(token) - }(i) + token := fmt.Sprintf("token-%d", i) + v.AddToken(token) + v.ValidateToken(token) + v.ValidateToken(fmt.Sprintf("absent-%d", i)) + if i%2 == 0 { + v.RemoveToken(token) + } + }() } wg.Wait() - // Validate concurrently - for i := 0; i < 100; i++ { - wg.Add(1) - go func(idx int) { - defer wg.Done() - token := fmt.Sprintf("token-%d", idx) - assert.True(t, validator.ValidateToken(token)) - }(i) - } - wg.Wait() - - // Remove concurrently - for i := 0; i < 50; i++ { - wg.Add(1) - go func(idx int) { - defer wg.Done() - token := fmt.Sprintf("token-%d", idx) - validator.RemoveToken(token) - }(i) - } - wg.Wait() - - // Verify removal - for i := 0; i < 50; i++ { - token := fmt.Sprintf("token-%d", i) - assert.False(t, validator.ValidateToken(token)) - } - for i := 50; i < 100; i++ { - token := fmt.Sprintf("token-%d", i) - assert.True(t, validator.ValidateToken(token)) + for i := range n { + eq(t, v.ValidateToken(fmt.Sprintf("token-%d", i)), i%2 != 0, fmt.Sprintf("token %d", i)) } + eq(t, len(v.tokens), n/2, "surviving entries") } -