v0.3.0 security and edge case improvement all around
This commit is contained in:
@@ -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
|
||||||
@@ -42,4 +54,4 @@ server.AddCredential(cred)
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
go test -v ./
|
go test -v ./
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
|||||||
+19
-2
@@ -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 {
|
||||||
@@ -172,4 +179,14 @@ func TestValidatePHCHashFormat(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
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,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
|
||||||
@@ -50,4 +50,4 @@ Utility functions for HTTP headers:
|
|||||||
token, _ := auth.ParseBearerToken(header)
|
token, _ := auth.ParseBearerToken(header)
|
||||||
|
|
||||||
Each module can be used independently without initializing other components.
|
Each module can be used independently without initializing other components.
|
||||||
*/
|
*/
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
// FILE: auth/errors.go
|
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -10,19 +9,19 @@ 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
|
||||||
var (
|
var (
|
||||||
ErrTokenMalformed = errors.New("token: malformed structure")
|
ErrTokenMalformed = errors.New("token: malformed structure")
|
||||||
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")
|
ErrTokenNoPublicKey = errors.New("token: public key required for verification")
|
||||||
ErrTokenNoPublicKey = errors.New("token: public key required for verification")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// JWT secret errors
|
// JWT secret errors
|
||||||
@@ -47,17 +46,17 @@ var (
|
|||||||
|
|
||||||
// SCRAM-specific errors
|
// SCRAM-specific errors
|
||||||
var (
|
var (
|
||||||
ErrSCRAMInvalidNonce = errors.New("scram: invalid nonce or expired handshake")
|
ErrSCRAMInvalidNonce = errors.New("scram: invalid nonce or expired handshake")
|
||||||
ErrSCRAMTimeout = errors.New("scram: handshake timeout")
|
ErrSCRAMTimeout = errors.New("scram: handshake timeout")
|
||||||
ErrSCRAMVerifyInProgress = errors.New("scram: verification already in progress")
|
ErrSCRAMVerifyInProgress = errors.New("scram: verification already in progress")
|
||||||
ErrSCRAMInvalidProof = errors.New("scram: invalid proof encoding")
|
ErrSCRAMInvalidProof = errors.New("scram: invalid proof encoding")
|
||||||
ErrSCRAMInvalidProofLen = errors.New("scram: invalid proof length")
|
ErrSCRAMInvalidProofLen = errors.New("scram: invalid proof length")
|
||||||
ErrSCRAMServerAuthFailed = errors.New("scram: server authentication failed")
|
ErrSCRAMServerAuthFailed = errors.New("scram: server authentication failed")
|
||||||
ErrSCRAMInvalidState = errors.New("scram: invalid handshake state")
|
ErrSCRAMInvalidState = errors.New("scram: invalid handshake state")
|
||||||
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")
|
|
||||||
)
|
|
||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,4 +1,3 @@
|
|||||||
// FILE: auth/http.go
|
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -58,4 +57,5 @@ func ExtractAuthType(header string) string {
|
|||||||
return header[:idx]
|
return header[:idx]
|
||||||
}
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+2
-2
@@ -1,4 +1,3 @@
|
|||||||
// FILE: auth/http_test.go
|
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -51,4 +50,5 @@ func TestHTTPAuthParsing(t *testing.T) {
|
|||||||
_, err = ParseBearerToken("Bearer ")
|
_, err = ParseBearerToken("Bearer ")
|
||||||
assert.Error(t, err)
|
assert.Error(t, err)
|
||||||
assert.Equal(t, ErrAuthEmptyBearerToken, err)
|
assert.Equal(t, ErrAuthEmptyBearerToken, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -312,4 +321,4 @@ func parseRSAPublicKey(pemBytes []byte) (*rsa.PublicKey, error) {
|
|||||||
return nil, ErrRSANotPublicKey
|
return nil, ErrRSANotPublicKey
|
||||||
}
|
}
|
||||||
return pubKey, nil
|
return pubKey, nil
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-2
@@ -1,4 +1,3 @@
|
|||||||
// FILE: auth/jwt_test.go
|
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -256,4 +255,5 @@ 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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() {
|
||||||
close(s.cleanupStop)
|
s.stopOnce.Do(func() {
|
||||||
s.cleanupTicker.Stop()
|
close(s.cleanupStop)
|
||||||
|
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]
|
||||||
Username: username,
|
if !exists {
|
||||||
ClientNonce: clientNonce,
|
t := s.decoyTemplate // mirror real parameter shape
|
||||||
ServerNonce: serverNonce,
|
if t.ArgonTime == 0 {
|
||||||
FullNonce: fullNonce,
|
t.ArgonTime, t.ArgonMemory, t.ArgonThreads = DefaultArgonTime, DefaultArgonMemory, DefaultArgonThreads
|
||||||
Credential: cred,
|
}
|
||||||
CreatedAt: time.Now(),
|
// Deterministic salt + stored decoy handshake so the final
|
||||||
verifying: 0,
|
// step fails with ErrInvalidCredentials, matching the wrong-password path.
|
||||||
|
decoy := &Credential{
|
||||||
|
Username: username,
|
||||||
|
Salt: s.decoySalt(username),
|
||||||
|
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,
|
||||||
|
Salt: base64.StdEncoding.EncodeToString(decoy.Salt),
|
||||||
|
ArgonTime: decoy.ArgonTime,
|
||||||
|
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)
|
|
||||||
}
|
|
||||||
+24
-9
@@ -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.
|
||||||
@@ -316,4 +331,4 @@ func TestScramExplicitTimeout(t *testing.T) {
|
|||||||
assert.ErrorIs(t, err, ErrSCRAMTimeout, "Client should reject after timeout")
|
assert.ErrorIs(t, err, ErrSCRAMTimeout, "Client should reject after timeout")
|
||||||
|
|
||||||
_ = originalTimeout // Suppress unused variable warning
|
_ = originalTimeout // Suppress unused variable warning
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-2
@@ -1,4 +1,3 @@
|
|||||||
// FILE: auth/token_test.go
|
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -78,4 +77,5 @@ func TestConcurrentTokenValidator(t *testing.T) {
|
|||||||
token := fmt.Sprintf("token-%d", i)
|
token := fmt.Sprintf("token-%d", i)
|
||||||
assert.True(t, validator.ValidateToken(token))
|
assert.True(t, validator.ValidateToken(token))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user