diff --git a/algorithm/argon2/hasher.go b/algorithm/argon2/hasher.go index 24b0211..d6e50fd 100644 --- a/algorithm/argon2/hasher.go +++ b/algorithm/argon2/hasher.go @@ -90,7 +90,9 @@ func (h *Hasher) Merge(hash *Hasher) { // Hash performs the hashing operation and returns either a argon2.Digest or an error. func (h *Hasher) Hash(password string) (digest algorithm.Digest, err error) { - h.defaults() + if err = h.validate(); err != nil { + return nil, fmt.Errorf(algorithm.ErrFmtHasherHash, AlgName, err) + } if digest, err = h.hash(password); err != nil { return nil, fmt.Errorf(algorithm.ErrFmtHasherHash, AlgName, err) @@ -112,7 +114,9 @@ func (h *Hasher) hash(password string) (hashed algorithm.Digest, err error) { // HashWithSalt overloads the Hash method allowing the user to provide a salt. It's recommended instead to configure the // salt size and let this be a random value generated using crypto/rand. func (h *Hasher) HashWithSalt(password string, salt []byte) (digest algorithm.Digest, err error) { - h.defaults() + if err = h.validate(); err != nil { + return nil, fmt.Errorf(algorithm.ErrFmtHasherHash, AlgName, err) + } if digest, err = h.hashWithSalt(password, salt); err != nil { return nil, fmt.Errorf(algorithm.ErrFmtHasherHash, AlgName, err) @@ -171,10 +175,22 @@ func (h *Hasher) Validate() (err error) { func (h *Hasher) validate() (err error) { h.defaults() - mMin := uint32(h.p) * MemoryMinParallelismMultiplier + // The parallelism and memory are validated using the values Digest.defaults will apply when they're unset, + // otherwise the memory could be rounded below the minimum for the parallelism and produce an undecodable digest. + p, m := uint32(h.p), h.m + + if p < ParallelismMin { + p = ParallelismDefault + } + + if m < MemoryMin { + m = MemoryDefault + } + + mMin := p * MemoryMinParallelismMultiplier - if h.m < mMin || h.m > MemoryMax { - return fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrParameterInvalid, "m", mMin, " (p * 8)", MemoryMax, h.m) + if m < mMin || m > MemoryMax { + return fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrParameterInvalid, "m", mMin, " (p * 8)", MemoryMax, m) } return nil diff --git a/algorithm/argon2/regression_test.go b/algorithm/argon2/regression_test.go index f8f4807..444f75e 100644 --- a/algorithm/argon2/regression_test.go +++ b/algorithm/argon2/regression_test.go @@ -1,6 +1,7 @@ package argon2 import ( + "fmt" "testing" "github.com/stretchr/testify/assert" @@ -82,3 +83,45 @@ func TestDecodeVariantRejectsOtherVariants(t *testing.T) { require.NoError(t, err) assert.Equal(t, encoded, digest.Encode()) } + +func TestNewRejectsMemoryBelowDefaultParallelismMinimum(t *testing.T) { + hasher, err := New(WithM(30)) + + assert.Nil(t, hasher) + assert.EqualError(t, err, "argon2 validation error: parameter is invalid: parameter 'm' must be between 32 (p * 8) and 4294967295 but is set to '30'") +} + +func TestHashRejectsMemoryBelowParallelismMinimumWithoutValidate(t *testing.T) { + hasher := &Hasher{} + + require.NoError(t, hasher.WithOptions(WithM(30))) + + digest, err := hasher.Hash("password") + + assert.Nil(t, digest) + assert.EqualError(t, err, "argon2 hashing error: parameter is invalid: parameter 'm' must be between 32 (p * 8) and 4294967295 but is set to '30'") + + digest, err = hasher.HashWithSalt("password", []byte("saltsaltsaltsalt")) + + assert.Nil(t, digest) + assert.EqualError(t, err, "argon2 hashing error: parameter is invalid: parameter 'm' must be between 32 (p * 8) and 4294967295 but is set to '30'") +} + +func TestHashedDigestsWithDefaultParallelismRoundTrip(t *testing.T) { + for _, m := range []uint32{32, 33, 47, 48, 64} { + t.Run(fmt.Sprintf("m=%d", m), func(t *testing.T) { + hasher, err := New(WithM(m)) + require.NoError(t, err) + + digest, err := hasher.Hash("password") + require.NoError(t, err) + + encoded := digest.Encode() + + decoded, err := Decode(encoded) + require.NoError(t, err, "encoded digest %q could not be decoded", encoded) + + assert.True(t, decoded.Match("password")) + }) + } +}