v0.3.1 minor refactor, tests changed to standard library

This commit is contained in:
2026-07-18 17:09:09 -04:00
parent 74434a0c75
commit 4b04334797
15 changed files with 2051 additions and 946 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
BSD 3-Clause License 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 Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met: modification, are permitted provided that the following conditions are met:
+10 -1
View File
@@ -53,5 +53,14 @@ server.AddCredential(cred)
## Testing ## Testing
```bash ```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
``` ```
+119 -99
View File
@@ -1,9 +1,9 @@
package auth package auth
import ( import (
"crypto/hmac"
"crypto/rand" "crypto/rand"
"crypto/sha256" "crypto/sha256"
"crypto/subtle"
"encoding/base64" "encoding/base64"
"fmt" "fmt"
"strings" "strings"
@@ -25,6 +25,17 @@ const (
MaxPHCHashLen = 256 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 // argonParams holds configurable Argon2id parameters
type argonParams struct { type argonParams struct {
time uint32 time uint32
@@ -86,9 +97,7 @@ func HashPassword(password string, opts ...Option) (string, error) {
} }
salt := make([]byte, params.saltLen) salt := make([]byte, params.saltLen)
if _, err := rand.Read(salt); err != nil { rand.Read(salt) // cryptographically secure random bytes
return "", fmt.Errorf("%w: %v", ErrSaltGenerationFailed, err)
}
hash := argon2.IDKey([]byte(password), salt, params.time, params.memory, params.threads, params.keyLen) 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, // PHC format for Argon2id. It validates structure, parameters, and encoding,
// but does not verify a password against the hash. // but does not verify a password against the hash.
func ValidatePHCHashFormat(phcHash string) error { func ValidatePHCHashFormat(phcHash string) error {
// Cap total input before any splitting or base64 decoding _, err := parsePHC(phcHash)
if len(phcHash) > MaxPHCHashLen { return err
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
} }
// parsed + verified PHC material, reused to avoid a second KDF pass // parsed + verified PHC material, reused to avoid a second KDF pass
type phcResult struct { type phcResult struct {
derived []byte // argon2.IDKey output; == SCRAM salted password when len == DefaultArgonKeyLen derived []byte
salt []byte expectedHash []byte
time uint32 salt []byte
memory uint32 time uint32
threads uint8 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 // verifyPHC validates format, bounds the password, runs the KDF once, and
// constant-time compares against the encoded digest. // constant-time compares against the encoded digest.
func verifyPHC(password, phcHash string) (*phcResult, error) { func verifyPHC(password, phcHash string) (*phcResult, error) {
if err := ValidatePHCHashFormat(phcHash); err != nil {
return nil, err
}
if len(password) > MaxPasswordLen { if len(password) > MaxPasswordLen {
return nil, ErrPasswordTooLong return nil, ErrPasswordTooLong
} }
parts := strings.Split(phcHash, "$") r, err := parsePHC(phcHash)
if err != nil {
return nil, err
}
r := &phcResult{} // Bound the KDF before it runs; the record is untrusted input
// Parse is guaranteed well-formed by ValidatePHCHashFormat above. if err := checkArgonCost(r.memory, r.time, r.threads); err != nil {
fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &r.memory, &r.time, &r.threads) return nil, err
}
// Encodings validated above; errors are unreachable. r.derived = argon2.IDKey([]byte(password), r.salt, r.time, r.memory, r.threads, uint32(len(r.expectedHash)))
r.salt, _ = base64.RawStdEncoding.DecodeString(parts[4]) if subtle.ConstantTimeCompare(r.derived, r.expectedHash) != 1 {
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) {
return nil, ErrInvalidCredentials return nil, ErrInvalidCredentials
} }
return r, nil 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
}
+382 -147
View File
@@ -1,192 +1,427 @@
package auth package auth
import ( import (
"bytes"
"crypto/sha256"
"encoding/base64" "encoding/base64"
"fmt"
"strings" "strings"
"sync" "sync"
"testing" "testing"
"github.com/stretchr/testify/assert" "golang.org/x/crypto/argon2"
"github.com/stretchr/testify/require"
) )
func TestPasswordHashing(t *testing.T) { func TestHashPasswordEncoding(t *testing.T) {
password := "testPassword123" const pw = "testPassword123"
hash, err := HashPassword(pw)
noErr(t, err, "HashPassword")
// Test hashing with default parameters parts := strings.Split(hash, "$")
hash, err := HashPassword(password) eq(t, len(parts), 6, "field count")
require.NoError(t, err, "Failed to hash password") 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 salt, err := base64.RawStdEncoding.DecodeString(parts[4])
assert.True(t, strings.HasPrefix(hash, "$argon2id$"), noErr(t, err, "salt decode")
"Hash should have argon2id prefix, got: %s", hash) eq(t, len(salt), DefaultArgonSaltLen, "salt length")
// Test verification with correct password digest, err := base64.RawStdEncoding.DecodeString(parts[5])
err = VerifyPassword(password, hash) noErr(t, err, "digest decode")
assert.NoError(t, err, "Failed to verify correct password") eq(t, len(digest), DefaultArgonKeyLen, "digest length")
// Test verification with incorrect password // the encoded digest must be reproducible from the encoded material
err = VerifyPassword("wrongPassword", hash) want := argon2.IDKey([]byte(pw), salt, DefaultArgonTime, DefaultArgonMemory, DefaultArgonThreads, DefaultArgonKeyLen)
assert.Error(t, err, "Verification should fail for incorrect password") eqBytes(t, digest, want, "digest")
assert.Equal(t, ErrInvalidCredentials, err)
// Test weak password isTrue(t, len(hash) <= MaxPHCHashLen, "default encoding within MaxPHCHashLen")
_, err = HashPassword("weak") noErr(t, ValidatePHCHashFormat(hash), "self-validation")
assert.Equal(t, ErrWeakPassword, err, "Should reject weak password") }
// Test with custom options func TestHashPasswordSaltUniqueness(t *testing.T) {
hash, err = HashPassword(password, seen := make(map[string]struct{}, 64)
WithTime(5), for range 64 {
WithMemory(128*1024), h, err := HashPassword("testPassword123", cheapArgon...)
WithThreads(8)) noErr(t, err, "HashPassword")
require.NoError(t, err) salt := strings.Split(h, "$")[4]
if _, dup := seen[salt]; dup {
err = VerifyPassword(password, hash) t.Fatalf("duplicate salt: %s", salt)
assert.NoError(t, err) }
seen[salt] = struct{}{}
// 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 TestEmptyPasswordAfterValidation(t *testing.T) { func TestVerifyPassword(t *testing.T) {
// Empty password should be rejected by length check 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("") _, 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 h8, err := HashPassword("12345678", cheapArgon...)
hash, err := HashPassword("12345678") noErr(t, err, "eight bytes")
require.NoError(t, err) noErr(t, VerifyPassword("12345678", h8), "verify eight bytes")
err = VerifyPassword("12345678", hash) // the minimum is measured in bytes, not runes
assert.NoError(t, err) 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) { func TestHashPasswordOptions(t *testing.T) {
password := "testPassword123" const pw = "testPassword123"
hash, err := HashPassword(password)
require.NoError(t, err)
// Test concurrent verification h, err := HashPassword(pw, WithTime(2), WithMemory(16*1024), WithThreads(2))
var wg sync.WaitGroup noErr(t, err, "custom parameters")
for i := 0; i < 10; i++ { eq(t, strings.Split(h, "$")[3], "m=16384,t=2,p=2", "encoded parameters")
wg.Add(1) noErr(t, VerifyPassword(pw, h), "verify custom parameters")
go func() {
defer wg.Done() // zero values are discarded by the option guards
err := VerifyPassword(password, hash) h, err = HashPassword(pw, WithMemory(testArgonMemory), WithTime(0), WithThreads(0))
assert.NoError(t, err) 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) { func TestValidatePHCHashFormat(t *testing.T) {
// Generate valid hash for testing ver := fmt.Sprintf("v=%d", argon2.Version)
validHash, err := HashPassword("testPassword123") const params = "m=65536,t=3,p=4"
require.NoError(t, err) okSalt, okDigest := rawB64(16), rawB64(32)
// Test valid hash build := func(alg, version, prm, salt, digest string) string {
err = ValidatePHCHashFormat(validHash) return "$" + alg + "$" + version + "$" + prm + "$" + salt + "$" + digest
assert.NoError(t, err, "Valid hash should pass validation") }
std := func(prm, salt, digest string) string { return build("argon2id", ver, prm, salt, digest) }
// Test malformed formats generated, err := HashPassword("testPassword123", cheapArgon...)
testCases := []struct { noErr(t, err, "HashPassword")
name string
hash string // largest structurally valid record: proves MaxPHCHashLen never binds
wantErr error 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}, {"empty", "", ErrPHCInvalidFormat},
{"not PHC format", "plaintext", ErrPHCInvalidFormat}, {"plaintext", "plaintext", ErrPHCInvalidFormat},
{"wrong prefix", "argon2id$v=19$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, {"missing leading separator", "argon2id$" + ver + "$" + params + "$" + okSalt + "$" + okDigest, ErrPHCInvalidFormat},
{"wrong algorithm", "$bcrypt$v=19$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, {"leading garbage", "x" + std(params, okSalt, okDigest), ErrPHCInvalidFormat},
{"missing version", "$argon2id$$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, {"leading space", " " + std(params, okSalt, okDigest), ErrPHCInvalidFormat},
{"wrong version", "$argon2id$v=1$m=65536,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, {"too few fields", "$argon2id$" + ver + "$" + params, ErrPHCInvalidFormat},
{"missing params", "$argon2id$v=19$$salt$hash", ErrPHCInvalidFormat}, {"too many fields", std(params, okSalt, okDigest) + "$extra", ErrPHCInvalidFormat},
{"invalid params format", "$argon2id$v=19$invalid$salt$hash", ErrPHCInvalidFormat}, {"over length cap", strings.Repeat("A", MaxPHCHashLen+1), 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}, {"algorithm bcrypt", build("bcrypt", ver, params, okSalt, okDigest), ErrPHCInvalidFormat},
{"zero threads", "$argon2id$v=19$m=65536,t=3,p=0$salt$hash", ErrPHCInvalidFormat}, {"algorithm argon2i", build("argon2i", ver, params, okSalt, okDigest), ErrPHCInvalidFormat},
{"excessive memory", "$argon2id$v=19$m=5000000,t=3,p=4$salt$hash", ErrPHCInvalidFormat}, {"algorithm argon2d", build("argon2d", ver, params, okSalt, okDigest), ErrPHCInvalidFormat},
{"excessive time", "$argon2id$v=19$m=65536,t=2000,p=4$salt$hash", ErrPHCInvalidFormat}, {"algorithm case", build("ARGON2ID", ver, params, okSalt, okDigest), 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$" + {"version empty", build("argon2id", "", params, okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$!!!invalid!!!", ErrPHCInvalidHash}, {"version malformed", build("argon2id", "version=19", params, okSalt, okDigest), ErrPHCInvalidFormat},
{"short salt", "$argon2id$v=19$m=65536,t=3,p=4$" + {"version leading zero", build("argon2id", "v=019", params, okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString([]byte("short")) + "$" + {"version negative", build("argon2id", "v=-19", params, okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString([]byte("hash1234567890123456")), ErrPHCInvalidSalt}, {"version too low", build("argon2id", "v=18", params, okSalt, okDigest), ErrPHCInvalidFormat},
{"short hash", "$argon2id$v=19$m=65536,t=3,p=4$" + {"version too high", build("argon2id", "v=20", params, okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$" +
base64.RawStdEncoding.EncodeToString([]byte("short")), ErrPHCInvalidHash}, {"parameters empty", std("", okSalt, okDigest), ErrPHCInvalidFormat},
{"too few parts", "$argon2id$v=19$m=65536,t=3,p=4", ErrPHCInvalidFormat}, {"parameters reordered", std("t=3,m=65536,p=4", okSalt, okDigest), ErrPHCInvalidFormat},
{"too many parts", "$argon2id$v=19$m=65536,t=3,p=4$salt$hash$extra", ErrPHCInvalidFormat}, {"parameters spaced", std("m=65536, t=3, p=4", okSalt, okDigest), ErrPHCInvalidFormat},
{"oversized salt", "$argon2id$v=19$m=65536,t=3,p=4$" + {"parameters trailing", std("m=65536,t=3,p=4,x=1", okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString(make([]byte, 128)) + "$" + {"parameters negative", std("m=-1,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString([]byte("hash1234567890123456")), ErrPHCInvalidSalt}, {"parameters overflow", std("m=99999999999999999999,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat},
{"oversized hash", "$argon2id$v=19$m=65536,t=3,p=4$" + {"zero time", std("m=65536,t=0,p=4", okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$" + {"zero memory", std("m=0,t=3,p=4", okSalt, okDigest), ErrPHCInvalidFormat},
base64.RawStdEncoding.EncodeToString(make([]byte, 128)), ErrPHCInvalidHash}, {"zero threads", std("m=65536,t=3,p=0", okSalt, okDigest), ErrPHCInvalidFormat},
{"oversized input", "$argon2id$v=19$m=65536,t=3,p=4$" + {"memory at cap", std("m=4194304,t=3,p=4", okSalt, okDigest), nil},
strings.Repeat("A", 512) + "$hash", ErrPHCInvalidFormat}, {"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) { t.Run(tc.name, func(t *testing.T) {
err := ValidatePHCHashFormat(tc.hash) 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{ for _, h := range []string{
"$argon2id$v=19$m=65536,t=0,p=4$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA", "", "$", "$$$$$", "$argon2id$", strings.Repeat("$", 1000),
"$argon2id$v=19$m=65536,t=3,p=0$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA", "$argon2id$v=19$m=65536,t=0,p=4$" + rawB64(16) + "$" + rawB64(32),
"$argon2id$v=19$garbage$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA", "$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")
}
+2
View File
@@ -42,6 +42,7 @@ var (
ErrPHCInvalidFormat = errors.New("phc: invalid format") ErrPHCInvalidFormat = errors.New("phc: invalid format")
ErrPHCInvalidSalt = errors.New("phc: invalid salt encoding") ErrPHCInvalidSalt = errors.New("phc: invalid salt encoding")
ErrPHCInvalidHash = errors.New("phc: invalid hash encoding") ErrPHCInvalidHash = errors.New("phc: invalid hash encoding")
ErrPHCCostTooHigh = errors.New("phc: cost parameters exceed verification limit")
) )
// SCRAM-specific errors // SCRAM-specific errors
@@ -57,6 +58,7 @@ var (
ErrSCRAMZeroParams = errors.New("scram: invalid Argon2 parameters") ErrSCRAMZeroParams = errors.New("scram: invalid Argon2 parameters")
ErrSCRAMSaltTooShort = errors.New("scram: salt must be at least 16 bytes") ErrSCRAMSaltTooShort = errors.New("scram: salt must be at least 16 bytes")
ErrSCRAMTooManyHandshakes = errors.New("scram: handshake capacity exceeded") ErrSCRAMTooManyHandshakes = errors.New("scram: handshake capacity exceeded")
ErrSCRAMParamsTooLarge = errors.New("scram: Argon2 parameters exceed limit")
) )
// Credential import/export errors // Credential import/export errors
+1 -7
View File
@@ -4,13 +4,7 @@ go 1.26.0
require ( require (
github.com/golang-jwt/jwt/v5 v5.3.1 github.com/golang-jwt/jwt/v5 v5.3.1
github.com/stretchr/testify v1.11.1
golang.org/x/crypto v0.54.0 golang.org/x/crypto v0.54.0
) )
require ( require golang.org/x/sys v0.47.0 // indirect
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
)
-24
View File
@@ -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 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= 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 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= 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 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= 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=
+10 -24
View File
@@ -7,34 +7,30 @@ import (
// ParseBasicAuth extracts username/password from Basic auth header // ParseBasicAuth extracts username/password from Basic auth header
func ParseBasicAuth(header string) (username, password string, err error) { func ParseBasicAuth(header string) (username, password string, err error) {
const prefix = "Basic " encoded, ok := strings.CutPrefix(header, "Basic ")
if !strings.HasPrefix(header, prefix) { if !ok {
return "", "", ErrAuthInvalidBasicFormat return "", "", ErrAuthInvalidBasicFormat
} }
encoded := strings.TrimPrefix(header, prefix)
decoded, err := base64.StdEncoding.DecodeString(encoded) decoded, err := base64.StdEncoding.DecodeString(encoded)
if err != nil { if err != nil {
return "", "", ErrAuthInvalidBasicEncoding return "", "", ErrAuthInvalidBasicEncoding
} }
credentials := string(decoded) username, password, ok = strings.Cut(string(decoded), ":")
idx := strings.IndexByte(credentials, ':') if !ok {
if idx < 0 {
return "", "", ErrAuthInvalidBasicCreds return "", "", ErrAuthInvalidBasicCreds
} }
return credentials[:idx], credentials[idx+1:], nil return username, password, nil
} }
// ParseBearerToken extracts token from Bearer auth header // ParseBearerToken extracts token from Bearer auth header
func ParseBearerToken(header string) (token string, err error) { func ParseBearerToken(header string) (token string, err error) {
const prefix = "Bearer " token, ok := strings.CutPrefix(header, "Bearer ")
if !strings.HasPrefix(header, prefix) { if !ok {
return "", ErrAuthInvalidBearerFormat return "", ErrAuthInvalidBearerFormat
} }
token = strings.TrimPrefix(header, prefix)
if token == "" { if token == "" {
return "", ErrAuthEmptyBearerToken return "", ErrAuthEmptyBearerToken
} }
@@ -44,18 +40,8 @@ func ParseBearerToken(header string) (token string, err error) {
// ExtractAuthType returns authentication type from header // ExtractAuthType returns authentication type from header
func ExtractAuthType(header string) string { func ExtractAuthType(header string) string {
if strings.HasPrefix(header, "Basic ") { if authType, _, ok := strings.Cut(header, " "); ok {
return "Basic" return authType
} }
if strings.HasPrefix(header, "Bearer ") { return "" // Matches original behavior if no space is found or string is empty
return "Bearer"
}
// Extract first word as auth type
idx := strings.IndexByte(header, ' ')
if idx > 0 {
return header[:idx]
}
return ""
} }
+114 -45
View File
@@ -2,53 +2,122 @@ package auth
import ( import (
"encoding/base64" "encoding/base64"
"fmt"
"strings"
"testing" "testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
) )
func TestHTTPAuthParsing(t *testing.T) { func basicHeader(payload string) string {
// Test Basic Auth return "Basic " + base64.StdEncoding.EncodeToString([]byte(payload))
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 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)
}
})
}
+30 -33
View File
@@ -111,6 +111,16 @@ func NewJWTRSA(privateKey *rsa.PrivateKey, opts ...JWTOption) (*JWT, error) {
return j, nil 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) // NewJWTVerifier creates JWT manager for verification only (RS256)
func NewJWTVerifier(publicKey *rsa.PublicKey, opts ...JWTOption) (*JWT, error) { func NewJWTVerifier(publicKey *rsa.PublicKey, opts ...JWTOption) (*JWT, error) {
if publicKey == nil { if publicKey == nil {
@@ -132,6 +142,16 @@ func NewJWTVerifier(publicKey *rsa.PublicKey, opts ...JWTOption) (*JWT, error) {
return j, nil 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 // GenerateToken creates signed JWT with claims
func (j *JWT) GenerateToken(userID string, claims map[string]any) (string, error) { func (j *JWT) GenerateToken(userID string, claims map[string]any) (string, error) {
if userID == "" { if userID == "" {
@@ -191,27 +211,26 @@ func (j *JWT) ValidateToken(tokenString string) (string, map[string]any, error)
func mapJWTError(err error) error { func mapJWTError(err error) error {
switch { switch {
case errors.Is(err, jwt.ErrTokenMalformed): 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): 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): 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): 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): 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): 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): 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: 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 // GenerateHS256Token creates HS256 JWT without manager instance
func GenerateHS256Token(secret []byte, userID string, claims map[string]any, lifetime time.Duration) (string, error) { func GenerateHS256Token(secret []byte, userID string, claims map[string]any, lifetime time.Duration) (string, error) {
if len(secret) < 32 { if len(secret) < 32 {
@@ -263,28 +282,6 @@ func ValidateHS256Token(secret []byte, tokenString string) (string, map[string]a
return claims.Subject, claims.Extra, nil 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. // parseRSAPrivateKey parses a PEM-encoded RSA private key.
func parseRSAPrivateKey(pemBytes []byte) (*rsa.PrivateKey, error) { func parseRSAPrivateKey(pemBytes []byte) (*rsa.PrivateKey, error) {
block, _ := pem.Decode(pemBytes) block, _ := pem.Decode(pemBytes)
+547 -200
View File
@@ -1,259 +1,606 @@
package auth package auth
import ( import (
"bytes"
"crypto/ed25519"
"crypto/hmac"
"crypto/rand" "crypto/rand"
"crypto/rsa" "crypto/rsa"
"crypto/sha256"
"crypto/x509" "crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem" "encoding/pem"
"errors"
"fmt"
"strings" "strings"
"sync"
"testing" "testing"
"time" "time"
"github.com/golang-jwt/jwt/v5" "github.com/golang-jwt/jwt/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
) )
func TestJWTHS256(t *testing.T) { var testSecret = []byte("test-secret-key-must-be-32-bytes")
secret := []byte("test-secret-key-must-be-32-bytes")
jwtMgr, err := NewJWT(secret)
require.NoError(t, err)
userID := "user123" func genRSAKey() *rsa.PrivateKey {
claims := map[string]any{ 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", "email": "test@example.com",
"role": "admin", "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 userID, claims, err := manager.ValidateToken(token)
token, err := jwtMgr.GenerateToken(userID, claims) noErr(t, err, "ValidateToken")
require.NoError(t, err) eq(t, userID, "user123", "user id")
assert.NotEmpty(t, token) 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 // nil claims must not emit an extra object
extractedUserID, extractedClaims, err := jwtMgr.ValidateToken(token) bare, err := manager.GenerateToken("user123", nil)
require.NoError(t, err) 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) func TestJWTSecretLength(t *testing.T) {
assert.Equal(t, "test@example.com", extractedClaims["email"]) _, err := NewJWT(nil)
assert.Equal(t, "admin", extractedClaims["role"]) 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) { func TestJWTRS256(t *testing.T) {
// Generate RSA key pair key := testRSAKey()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
// Test with private key (can sign and verify) signer, err := NewJWTRSA(key)
jwtMgr, err := NewJWTRSA(privateKey) noErr(t, err, "NewJWTRSA")
require.NoError(t, err) token, err := signer.GenerateToken("user456", map[string]any{"scope": "read:all"})
noErr(t, err, "GenerateToken")
userID := "user456" header, _ := jwtParts(t, token)
claims := map[string]any{ eq(t, str(t, header, "alg"), "RS256", "alg")
"scope": "read:all",
}
// Generate token userID, claims, err := signer.ValidateToken(token)
token, err := jwtMgr.GenerateToken(userID, claims) noErr(t, err, "self validation")
require.NoError(t, err) eq(t, userID, "user456", "user id")
assert.NotEmpty(t, token) eq(t, str(t, claims, "scope"), "read:all", "scope claim")
// Validate with same manager verifier, err := NewJWTVerifier(&key.PublicKey)
extractedUserID, extractedClaims, err := jwtMgr.ValidateToken(token) noErr(t, err, "NewJWTVerifier")
require.NoError(t, err) userID, _, err = verifier.ValidateToken(token)
assert.Equal(t, userID, extractedUserID) noErr(t, err, "verifier validation")
assert.Equal(t, "read:all", extractedClaims["scope"]) eq(t, userID, "user456", "user id from verifier")
// Test with verifier only (public key) _, err = verifier.GenerateToken("user456", nil)
verifier, err := NewJWTVerifier(&privateKey.PublicKey) errIs(t, err, ErrTokenNoPrivateKey, "verifier must not sign")
require.NoError(t, err)
// Should validate token // an unrelated key must not verify
extractedUserID, _, err = verifier.ValidateToken(token) foreign, err := NewJWTVerifier(&testRSAKeyAlt().PublicKey)
require.NoError(t, err) noErr(t, err, "NewJWTVerifier foreign")
assert.Equal(t, userID, extractedUserID) _, _, err = foreign.ValidateToken(token)
errIs(t, err, ErrTokenInvalidSignature, "foreign public key")
// Should not generate token _, err = NewJWTRSA(nil)
_, err = verifier.GenerateToken(userID, claims) errIs(t, err, ErrTokenNoPrivateKey, "nil private key")
assert.Equal(t, ErrTokenNoPrivateKey, err) _, err = NewJWTVerifier(nil)
errIs(t, err, ErrTokenNoPublicKey, "nil public key")
} }
func TestJWTOptions(t *testing.T) { func TestJWTAlgorithmEnforcement(t *testing.T) {
secret := []byte("test-secret-key-must-be-32-bytes") key := testRSAKey()
hs, err := NewJWT(testSecret)
noErr(t, err, "NewJWT")
rs, err := NewJWTRSA(key)
noErr(t, err, "NewJWTRSA")
// Test custom lifetime hsToken, err := hs.GenerateToken("u", nil)
jwtMgr, err := NewJWT(secret, noErr(t, err, "HS256 token")
WithTokenLifetime(1*time.Hour), 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"), WithIssuer("test-issuer"),
WithAudience([]string{"api.example.com"}), WithAudience([]string{"api.example.com"}),
) )
require.NoError(t, err) noErr(t, err, "NewJWT")
token, err := jwtMgr.GenerateToken("user1", nil) token, err := manager.GenerateToken("user1", nil)
require.NoError(t, err) noErr(t, err, "GenerateToken")
// Parse token to check claims _, payload := jwtParts(t, token)
parsed, _ := jwt.Parse(token, func(token *jwt.Token) (any, error) { eq(t, str(t, payload, "iss"), "test-issuer", "issuer")
return secret, nil 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) superset := signHS256(t, testSecret, defaultHeader(), map[string]any{
"sub": "u", "iss": "test-issuer", "exp": time.Now().Add(time.Hour).Unix(),
// Check issuer "aud": []string{"other.example.com", "api.example.com"},
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(),
}) })
tokenString, err := token.SignedString(secret) _, _, err = manager.ValidateToken(superset)
require.NoError(t, err) noErr(t, err, "expected audience among others")
// Should fail immediately (not valid yet) // an unconstrained manager imposes neither claim
_, _, err = jwtMgr.ValidateToken(tokenString) plain, err := NewJWT(testSecret)
assert.ErrorIs(t, err, ErrTokenNotYetValid) noErr(t, err, "NewJWT plain")
_, _, err = plain.ValidateToken(token)
// Create manager with leeway noErr(t, err, "unconstrained validation")
jwtMgrWithLeeway, err := NewJWT(secret, WithLeeway(5*time.Second))
require.NoError(t, err)
// Should pass with leeway
_, _, err = jwtMgrWithLeeway.ValidateToken(tokenString)
assert.NoError(t, err)
} }
func TestStandaloneFunctions(t *testing.T) { func TestJWTUnenforcedClaims(t *testing.T) {
secret := []byte("test-secret-key-must-be-32-bytes") // Documented gaps: sub is not required and iat is not verified.
userID := "standalone-user" // Callers must reject an empty user id themselves.
claims := map[string]any{"test": "value"} manager, err := NewJWT(testSecret)
noErr(t, err, "NewJWT")
// Generate token token := signHS256(t, testSecret, defaultHeader(), map[string]any{
token, err := GenerateHS256Token(secret, userID, claims, 1*time.Hour) "exp": time.Now().Add(time.Hour).Unix(),
require.NoError(t, err) "iat": time.Now().Add(24 * time.Hour).Unix(),
})
// Validate token userID, claims, err := manager.ValidateToken(token)
extractedUserID, extractedClaims, err := ValidateHS256Token(secret, token) noErr(t, err, "token without subject")
require.NoError(t, err) eq(t, userID, "", "empty subject accepted")
eq(t, len(claims), 0, "no extra claims")
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)
} }
func TestJWTRSAFromPEM(t *testing.T) { func TestJWTOptionGuards(t *testing.T) {
// 1. Generate a new RSA key pair for this test manager, err := NewJWT(testSecret,
privateKey, err := rsa.GenerateKey(rand.Reader, 2048) WithTokenLifetime(0), WithTokenLifetime(-time.Hour), WithLeeway(-time.Second))
require.NoError(t, err) 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 manager, err = NewJWT(testSecret, WithLeeway(0))
privateKeyPEM := pem.EncodeToMemory(&pem.Block{ noErr(t, err, "NewJWT zero leeway")
Type: "RSA PRIVATE KEY", eq(t, manager.leeway, time.Duration(0), "zero leeway applied")
Bytes: x509.MarshalPKCS1PrivateKey(privateKey),
})
// 3. Encode the public key to PEM format // options apply to every constructor
publicKeyBytes, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey) verifier, err := NewJWTVerifier(&testRSAKey().PublicKey, WithIssuer("iss"), WithLeeway(time.Minute))
require.NoError(t, err) noErr(t, err, "NewJWTVerifier")
publicKeyPEM := pem.EncodeToMemory(&pem.Block{ eq(t, verifier.issuer, "iss", "issuer")
Type: "PUBLIC KEY", eq(t, verifier.leeway, time.Minute, "leeway")
Bytes: publicKeyBytes, }
})
// 4. Test the PEM constructor for the signer func TestJWTStandaloneFunctions(t *testing.T) {
jwtMgr, err := NewJWTRSAFromPEM(privateKeyPEM) token, err := GenerateHS256Token(testSecret, "standalone-user",
require.NoError(t, err) map[string]any{"test": "value", "count": 42}, time.Hour)
noErr(t, err, "GenerateHS256Token")
token, err := jwtMgr.GenerateToken("user-from-pem", nil) userID, claims, err := ValidateHS256Token(testSecret, token)
require.NoError(t, err) noErr(t, err, "ValidateHS256Token")
assert.NotEmpty(t, token) 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 _, _, err = ValidateHS256Token(bytes.Repeat([]byte("x"), 32), token)
verifier, err := NewJWTVerifierFromPEM(publicKeyPEM) errIs(t, err, ErrTokenInvalidSignature, "wrong secret")
require.NoError(t, err)
userID, _, err := verifier.ValidateToken(token) expired, err := GenerateHS256Token(testSecret, "u", nil, -time.Hour)
require.NoError(t, err) noErr(t, err, "expired token")
assert.Equal(t, "user-from-pem", userID) _, _, 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")) _, 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")) _, 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)
}
}
}
+16 -14
View File
@@ -154,12 +154,18 @@ func ImportCredential(data map[string]any) (*Credential, error) {
if argonTime == 0 || argonMemory == 0 || argonThreads == 0 { if argonTime == 0 || argonMemory == 0 || argonThreads == 0 {
return nil, ErrSCRAMZeroParams return nil, ErrSCRAMZeroParams
} }
if err := checkArgonCost(argonMemory, argonTime, argonThreads); err != nil {
return nil, ErrSCRAMParamsTooLarge
}
if len(salt) < 16 { if len(salt) < 16 {
return nil, ErrSCRAMSaltTooShort return nil, ErrSCRAMSaltTooShort
} }
if len(storedKey) != sha256.Size || len(serverKey) != sha256.Size { if len(storedKey) != sha256.Size {
return nil, ErrCredInvalidStoredKey return nil, ErrCredInvalidStoredKey
} }
if len(serverKey) != sha256.Size {
return nil, ErrCredInvalidServerKey
}
return &Credential{ return &Credential{
Username: username, Username: username,
@@ -405,7 +411,8 @@ func (s *ScramServer) ProcessClientFinalMessage(fullNonce, clientProof string) (
if len(clientProofBytes) != len(clientSignature) { if len(clientProofBytes) != len(clientSignature) {
return ServerFinalMessage{}, ErrSCRAMInvalidProofLen return ServerFinalMessage{}, ErrSCRAMInvalidProofLen
} }
clientKey := xorBytes(clientProofBytes, clientSignature) clientKey := make([]byte, len(clientProofBytes))
subtle.XORBytes(clientKey, clientProofBytes, clientSignature)
// Verify by computing StoredKey // Verify by computing StoredKey
computedStoredKey := sha256.Sum256(clientKey) 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 { if msg.ArgonTime == 0 || msg.ArgonMemory == 0 || msg.ArgonThreads == 0 {
return ClientFinalRequest{}, ErrSCRAMZeroParams 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 // Derive keys using Argon2id
saltedPassword := argon2.IDKey([]byte(c.Password), salt, msg.ArgonTime, msg.ArgonMemory, msg.ArgonThreads, 32) 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 // Compute client proof
clientSignature := computeHMAC(storedKey[:], []byte(c.authMessage)) 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 // Store server key for verification
c.serverKey = serverKey c.serverKey = serverKey
@@ -589,14 +602,3 @@ func computeHMAC(key, message []byte) []byte {
mac.Write(message) mac.Write(message)
return mac.Sum(nil) 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
}
+661 -295
View File
@@ -1,334 +1,700 @@
package auth package auth
import ( import (
"bytes"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt" "fmt"
"math"
"strings"
"sync" "sync"
"testing" "testing"
"time" "time"
"github.com/stretchr/testify/assert" "golang.org/x/crypto/argon2"
"github.com/stretchr/testify/require"
) )
// setupScramTest is a helper to initialize a server and user credential for testing. func newTestServer(t *testing.T) *ScramServer {
// It performs the full Argon2 -> SCRAM migration workflow. t.Helper()
func setupScramTest(t *testing.T) (server *ScramServer, username, password string, cred *Credential) { s := NewScramServer()
username = "testuser" t.Cleanup(s.Stop)
password = "SecurePassword123" return s
// 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
} }
// TestScram_FullRoundtrip_Success simulates a complete, successful authentication handshake. func testCredential(t *testing.T, username, password string) *Credential {
func TestScram_FullRoundtrip_Success(t *testing.T) { t.Helper()
server, username, password, _ := setupScramTest(t) phcHash, err := HashPassword(password, cheapArgon...)
client := NewScramClient(username, password) noErr(t, err, "HashPassword")
cred, err := MigrateFromPHC(username, password, phcHash)
// --- Step 1: Client sends its first message --- noErr(t, err, "MigrateFromPHC")
clientFirst, err := client.StartAuthentication() return cred
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")
} }
// TestScram_FullRoundtrip_WrongPassword ensures authentication fails with an incorrect password. func setupScram(t *testing.T) (*ScramServer, string, string, *Credential) {
func TestScram_FullRoundtrip_WrongPassword(t *testing.T) { t.Helper()
server, username, _, _ := setupScramTest(t) const username, password = "testuser", "SecurePassword123"
defer server.Stop() cred := testCredential(t, username, password)
// Create a client with the WRONG password s := newTestServer(t)
client := NewScramClient(username, "WrongPassword!!!") s.AddCredential(cred)
return s, username, password, cred
// 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")
} }
// TestScram_FullRoundtrip_UserNotFound tests for user enumeration protection. func handshakeCount(s *ScramServer) int {
// The server must be indistinguishable from the wrong-password path: no error at s.mu.RLock()
// first message, stable decoy salt across probes, ErrInvalidCredentials at proof. defer s.mu.RUnlock()
func TestScram_FullRoundtrip_UserNotFound(t *testing.T) { return len(s.handshakes)
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")
} }
// TestScram_InvalidNonce simulates a replay attack or message mismatch. // startHandshake drives a fresh client to the point where a proof is pending.
func TestScram_InvalidNonce(t *testing.T) { func startHandshake(t *testing.T, s *ScramServer, username, password string) ClientFinalRequest {
server, username, password, _ := setupScramTest(t) t.Helper()
client := NewScramClient(username, password) c := NewScramClient(username, password)
first, err := c.StartAuthentication()
// Perform the first part of the handshake noErr(t, err, "StartAuthentication")
clientFirst, _ := client.StartAuthentication() serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce)
serverFirst, _ := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) noErr(t, err, "ProcessClientFirstMessage")
clientFinal, _ := client.ProcessServerFirstMessage(serverFirst) final, err := c.ProcessServerFirstMessage(serverFirst)
noErr(t, err, "ProcessServerFirstMessage")
// Attempt to finalize with a completely different nonce return final
_, 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")
} }
// TestScram_CredentialImportExport verifies that credentials can be serialized and deserialized correctly. func runHandshake(s *ScramServer, c *ScramClient) error {
func TestScram_CredentialImportExport(t *testing.T) { first, err := c.StartAuthentication()
_, _, _, originalCred := setupScramTest(t) if err != nil {
return err
// 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"),
} }
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 func TestScramRoundtrip(t *testing.T) {
var nonces []string s, user, pw, _ := setupScram(t)
for i := 0; i < 5; i++ { noErr(t, runHandshake(s, NewScramClient(user, pw)), "handshake")
clientNonce := fmt.Sprintf("client-nonce-%d", i) eq(t, handshakeCount(s), 0, "handshake retained after success")
msg, err := server.ProcessClientFirstMessage("testuser", clientNonce) }
require.NoError(t, err)
nonces = append(nonces, msg.FullNonce) 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 // failure surfaces only at the proof step, as for a wrong password
server.mu.RLock() final, err := c.ProcessServerFirstMessage(serverFirst)
assert.Len(t, server.handshakes, 5) noErr(t, err, "ProcessServerFirstMessage")
server.mu.RUnlock() _, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
errIs(t, err, ErrInvalidCredentials, "unknown user")
}
// Manually set old timestamp for first 3 handshakes func TestScramDecoySaltIsolation(t *testing.T) {
server.mu.Lock() // decoyKey is per-instance: a shared decoy salt would be a global oracle
oldTime := time.Now().Add(-2 * ScramHandshakeTimeout) // for account existence across a cluster.
count := 0 a, b := newTestServer(t), newTestServer(t)
for nonce := range server.handshakes { first, err := a.ProcessClientFirstMessage("ghost", "n1")
if count < 3 { noErr(t, err, "server a")
server.handshakes[nonce].CreatedAt = oldTime second, err := b.ProcessClientFirstMessage("ghost", "n2")
count++ 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() eq(t, handshakeCount(s), 0, "handshakes leaked")
// 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()
} }
// TestScramConcurrentSameUser verifies multiple concurrent authentications for same user func TestScramConcurrentMixedTraffic(t *testing.T) {
func TestScramConcurrentSameUser(t *testing.T) { s := newTestServer(t)
server, username, password, _ := setupScramTest(t) creds := make([]*Credential, 4)
defer server.Stop() for i := range creds {
creds[i] = testCredential(t, fmt.Sprintf("user-%d", i), "SecurePassword123")
// Number of concurrent authentication attempts }
numAttempts := 10
results := make(chan error, numAttempts)
var wg sync.WaitGroup var wg sync.WaitGroup
for i := 0; i < numAttempts; i++ { for _, cred := range creds {
wg.Add(1) wg.Add(1)
go func(attempt int) { go func() {
defer wg.Done() defer wg.Done()
s.AddCredential(cred)
// Each goroutine performs full authentication }()
client := NewScramClient(username, password) }
for i := range 16 {
// Step 1: Client first wg.Add(1)
clientFirst, err := client.StartAuthentication() go func() {
if err != nil { defer wg.Done()
results <- err // registration races against lookup; both paths take s.mu
return _, _ = 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))
}()
// 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)
} }
wg.Wait() wg.Wait()
close(results) }
// Verify all attempts succeeded func FuzzImportCredential(f *testing.F) {
successCount := 0 cred, err := DeriveCredential("u", "SecurePassword123", make([]byte, 16),
for err := range results { testArgonTime, testArgonMemory, testArgonThreads)
if err == nil { if err != nil {
successCount++ f.Fatal(err)
} else { }
t.Logf("Auth attempt failed: %v", 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
} }
+85
View File
@@ -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))
}
+73 -56
View File
@@ -2,80 +2,97 @@ package auth
import ( import (
"fmt" "fmt"
"strings"
"sync" "sync"
"testing" "testing"
"github.com/stretchr/testify/assert"
) )
func TestSimpleTokenValidator(t *testing.T) { func TestSimpleTokenValidator(t *testing.T) {
validator := NewSimpleTokenValidator() v := NewSimpleTokenValidator()
const first, second = "test-token-123", "test-token-456"
token1 := "test-token-123" isTrue(t, !v.ValidateToken(first), "empty validator rejects")
token2 := "test-token-456"
// Add tokens v.AddToken(first)
validator.AddToken(token1) v.AddToken(second)
validator.AddToken(token2) 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 // matching is exact
assert.True(t, validator.ValidateToken(token1)) isTrue(t, !v.ValidateToken(first+"x"), "suffix")
assert.True(t, validator.ValidateToken(token2)) 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 v.RemoveToken(first)
assert.False(t, validator.ValidateToken("invalid-token")) isTrue(t, !v.ValidateToken(first), "removed token")
isTrue(t, v.ValidateToken(second), "surviving token")
// Remove token // removing an absent token is a no-op
validator.RemoveToken(token1) v.RemoveToken("never-added")
assert.False(t, validator.ValidateToken(token1)) eq(t, len(v.tokens), 1, "entry count after no-op removal")
assert.True(t, validator.ValidateToken(token2))
// 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) { func TestSimpleTokenValidatorKeying(t *testing.T) {
validator := NewSimpleTokenValidator() 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 var wg sync.WaitGroup
for i := 0; i < 100; i++ { for i := range n {
wg.Add(1) wg.Add(1)
go func(idx int) { go func() {
defer wg.Done() defer wg.Done()
token := fmt.Sprintf("token-%d", idx) token := fmt.Sprintf("token-%d", i)
validator.AddToken(token) v.AddToken(token)
}(i) v.ValidateToken(token)
v.ValidateToken(fmt.Sprintf("absent-%d", i))
if i%2 == 0 {
v.RemoveToken(token)
}
}()
} }
wg.Wait() wg.Wait()
// Validate concurrently for i := range n {
for i := 0; i < 100; i++ { eq(t, v.ValidateToken(fmt.Sprintf("token-%d", i)), i%2 != 0, fmt.Sprintf("token %d", 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))
} }
eq(t, len(v.tokens), n/2, "surviving entries")
} }