v0.3.0 security and edge case improvement all around

This commit is contained in:
2026-07-18 08:48:48 -04:00
parent aafa680a35
commit 74434a0c75
15 changed files with 350 additions and 200 deletions
+12
View File
@@ -23,11 +23,23 @@ userID, claims, _ := jwtMgr.ValidateToken(token)
// SCRAM authentication // SCRAM authentication
server := auth.NewScramServer() server := auth.NewScramServer()
defer server.Stop()
phcHash, _ := auth.HashPassword("password123") phcHash, _ := auth.HashPassword("password123")
cred, _ := auth.MigrateFromPHC("user", "password123", phcHash) cred, _ := auth.MigrateFromPHC("user", "password123", phcHash)
server.AddCredential(cred) server.AddCredential(cred)
``` ```
### SCRAM contract notes
- Unknown usernames succeed at `ProcessClientFirstMessage` and fail at
`ProcessClientFinalMessage` with `ErrInvalidCredentials`. This is deliberate
user-enumeration protection. Do not log the first message as an auth success.
- Decoy Argon2 parameters mirror the most recently added credential. Provision
all credentials in a deployment with identical parameters, or the decoy shape
becomes a distinguisher.
- Passwords are bounded by `MaxPasswordLen` (1024 bytes) at every KDF entry
point.
## Package Structure ## Package Structure
- `doc.go` - Overview and package documentation - `doc.go` - Overview and package documentation
+82 -44
View File
@@ -1,9 +1,9 @@
// FILE: auth/argon2.go
package auth package auth
import ( import (
"crypto/hmac"
"crypto/rand" "crypto/rand"
"crypto/subtle" "crypto/sha256"
"encoding/base64" "encoding/base64"
"fmt" "fmt"
"strings" "strings"
@@ -18,6 +18,11 @@ const (
DefaultArgonThreads = 4 DefaultArgonThreads = 4
DefaultArgonSaltLen = 16 DefaultArgonSaltLen = 16
DefaultArgonKeyLen = 32 DefaultArgonKeyLen = 32
MaxPasswordLen = 1024
// upper bounds for untrusted PHC input
MaxArgonSaltLen = 64
MaxArgonKeyLen = 64
MaxPHCHashLen = 256
) )
// argonParams holds configurable Argon2id parameters // argonParams holds configurable Argon2id parameters
@@ -64,6 +69,9 @@ func HashPassword(password string, opts ...Option) (string, error) {
if len(password) < 8 { if len(password) < 8 {
return "", ErrWeakPassword return "", ErrWeakPassword
} }
if len(password) > MaxPasswordLen {
return "", ErrPasswordTooLong
}
params := &argonParams{ params := &argonParams{
time: DefaultArgonTime, time: DefaultArgonTime,
@@ -92,62 +100,50 @@ func HashPassword(password string, opts ...Option) (string, error) {
// VerifyPassword checks password against PHC-format hash (standalone) // VerifyPassword checks password against PHC-format hash (standalone)
func VerifyPassword(password, phcHash string) error { func VerifyPassword(password, phcHash string) error {
parts := strings.Split(phcHash, "$") _, err := verifyPHC(password, phcHash)
if len(parts) != 6 || parts[1] != "argon2id" { return err
return ErrPHCInvalidFormat
}
var memory, time uint32
var threads uint8
fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads)
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
if err != nil {
return fmt.Errorf("%w: %v", ErrPHCInvalidSalt, err)
}
expectedHash, err := base64.RawStdEncoding.DecodeString(parts[5])
if err != nil {
return fmt.Errorf("%w: %v", ErrPHCInvalidHash, err)
}
computedHash := argon2.IDKey([]byte(password), salt, time, memory, threads, uint32(len(expectedHash)))
if subtle.ConstantTimeCompare(computedHash, expectedHash) != 1 {
return ErrInvalidCredentials
}
return nil
} }
// MigrateFromPHC converts PHC hash to SCRAM credential // MigrateFromPHC converts PHC hash to SCRAM credential
func MigrateFromPHC(username, password, phcHash string) (*Credential, error) { func MigrateFromPHC(username, password, phcHash string) (*Credential, error) {
parts := strings.Split(phcHash, "$") r, err := verifyPHC(password, phcHash)
if len(parts) != 6 || parts[1] != "argon2id" {
return nil, ErrPHCInvalidFormat
}
var memory, time uint32
var threads uint8
fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads)
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
if err != nil { if err != nil {
return nil, ErrPHCInvalidSalt
}
// Use standalone function for verification
if err := VerifyPassword(password, phcHash); err != nil {
return nil, err return nil, err
} }
if len(r.derived) == DefaultArgonKeyLen {
return credentialFromSaltedPassword(username, r.derived, r.salt, r.time, r.memory, r.threads), nil
}
// Non-standard digest length: derive at the required key length.
return DeriveCredential(username, password, r.salt, r.time, r.memory, r.threads)
}
return DeriveCredential(username, password, salt, time, memory, threads) // key derivation split from the KDF so callers holding a salted
// password can build a credential without re-running Argon2.
func credentialFromSaltedPassword(username string, saltedPassword, salt []byte, time, memory uint32, threads uint8) *Credential {
clientKey := computeHMAC(saltedPassword, []byte("Client Key"))
serverKey := computeHMAC(saltedPassword, []byte("Server Key"))
storedKey := sha256.Sum256(clientKey)
return &Credential{
Username: username,
Salt: salt,
ArgonTime: time,
ArgonMemory: memory,
ArgonThreads: threads,
StoredKey: storedKey[:],
ServerKey: serverKey,
}
} }
// ValidatePHCHashFormat checks if a hash string has a valid and complete // ValidatePHCHashFormat checks if a hash string has a valid and complete
// 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
if len(phcHash) > MaxPHCHashLen {
return fmt.Errorf("%w: encoded hash exceeds %d bytes", ErrPHCInvalidFormat, MaxPHCHashLen)
}
parts := strings.Split(phcHash, "$") parts := strings.Split(phcHash, "$")
if len(parts) != 6 { if len(parts) != 6 {
return fmt.Errorf("%w: expected 6 parts, got %d", ErrPHCInvalidFormat, len(parts)) return fmt.Errorf("%w: expected 6 parts, got %d", ErrPHCInvalidFormat, len(parts))
@@ -203,6 +199,9 @@ func ValidatePHCHashFormat(phcHash string) error {
if len(salt) < 8 { // Minimum safe salt length if len(salt) < 8 { // Minimum safe salt length
return fmt.Errorf("%w: salt too short (%d bytes)", ErrPHCInvalidSalt, len(salt)) 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 // Validate hash encoding
hash, err := base64.RawStdEncoding.DecodeString(parts[5]) hash, err := base64.RawStdEncoding.DecodeString(parts[5])
@@ -212,6 +211,45 @@ func ValidatePHCHashFormat(phcHash string) error {
if len(hash) < 16 { // Minimum hash length if len(hash) < 16 { // Minimum hash length
return fmt.Errorf("%w: hash too short (%d bytes)", ErrPHCInvalidHash, len(hash)) 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 return nil
} }
// parsed + verified PHC material, reused to avoid a second KDF pass
type phcResult struct {
derived []byte // argon2.IDKey output; == SCRAM salted password when len == DefaultArgonKeyLen
salt []byte
time uint32
memory uint32
threads uint8
}
// verifyPHC validates format, bounds the password, runs the KDF once, and
// constant-time compares against the encoded digest.
func verifyPHC(password, phcHash string) (*phcResult, error) {
if err := ValidatePHCHashFormat(phcHash); err != nil {
return nil, err
}
if len(password) > MaxPasswordLen {
return nil, ErrPasswordTooLong
}
parts := strings.Split(phcHash, "$")
r := &phcResult{}
// Parse is guaranteed well-formed by ValidatePHCHashFormat above.
fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &r.memory, &r.time, &r.threads)
// Encodings validated above; errors are unreachable.
r.salt, _ = base64.RawStdEncoding.DecodeString(parts[4])
expected, _ := base64.RawStdEncoding.DecodeString(parts[5])
r.derived = argon2.IDKey([]byte(password), r.salt, r.time, r.memory, r.threads, uint32(len(expected)))
if !hmac.Equal(r.derived, expected) {
return nil, ErrInvalidCredentials
}
return r, nil
}
+18 -1
View File
@@ -1,4 +1,3 @@
// FILE: auth/argon2_test.go
package auth package auth
import ( import (
@@ -154,6 +153,14 @@ func TestValidatePHCHashFormat(t *testing.T) {
base64.RawStdEncoding.EncodeToString([]byte("short")), ErrPHCInvalidHash}, base64.RawStdEncoding.EncodeToString([]byte("short")), ErrPHCInvalidHash},
{"too few parts", "$argon2id$v=19$m=65536,t=3,p=4", ErrPHCInvalidFormat}, {"too few parts", "$argon2id$v=19$m=65536,t=3,p=4", ErrPHCInvalidFormat},
{"too many parts", "$argon2id$v=19$m=65536,t=3,p=4$salt$hash$extra", ErrPHCInvalidFormat}, {"too many parts", "$argon2id$v=19$m=65536,t=3,p=4$salt$hash$extra", ErrPHCInvalidFormat},
{"oversized salt", "$argon2id$v=19$m=65536,t=3,p=4$" +
base64.RawStdEncoding.EncodeToString(make([]byte, 128)) + "$" +
base64.RawStdEncoding.EncodeToString([]byte("hash1234567890123456")), ErrPHCInvalidSalt},
{"oversized hash", "$argon2id$v=19$m=65536,t=3,p=4$" +
base64.RawStdEncoding.EncodeToString([]byte("salt12345678")) + "$" +
base64.RawStdEncoding.EncodeToString(make([]byte, 128)), ErrPHCInvalidHash},
{"oversized input", "$argon2id$v=19$m=65536,t=3,p=4$" +
strings.Repeat("A", 512) + "$hash", ErrPHCInvalidFormat},
} }
for _, tc := range testCases { for _, tc := range testCases {
@@ -173,3 +180,13 @@ func TestValidatePHCHashFormat(t *testing.T) {
err = VerifyPassword("testPassword123", validHash) err = VerifyPassword("testPassword123", validHash)
assert.NoError(t, err, "Validated hash should still work for password verification") assert.NoError(t, err, "Validated hash should still work for password verification")
} }
func TestVerifyPassword_MalformedParamsNoPanic(t *testing.T) {
for _, h := range []string{
"$argon2id$v=19$m=65536,t=0,p=4$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA",
"$argon2id$v=19$m=65536,t=3,p=0$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA",
"$argon2id$v=19$garbage$c2FsdHNhbHRzYWx0MTI$aGFzaGhhc2hoYXNoaGFzaA",
} {
assert.Error(t, VerifyPassword("whatever", h))
}
}
+1 -1
View File
@@ -1,4 +1,3 @@
// FILE: auth/doc.go
package auth package auth
/* /*
@@ -37,6 +36,7 @@ Server and client implementation for SCRAM:
// Server // Server
server := auth.NewScramServer() server := auth.NewScramServer()
defer server.Stop()
server.AddCredential(credential) server.AddCredential(credential)
// Client // Client
+2 -8
View File
@@ -1,4 +1,3 @@
// FILE: auth/errors.go
package auth package auth
import ( import (
@@ -10,6 +9,7 @@ import (
var ( var (
ErrInvalidCredentials = errors.New("invalid credentials") ErrInvalidCredentials = errors.New("invalid credentials")
ErrWeakPassword = errors.New("password must be at least 8 characters") ErrWeakPassword = errors.New("password must be at least 8 characters")
ErrPasswordTooLong = errors.New("password must be at most 1024 characters")
) )
// JWT-specific errors // JWT-specific errors
@@ -18,7 +18,6 @@ var (
ErrTokenExpired = errors.New("token: expired") ErrTokenExpired = errors.New("token: expired")
ErrTokenNotYetValid = errors.New("token: not yet valid") ErrTokenNotYetValid = errors.New("token: not yet valid")
ErrTokenInvalidSignature = errors.New("token: invalid signature") ErrTokenInvalidSignature = errors.New("token: invalid signature")
ErrTokenAlgorithmMismatch = errors.New("token: algorithm mismatch")
ErrTokenMissingClaim = errors.New("token: missing required claim") ErrTokenMissingClaim = errors.New("token: missing required claim")
ErrTokenEmptyUserID = errors.New("token: empty user ID") ErrTokenEmptyUserID = errors.New("token: empty user ID")
ErrTokenNoPrivateKey = errors.New("token: private key required for signing") ErrTokenNoPrivateKey = errors.New("token: private key required for signing")
@@ -57,7 +56,7 @@ var (
ErrSCRAMInvalidSalt = errors.New("scram: invalid salt encoding") ErrSCRAMInvalidSalt = errors.New("scram: invalid salt encoding")
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")
ErrSCRAMNonceGenFailed = errors.New("scram: failed to generate nonce") ErrSCRAMTooManyHandshakes = errors.New("scram: handshake capacity exceeded")
) )
// Credential import/export errors // Credential import/export errors
@@ -88,8 +87,3 @@ var (
var ( var (
ErrSaltGenerationFailed = errors.New("failed to generate salt") ErrSaltGenerationFailed = errors.New("failed to generate salt")
) )
// Key generation errors
var (
ErrRSAKeyGenFailed = errors.New("failed to generate RSA key")
)
+4 -4
View File
@@ -1,16 +1,16 @@
module github.com/lixenwraith/auth module github.com/lixenwraith/auth
go 1.25.3 go 1.26.0
require ( require (
github.com/golang-jwt/jwt/v5 v5.3.0 github.com/golang-jwt/jwt/v5 v5.3.1
github.com/stretchr/testify v1.11.1 github.com/stretchr/testify v1.11.1
golang.org/x/crypto v0.43.0 golang.org/x/crypto v0.54.0
) )
require ( require (
github.com/davecgh/go-spew v1.1.1 // indirect github.com/davecgh/go-spew v1.1.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
golang.org/x/sys v0.37.0 // indirect golang.org/x/sys v0.47.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect
) )
+14
View File
@@ -2,14 +2,28 @@ 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/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 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.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= 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 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= 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 h1:dduJYIi3A3KOfdGOHX8AVZ/jGiyPa3IbBozJ5kNuE04=
golang.org/x/crypto v0.43.0/go.mod h1:BFbav4mRNlXJL4wNeejLpWxB7wMbc79PdRGhWKncxR0= golang.org/x/crypto v0.43.0/go.mod h1:BFbav4mRNlXJL4wNeejLpWxB7wMbc79PdRGhWKncxR0=
golang.org/x/crypto v0.52.0 h1:RMs7fP2rXdep0CftQlK8Uf+kibLm7qkCcradZWYz988=
golang.org/x/crypto v0.52.0/go.mod h1:1QgfPxDqh0T2M/elOJtp9RvuR95kVjir0e6/BvEmGbc=
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ= golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ=
golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY=
golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= 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 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
+1 -1
View File
@@ -1,4 +1,3 @@
// FILE: auth/http.go
package auth package auth
import ( import (
@@ -59,3 +58,4 @@ func ExtractAuthType(header string) string {
} }
return "" return ""
} }
+1 -1
View File
@@ -1,4 +1,3 @@
// FILE: auth/http_test.go
package auth package auth
import ( import (
@@ -52,3 +51,4 @@ func TestHTTPAuthParsing(t *testing.T) {
assert.Error(t, err) assert.Error(t, err)
assert.Equal(t, ErrAuthEmptyBearerToken, err) assert.Equal(t, ErrAuthEmptyBearerToken, err)
} }
+15 -6
View File
@@ -1,4 +1,3 @@
// FILE: auth/jwt.go
package auth package auth
import ( import (
@@ -206,10 +205,7 @@ func mapJWTError(err error) error {
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)
default: default:
// Check for algorithm mismatch in error message // Alg rejection (WithValidMethods) surfaces as ErrTokenSignatureInvalid.
if errors.Is(err, jwt.ErrTokenSignatureInvalid) {
return fmt.Errorf("%w : %w", ErrTokenAlgorithmMismatch, err)
}
return fmt.Errorf("%w : %w", ErrTokenMalformed, err) return fmt.Errorf("%w : %w", ErrTokenMalformed, err)
} }
} }
@@ -221,12 +217,16 @@ func GenerateHS256Token(secret []byte, userID string, claims map[string]any, lif
if len(secret) < 32 { if len(secret) < 32 {
return "", ErrSecretTooShort return "", ErrSecretTooShort
} }
if userID == "" {
return "", ErrTokenEmptyUserID
}
now := time.Now() now := time.Now()
token := jwt.NewWithClaims(jwt.SigningMethodHS256, customClaims{ token := jwt.NewWithClaims(jwt.SigningMethodHS256, customClaims{
RegisteredClaims: jwt.RegisteredClaims{ RegisteredClaims: jwt.RegisteredClaims{
Subject: userID, Subject: userID,
IssuedAt: jwt.NewNumericDate(now), IssuedAt: jwt.NewNumericDate(now),
NotBefore: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(lifetime)), ExpiresAt: jwt.NewNumericDate(now.Add(lifetime)),
}, },
Extra: claims, Extra: claims,
@@ -244,6 +244,7 @@ func ValidateHS256Token(secret []byte, tokenString string) (string, map[string]a
parser := jwt.NewParser( parser := jwt.NewParser(
jwt.WithValidMethods([]string{"HS256"}), jwt.WithValidMethods([]string{"HS256"}),
jwt.WithLeeway(DefaultLeeway), jwt.WithLeeway(DefaultLeeway),
jwt.WithExpirationRequired(),
) )
token, err := parser.ParseWithClaims(tokenString, &customClaims{}, func(token *jwt.Token) (any, error) { token, err := parser.ParseWithClaims(tokenString, &customClaims{}, func(token *jwt.Token) (any, error) {
@@ -290,10 +291,18 @@ func parseRSAPrivateKey(pemBytes []byte) (*rsa.PrivateKey, error) {
if block == nil { if block == nil {
return nil, ErrRSAInvalidPEM return nil, ErrRSAInvalidPEM
} }
key, err := x509.ParsePKCS1PrivateKey(block.Bytes) if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
return key, nil
}
// PKCS8 fallback
keyAny, err := x509.ParsePKCS8PrivateKey(block.Bytes)
if err != nil { if err != nil {
return nil, ErrRSAInvalidPrivateKey return nil, ErrRSAInvalidPrivateKey
} }
key, ok := keyAny.(*rsa.PrivateKey)
if !ok {
return nil, ErrRSAInvalidPrivateKey
}
return key, nil return key, nil
} }
+1 -1
View File
@@ -1,4 +1,3 @@
// FILE: auth/jwt_test.go
package auth package auth
import ( import (
@@ -257,3 +256,4 @@ func TestJWTRSAFromPEM(t *testing.T) {
_, err = NewJWTVerifierFromPEM([]byte("invalid pem data")) _, err = NewJWTVerifierFromPEM([]byte("invalid pem data"))
assert.ErrorIs(t, err, ErrRSAInvalidPEM) assert.ErrorIs(t, err, ErrRSAInvalidPEM)
} }
+132 -75
View File
@@ -1,4 +1,3 @@
// FILE: auth/scram.go
package auth package auth
import ( import (
@@ -8,6 +7,7 @@ import (
"crypto/subtle" "crypto/subtle"
"encoding/base64" "encoding/base64"
"fmt" "fmt"
"math"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
@@ -21,7 +21,11 @@ const (
// ScramHandshakeTimeout defines maximum time for completing SCRAM handshake // ScramHandshakeTimeout defines maximum time for completing SCRAM handshake
ScramHandshakeTimeout = 30 * time.Second ScramHandshakeTimeout = 30 * time.Second
// ScramCleanupInterval defines how often expired handshakes are cleaned // ScramCleanupInterval defines how often expired handshakes are cleaned
ScramCleanupInterval = 60 * time.Second ScramCleanupInterval = 15 * time.Second
// ScramMaxHandshakes bounds concurrent in-flight handshakes. Caps memory,
// not compute: per client-first server cost is one HMAC. Rate limiting
// upstream remains the control for connection floods.
ScramMaxHandshakes = 4096
) )
// Credential stores SCRAM authentication data // Credential stores SCRAM authentication data
@@ -79,13 +83,20 @@ func ImportCredential(data map[string]any) (*Credential, error) {
} }
switch v := val.(type) { switch v := val.(type) {
case float64: case float64:
// out-of-range float→int conversion is undefined in Go
if v < 0 || v > math.MaxUint32 || v != math.Trunc(v) {
return 0, fmt.Errorf("%w: %s", ErrCredInvalidType, key)
}
return uint32(v), nil return uint32(v), nil
case int: case int:
if v < 0 || int64(v) > math.MaxUint32 {
return 0, fmt.Errorf("%w: %s", ErrCredInvalidType, key)
}
return uint32(v), nil return uint32(v), nil
case uint32: case uint32:
return v, nil return v, nil
default: default:
return 0, fmt.Errorf("invalid type for %s", key) return 0, fmt.Errorf("%w: %s", ErrCredInvalidType, key)
} }
} }
@@ -106,8 +117,14 @@ func ImportCredential(data map[string]any) (*Credential, error) {
var argonThreads uint8 var argonThreads uint8
switch v := threadsVal.(type) { switch v := threadsVal.(type) {
case float64: case float64:
if v < 0 || v > math.MaxUint8 || v != math.Trunc(v) {
return nil, fmt.Errorf("%w: argon_threads", ErrCredInvalidType)
}
argonThreads = uint8(v) argonThreads = uint8(v)
case int: case int:
if v < 0 || v > math.MaxUint8 {
return nil, fmt.Errorf("%w: argon_threads", ErrCredInvalidType)
}
argonThreads = uint8(v) argonThreads = uint8(v)
case uint8: case uint8:
argonThreads = v argonThreads = v
@@ -133,6 +150,17 @@ func ImportCredential(data map[string]any) (*Credential, error) {
return nil, fmt.Errorf("%w: %v", ErrCredInvalidServerKey, err) return nil, fmt.Errorf("%w: %v", ErrCredInvalidServerKey, err)
} }
// Post-decode validation
if argonTime == 0 || argonMemory == 0 || argonThreads == 0 {
return nil, ErrSCRAMZeroParams
}
if len(salt) < 16 {
return nil, ErrSCRAMSaltTooShort
}
if len(storedKey) != sha256.Size || len(serverKey) != sha256.Size {
return nil, ErrCredInvalidStoredKey
}
return &Credential{ return &Credential{
Username: username, Username: username,
Salt: salt, Salt: salt,
@@ -154,23 +182,12 @@ func DeriveCredential(username, password string, salt []byte, time, memory uint3
return nil, ErrSCRAMZeroParams return nil, ErrSCRAMZeroParams
} }
// Derive salted password using Argon2id if len(password) > MaxPasswordLen {
return nil, ErrPasswordTooLong
}
saltedPassword := argon2.IDKey([]byte(password), salt, time, memory, threads, DefaultArgonKeyLen) saltedPassword := argon2.IDKey([]byte(password), salt, time, memory, threads, DefaultArgonKeyLen)
return credentialFromSaltedPassword(username, saltedPassword, salt, time, memory, threads), nil
// Derive keys
clientKey := computeHMAC(saltedPassword, []byte("Client Key"))
serverKey := computeHMAC(saltedPassword, []byte("Server Key"))
storedKey := sha256.Sum256(clientKey)
return &Credential{
Username: username,
Salt: salt,
ArgonTime: time,
ArgonMemory: memory,
ArgonThreads: threads,
StoredKey: storedKey[:],
ServerKey: serverKey,
}, nil
} }
// HandshakeState tracks ongoing authentication // HandshakeState tracks ongoing authentication
@@ -181,37 +198,57 @@ type HandshakeState struct {
FullNonce string FullNonce string
Credential *Credential Credential *Credential
CreatedAt time.Time CreatedAt time.Time
verifying int32 // Atomic flag to prevent race during verification verifying atomic.Int32 // Atomic flag to prevent race during verification
} }
// ScramServer handles server-side SCRAM authentication // ScramServer handles server-side SCRAM authentication
type ScramServer struct { type ScramServer struct {
credentials map[string]*Credential credentials map[string]*Credential
handshakes map[string]*HandshakeState handshakes map[string]*HandshakeState
decoyKey []byte // HMAC key for stable decoy salts
decoyTemplate Credential // param/salt-length shape mirrored to unknown users
mu sync.RWMutex mu sync.RWMutex
cleanupTicker *time.Ticker cleanupTicker *time.Ticker
cleanupStop chan struct{} cleanupStop chan struct{}
stopOnce sync.Once
} }
// NewScramServer creates SCRAM server // NewScramServer creates SCRAM server
func NewScramServer() *ScramServer { func NewScramServer() *ScramServer {
decoyKey := make([]byte, 32)
rand.Read(decoyKey)
s := &ScramServer{ s := &ScramServer{
credentials: make(map[string]*Credential), credentials: make(map[string]*Credential),
handshakes: make(map[string]*HandshakeState), handshakes: make(map[string]*HandshakeState),
decoyKey: decoyKey,
cleanupTicker: time.NewTicker(ScramCleanupInterval), cleanupTicker: time.NewTicker(ScramCleanupInterval),
cleanupStop: make(chan struct{}), cleanupStop: make(chan struct{}),
} }
// Start background cleanup goroutine
go s.cleanupLoop() go s.cleanupLoop()
return s return s
} }
// decoySalt generates stable decoy salt; indistinguishable across repeated probes
func (s *ScramServer) decoySalt(username string) []byte {
n := len(s.decoyTemplate.Salt)
if n < 16 {
n = DefaultArgonSaltLen
}
out := make([]byte, 0, n)
for i := 0; len(out) < n; i++ {
out = append(out, computeHMAC(s.decoyKey, fmt.Appendf(nil, "%s|%d", username, i))...)
}
return out[:n]
}
// Stop gracefully shuts down the server and cleanup goroutine // Stop gracefully shuts down the server and cleanup goroutine
func (s *ScramServer) Stop() { func (s *ScramServer) Stop() {
s.stopOnce.Do(func() {
close(s.cleanupStop) close(s.cleanupStop)
s.cleanupTicker.Stop() s.cleanupTicker.Stop()
})
} }
// cleanupLoop runs periodic cleanup of expired handshakes // cleanupLoop runs periodic cleanup of expired handshakes
@@ -226,57 +263,86 @@ func (s *ScramServer) cleanupLoop() {
} }
} }
// cleanupExpiredHandshakes removes handshakes older than timeout // locking split from sweep logic
func (s *ScramServer) cleanupExpiredHandshakes() { func (s *ScramServer) cleanupExpiredHandshakes() {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
s.evictExpiredLocked()
}
// evictExpiredLocked removes timed-out handshakes. Caller holds s.mu.
func (s *ScramServer) evictExpiredLocked() {
cutoff := time.Now().Add(-ScramHandshakeTimeout) cutoff := time.Now().Add(-ScramHandshakeTimeout)
for nonce, state := range s.handshakes { for nonce, state := range s.handshakes {
if state.CreatedAt.Before(cutoff) && atomic.LoadInt32(&state.verifying) == 0 { if state.CreatedAt.Before(cutoff) && state.verifying.Load() == 0 {
delete(s.handshakes, nonce) delete(s.handshakes, nonce)
} }
} }
} }
// ProcessClientFirstMessage processes initial auth request // ProcessClientFirstMessage processes initial auth request
//
// An unknown username does NOT produce an error here. The server returns
// a deterministic decoy salt and stores a decoy handshake so that failure
// surfaces only at ProcessClientFinalMessage as ErrInvalidCredentials, matching
// the wrong-password path. Callers must not treat a successful return as
// evidence that the account exists, and must not log it as an auth success.
//
// ErrSCRAMTooManyHandshakes is returned when the in-flight handshake cap is
// reached; the cap is applied before credential lookup so the rejection path is
// identical for known and unknown users.
func (s *ScramServer) ProcessClientFirstMessage(username, clientNonce string) (ServerFirstMessage, error) { func (s *ScramServer) ProcessClientFirstMessage(username, clientNonce string) (ServerFirstMessage, error) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
// Check if user exists // opportunistic sweep, then hard cap. Applied before the credential
cred, exists := s.credentials[username] // lookup so the rejection path is identical for known and unknown users.
if !exists { if len(s.handshakes) >= ScramMaxHandshakes {
// Prevent user enumeration - still generate response s.evictExpiredLocked()
salt := make([]byte, 16) if len(s.handshakes) >= ScramMaxHandshakes {
rand.Read(salt) return ServerFirstMessage{}, ErrSCRAMTooManyHandshakes
serverNonce := generateNonce() }
return ServerFirstMessage{
FullNonce: clientNonce + serverNonce,
Salt: base64.StdEncoding.EncodeToString(salt),
ArgonTime: DefaultArgonTime,
ArgonMemory: DefaultArgonMemory,
ArgonThreads: DefaultArgonThreads,
}, ErrInvalidCredentials
} }
// Generate server nonce // Generate server nonce
serverNonce := generateNonce() serverNonce := rand.Text()
fullNonce := clientNonce + serverNonce fullNonce := clientNonce + serverNonce
// Store handshake state // Check if user exists
state := &HandshakeState{ cred, exists := s.credentials[username]
if !exists {
t := s.decoyTemplate // mirror real parameter shape
if t.ArgonTime == 0 {
t.ArgonTime, t.ArgonMemory, t.ArgonThreads = DefaultArgonTime, DefaultArgonMemory, DefaultArgonThreads
}
// Deterministic salt + stored decoy handshake so the final
// step fails with ErrInvalidCredentials, matching the wrong-password path.
decoy := &Credential{
Username: username, Username: username,
ClientNonce: clientNonce, Salt: s.decoySalt(username),
ServerNonce: serverNonce, ArgonTime: t.ArgonTime,
ArgonMemory: t.ArgonMemory,
ArgonThreads: t.ArgonThreads,
StoredKey: make([]byte, sha256.Size), // never matches a real proof
ServerKey: make([]byte, sha256.Size),
}
s.handshakes[fullNonce] = &HandshakeState{
Username: username, ClientNonce: clientNonce, ServerNonce: serverNonce,
FullNonce: fullNonce, Credential: decoy, CreatedAt: time.Now(),
}
return ServerFirstMessage{
FullNonce: fullNonce, FullNonce: fullNonce,
Credential: cred, Salt: base64.StdEncoding.EncodeToString(decoy.Salt),
CreatedAt: time.Now(), ArgonTime: decoy.ArgonTime,
verifying: 0, ArgonMemory: decoy.ArgonMemory,
ArgonThreads: decoy.ArgonThreads,
}, nil // No early error → same control flow as valid user
} }
s.handshakes[fullNonce] = state
s.handshakes[fullNonce] = &HandshakeState{
Username: username, ClientNonce: clientNonce, ServerNonce: serverNonce,
FullNonce: fullNonce, Credential: cred, CreatedAt: time.Now(),
}
return ServerFirstMessage{ return ServerFirstMessage{
FullNonce: fullNonce, FullNonce: fullNonce,
Salt: base64.StdEncoding.EncodeToString(cred.Salt), Salt: base64.StdEncoding.EncodeToString(cred.Salt),
@@ -288,20 +354,21 @@ func (s *ScramServer) ProcessClientFirstMessage(username, clientNonce string) (S
// ProcessClientFinalMessage verifies client proof // ProcessClientFinalMessage verifies client proof
func (s *ScramServer) ProcessClientFinalMessage(fullNonce, clientProof string) (ServerFinalMessage, error) { func (s *ScramServer) ProcessClientFinalMessage(fullNonce, clientProof string) (ServerFinalMessage, error) {
s.mu.RLock() // ookup + CAS under one write lock; closes the sweep race
s.mu.Lock()
state, exists := s.handshakes[fullNonce] state, exists := s.handshakes[fullNonce]
s.mu.RUnlock()
if !exists { if !exists {
s.mu.Unlock()
return ServerFinalMessage{}, ErrSCRAMInvalidNonce return ServerFinalMessage{}, ErrSCRAMInvalidNonce
} }
ok := state.verifying.CompareAndSwap(0, 1)
// Mark as verifying to prevent deletion race s.mu.Unlock()
if !atomic.CompareAndSwapInt32(&state.verifying, 0, 1) { if !ok {
return ServerFinalMessage{}, ErrSCRAMVerifyInProgress return ServerFinalMessage{}, ErrSCRAMVerifyInProgress
} }
defer func() { defer func() {
atomic.StoreInt32(&state.verifying, 0) state.verifying.Store(0)
// Safe to delete after verification completes // Safe to delete after verification completes
s.mu.Lock() s.mu.Lock()
delete(s.handshakes, fullNonce) delete(s.handshakes, fullNonce)
@@ -313,7 +380,6 @@ func (s *ScramServer) ProcessClientFinalMessage(fullNonce, clientProof string) (
return ServerFinalMessage{}, ErrSCRAMTimeout return ServerFinalMessage{}, ErrSCRAMTimeout
} }
// [rest of verification logic unchanged]
// Decode client proof // Decode client proof
clientProofBytes, err := base64.StdEncoding.DecodeString(clientProof) clientProofBytes, err := base64.StdEncoding.DecodeString(clientProof)
if err != nil { if err != nil {
@@ -361,14 +427,11 @@ func (s *ScramServer) AddCredential(cred *Credential) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
s.credentials[cred.Username] = cred s.credentials[cred.Username] = cred
} s.decoyTemplate = Credential{
Salt: make([]byte, len(cred.Salt)),
func (s *ScramServer) cleanupHandshakes() { ArgonTime: cred.ArgonTime,
cutoff := time.Now().Add(-60 * time.Second) ArgonMemory: cred.ArgonMemory,
for nonce, state := range s.handshakes { ArgonThreads: cred.ArgonThreads,
if state.CreatedAt.Before(cutoff) && atomic.LoadInt32(&state.verifying) == 0 {
delete(s.handshakes, nonce)
}
} }
} }
@@ -393,14 +456,15 @@ func NewScramClient(username, password string) *ScramClient {
// StartAuthentication generates initial client message // StartAuthentication generates initial client message
func (c *ScramClient) StartAuthentication() (ClientFirstRequest, error) { func (c *ScramClient) StartAuthentication() (ClientFirstRequest, error) {
// Reject oversized password before the handshake commits to a KDF pass
if len(c.Password) > MaxPasswordLen {
return ClientFirstRequest{}, ErrPasswordTooLong
}
c.startTime = time.Now() c.startTime = time.Now()
// Generate client nonce // Generate client nonce
nonce := make([]byte, 32) c.clientNonce = rand.Text()
if _, err := rand.Read(nonce); err != nil {
return ClientFirstRequest{}, ErrSCRAMNonceGenFailed
}
c.clientNonce = base64.StdEncoding.EncodeToString(nonce)
return ClientFirstRequest{ return ClientFirstRequest{
Username: c.Username, Username: c.Username,
@@ -460,7 +524,6 @@ func (c *ScramClient) VerifyServerFinalMessage(msg ServerFinalMessage) error {
return ErrSCRAMTimeout return ErrSCRAMTimeout
} }
// [rest unchanged]
if c.authMessage == "" || c.serverKey == nil { if c.authMessage == "" || c.serverKey == nil {
return ErrSCRAMInvalidState return ErrSCRAMInvalidState
} }
@@ -537,9 +600,3 @@ func xorBytes(a, b []byte) []byte {
} }
return result return result
} }
func generateNonce() string {
b := make([]byte, 32)
rand.Read(b)
return base64.StdEncoding.EncodeToString(b)
}
+23 -8
View File
@@ -1,4 +1,3 @@
// FILE: auth/scram_test.go
package auth package auth
import ( import (
@@ -27,6 +26,7 @@ func setupScramTest(t *testing.T) (server *ScramServer, username, password strin
// 3. Create a server and add the new credential. // 3. Create a server and add the new credential.
server = NewScramServer() server = NewScramServer()
t.Cleanup(server.Stop)
server.AddCredential(cred) server.AddCredential(cred)
return server, username, password, cred return server, username, password, cred
@@ -64,6 +64,7 @@ func TestScram_FullRoundtrip_Success(t *testing.T) {
// TestScram_FullRoundtrip_WrongPassword ensures authentication fails with an incorrect password. // TestScram_FullRoundtrip_WrongPassword ensures authentication fails with an incorrect password.
func TestScram_FullRoundtrip_WrongPassword(t *testing.T) { func TestScram_FullRoundtrip_WrongPassword(t *testing.T) {
server, username, _, _ := setupScramTest(t) server, username, _, _ := setupScramTest(t)
defer server.Stop()
// Create a client with the WRONG password // Create a client with the WRONG password
client := NewScramClient(username, "WrongPassword!!!") client := NewScramClient(username, "WrongPassword!!!")
@@ -85,22 +86,36 @@ func TestScram_FullRoundtrip_WrongPassword(t *testing.T) {
} }
// TestScram_FullRoundtrip_UserNotFound tests for user enumeration protection. // TestScram_FullRoundtrip_UserNotFound tests for user enumeration protection.
// The server should not reveal whether a user exists or not in its first message. // The server must be indistinguishable from the wrong-password path: no error at
// first message, stable decoy salt across probes, ErrInvalidCredentials at proof.
func TestScram_FullRoundtrip_UserNotFound(t *testing.T) { func TestScram_FullRoundtrip_UserNotFound(t *testing.T) {
server, _, _, _ := setupScramTest(t) server, _, _, _ := setupScramTest(t)
defer server.Stop()
client := NewScramClient("unknown_user", "any_password") client := NewScramClient("unknown_user", "any_password")
clientFirst, err := client.StartAuthentication() clientFirst, err := client.StartAuthentication()
require.NoError(t, err) require.NoError(t, err)
// --- Step 2: Server should return an error, but also a FAKE response --- // unknown user must not error here
// This prevents an attacker from knowing if the user exists based on the response structure.
serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce) serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
assert.ErrorIs(t, err, ErrInvalidCredentials, "Server should return an error for an unknown user") require.NoError(t, err, "unknown user must not be signalled at first message")
assert.NotEmpty(t, serverFirst.FullNonce, "Server must still provide a nonce to prevent enumeration") assert.NotEmpty(t, serverFirst.FullNonce)
assert.NotEmpty(t, serverFirst.Salt, "Server must still provide a salt to prevent enumeration") assert.NotEmpty(t, serverFirst.Salt)
t.Log("SCRAM correctly protected against user enumeration") // 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. // TestScram_InvalidNonce simulates a replay attack or message mismatch.
+10 -16
View File
@@ -1,50 +1,44 @@
// FILE: auth/token.go
package auth package auth
import ( import (
"crypto/subtle" "crypto/sha256"
"sync" "sync"
) )
// SimpleTokenValidator implements in-memory token validation // SimpleTokenValidator implements in-memory token validation
type SimpleTokenValidator struct { type SimpleTokenValidator struct {
tokens map[string]struct{} tokens map[[32]byte]struct{} // keyed by SHA-256(token)
mu sync.RWMutex mu sync.RWMutex
} }
// NewSimpleTokenValidator creates token validator // NewSimpleTokenValidator creates token validator
func NewSimpleTokenValidator() *SimpleTokenValidator { func NewSimpleTokenValidator() *SimpleTokenValidator {
return &SimpleTokenValidator{ return &SimpleTokenValidator{
tokens: make(map[string]struct{}), tokens: make(map[[32]byte]struct{}),
} }
} }
// ValidateToken checks if token is valid // ValidateToken checks if token is valid
func (v *SimpleTokenValidator) ValidateToken(token string) bool { func (v *SimpleTokenValidator) ValidateToken(token string) bool {
h := sha256.Sum256([]byte(token))
v.mu.RLock() v.mu.RLock()
defer v.mu.RUnlock() defer v.mu.RUnlock()
_, ok := v.tokens[h]
// Constant-time comparison for each stored token return ok
for storedToken := range v.tokens {
if subtle.ConstantTimeEq(int32(len(token)), int32(len(storedToken))) == 1 {
if subtle.ConstantTimeCompare([]byte(token), []byte(storedToken)) == 1 {
return true
}
}
}
return false
} }
// AddToken adds token to validator // AddToken adds token to validator
func (v *SimpleTokenValidator) AddToken(token string) { func (v *SimpleTokenValidator) AddToken(token string) {
h := sha256.Sum256([]byte(token))
v.mu.Lock() v.mu.Lock()
defer v.mu.Unlock() defer v.mu.Unlock()
v.tokens[token] = struct{}{} v.tokens[h] = struct{}{}
} }
// RemoveToken removes token from validator // RemoveToken removes token from validator
func (v *SimpleTokenValidator) RemoveToken(token string) { func (v *SimpleTokenValidator) RemoveToken(token string) {
h := sha256.Sum256([]byte(token))
v.mu.Lock() v.mu.Lock()
defer v.mu.Unlock() defer v.mu.Unlock()
delete(v.tokens, token) delete(v.tokens, h)
} }
+1 -1
View File
@@ -1,4 +1,3 @@
// FILE: auth/token_test.go
package auth package auth
import ( import (
@@ -79,3 +78,4 @@ func TestConcurrentTokenValidator(t *testing.T) {
assert.True(t, validator.ValidateToken(token)) assert.True(t, validator.ValidateToken(token))
} }
} }