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