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