diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 5265baa..c8a28f9 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -31,7 +31,7 @@ jobs: go get -v -t ./... - name: 'Test' run: | - go test -coverprofile=coverage.txt -v ./... + go test -race -coverprofile=coverage.txt -covermode=atomic -v ./... - name: 'Coverage' uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f # v7.0.0 with: @@ -65,4 +65,41 @@ jobs: go build -v ./... - name: 'Test' run: | - go test -v ./... + go test -race -shuffle=on -v ./... + fuzz: + name: 'Fuzz (${{ matrix.target.label }})' + runs-on: 'ubuntu-latest' + strategy: + matrix: + target: + - package: '.' + func: 'FuzzDecode' + label: 'Decode' + - package: '.' + func: 'FuzzNormalize' + label: 'Normalize' + - package: './algorithm/scrypt' + func: 'FuzzDecodeAndMatch' + label: 'Scrypt Decode and Match' + fail-fast: false + steps: + - name: 'Harden Runner' + uses: step-security/harden-runner@05e31511f85b41b11d1cf0ef85d0992719546e2c # v2.21.0 + with: + egress-policy: 'audit' + - name: 'Set up Go' + uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 + with: + go-version: '1.27' + - name: 'Checkout' + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - name: 'Get Dependencies' + run: | + go get -v -t ./... + - name: 'Fuzz ${{ matrix.target.func }}' + run: | + go test -run '^$' -fuzz '^${{ matrix.target.func }}$' -fuzztime 120s ${{ matrix.target.package }} + - name: 'Show Failing Corpus' + if: failure() + run: | + find . -path '*/testdata/fuzz/*' -type f -print -exec cat {} \; diff --git a/.gitignore b/.gitignore index 35c6697..203679f 100644 --- a/.gitignore +++ b/.gitignore @@ -2,4 +2,6 @@ graphify-out/ # Added by ggshield -.cache_ggshield \ No newline at end of file +.cache_ggshield + +coverage.out \ No newline at end of file diff --git a/algorithm/argon2/decoder.go b/algorithm/argon2/decoder.go index 30b3b78..a53fc4e 100644 --- a/algorithm/argon2/decoder.go +++ b/algorithm/argon2/decoder.go @@ -162,16 +162,16 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return nil, fmt.Errorf("%w: key has 0 bytes", algorithm.ErrEncodedHashKeyEncoding) } - if decoded.t == 0 { - decoded.t = 1 + if decoded.t < IterationsMin { + return nil, fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, oT, IterationsMin, "", uint32(IterationsMax), decoded.t) } - if decoded.p == 0 { - decoded.p = 4 + if decoded.p < ParallelismMin || decoded.p > ParallelismMax { + return nil, fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, oP, ParallelismMin, "", uint32(ParallelismMax), decoded.p) } - if decoded.m == 0 { - decoded.m = 32 * 1024 + if decoded.m < MemoryMin { + return nil, fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, oM, MemoryMin, "", MemoryMax, decoded.m) } return decoded, nil diff --git a/algorithm/argon2/regression_test.go b/algorithm/argon2/regression_test.go new file mode 100644 index 0000000..eb11750 --- /dev/null +++ b/algorithm/argon2/regression_test.go @@ -0,0 +1,83 @@ +package argon2 + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDecodeRejectsParametersItCannotHonour(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"ZeroIterations", "$argon2id$v=19$m=8,t=0,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + {"ZeroParallelism", "$argon2id$v=19$m=8,t=1,p=0$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + {"ZeroMemory", "$argon2id$v=19$m=0,t=1,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + {"MemoryBelowMinimum", "$argon2id$v=19$m=1,t=1,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + {"ParallelismAboveMaximum", "$argon2id$v=19$m=8,t=1,p=16777216$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + digest, err := Decode(tc.digest) + + assert.Nil(t, digest) + assert.Error(t, err) + }) + } +} + +func TestDecodePreservesParameters(t *testing.T) { + testCases := []string{ + "$argon2id$v=19$m=8,t=1,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$argon2i$v=19$m=65536,t=3,p=4$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$argon2d$v=19$m=2097152,t=1,p=4$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + } + + for _, encoded := range testCases { + t.Run(encoded, func(t *testing.T) { + digest, err := Decode(encoded) + + require.NoError(t, err) + assert.Equal(t, encoded, digest.Encode()) + }) + } +} + +func TestHashedDigestsRoundTrip(t *testing.T) { + for _, variant := range []Variant{VariantID, VariantI, VariantD} { + t.Run(variant.String(), func(t *testing.T) { + hasher, err := New(WithVariant(variant), WithProfileRFC9106LowMemory()) + require.NoError(t, err) + require.NoError(t, hasher.Validate()) + + 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.Equal(t, encoded, decoded.Encode()) + assert.True(t, decoded.Match("password")) + assert.False(t, decoded.Match("incorrect")) + }) + } +} + +func TestDecodeVariantRejectsOtherVariants(t *testing.T) { + const encoded = "$argon2i$v=19$m=8,t=1,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5" + + digest, err := DecodeVariant(VariantID)(encoded) + + assert.Nil(t, digest) + assert.Error(t, err) + + digest, err = DecodeVariant(VariantI)(encoded) + + require.NoError(t, err) + assert.Equal(t, encoded, digest.Encode()) +} diff --git a/algorithm/bcrypt/const.go b/algorithm/bcrypt/const.go index 7d3d1ba..50dea1e 100644 --- a/algorithm/bcrypt/const.go +++ b/algorithm/bcrypt/const.go @@ -5,8 +5,9 @@ import ( ) const ( - // EncodingFmt is the encoding format for this algorithm. - EncodingFmt = "$%s$%d$%s%s" + // EncodingFmt is the encoding format for this algorithm. The cost is zero padded to two digits as the bcrypt + // modular crypt format always represents it that way, and other implementations reject a single digit cost. + EncodingFmt = "$%s$%02d$%s%s" // EncodingFmtSHA256 is the encoding format for the SHA256 variant of this algorithm. EncodingFmtSHA256 = "$%s$v=2,t=%s,r=%d$%s$%s" diff --git a/algorithm/bcrypt/decoder.go b/algorithm/bcrypt/decoder.go index b3365c2..3dc5446 100644 --- a/algorithm/bcrypt/decoder.go +++ b/algorithm/bcrypt/decoder.go @@ -126,6 +126,10 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return nil, fmt.Errorf("%w: iterations could not be parsed: %v", algorithm.ErrEncodedHashInvalidOptionValue, err) } + if err = validateCost(decoded.iterations); err != nil { + return nil, err + } + switch n, i := len(parts[1]), bcrypt.EncodedSaltSize+bcrypt.EncodedHashSize; n { case i: break @@ -181,6 +185,10 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return nil, fmt.Errorf("%w: option '%s' has invalid value '%s': %v", algorithm.ErrEncodedHashInvalidOptionValue, param.Key, param.Value, err) } } + + if err = validateCost(decoded.iterations); err != nil { + return nil, err + } } if decoded.salt, err = bcrypt.Base64Decode(salt); err != nil { @@ -195,3 +203,11 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return decoded, nil } + +func validateCost(cost int) (err error) { + if cost < bcrypt.MinCost || cost > bcrypt.MaxCost { + return fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, "cost", bcrypt.MinCost, "", bcrypt.MaxCost, cost) + } + + return nil +} diff --git a/algorithm/bcrypt/hasher.go b/algorithm/bcrypt/hasher.go index 3ace1de..4fcbf2f 100644 --- a/algorithm/bcrypt/hasher.go +++ b/algorithm/bcrypt/hasher.go @@ -27,15 +27,7 @@ func New(opts ...Opt) (hasher *Hasher, err error) { // NewSHA256 returns a new bcrypt.Hasher with the provided functional options applied as well as the bcrypt.VariantSHA256 // applied via the bcrypt.WithVariant bcrypt.Opt. func NewSHA256(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantSHA256)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantSHA256)}, opts...)...) } // Hasher is a crypt.Hash for bcrypt which can be initialized via bcrypt.New using a functional options pattern. diff --git a/algorithm/bcrypt/regression_test.go b/algorithm/bcrypt/regression_test.go new file mode 100644 index 0000000..e09b1f7 --- /dev/null +++ b/algorithm/bcrypt/regression_test.go @@ -0,0 +1,108 @@ +package bcrypt + +import ( + "testing" + + xbcrypt "github.com/go-crypt/x/bcrypt" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDecodeRejectsUnusableCost(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"StandardNegative", "$2b$-1$" + validStandardKey}, + {"StandardZero", "$2b$00$" + validStandardKey}, + {"StandardBelowMinimum", "$2b$03$" + validStandardKey}, + {"StandardAboveMaximum", "$2b$99$" + validStandardKey}, + {"SHA256Zero", "$bcrypt-sha256$v=2,t=2b,r=0$" + validSHA256Salt + "$" + validSHA256Key}, + {"SHA256AboveMaximum", "$bcrypt-sha256$v=2,t=2b,r=99$" + validSHA256Salt + "$" + validSHA256Key}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + digest, err := Decode(tc.digest) + + assert.Nil(t, digest) + assert.Error(t, err) + }) + } +} + +func TestDecodeAcceptsLegacyCosts(t *testing.T) { + for _, cost := range []string{"04", "05", "09", "31"} { + t.Run(cost, func(t *testing.T) { + encoded := "$2b$" + cost + "$" + validStandardKey + + digest, err := Decode(encoded) + + require.NoError(t, err) + assert.Equal(t, encoded, digest.Encode()) + }) + } +} + +func TestDecodeAcceptsEveryCostTheKeyDerivationAccepts(t *testing.T) { + assert.NoError(t, validateCost(xbcrypt.MinCost)) + assert.NoError(t, validateCost(xbcrypt.MaxCost)) + assert.Error(t, validateCost(xbcrypt.MinCost-1)) + assert.Error(t, validateCost(xbcrypt.MaxCost+1)) +} + +func TestHashedDigestsRoundTrip(t *testing.T) { + testCases := []struct { + name string + new func() (*Hasher, error) + }{ + {"Standard", func() (*Hasher, error) { return New(WithCost(IterationsMin)) }}, + {"SHA256", func() (*Hasher, error) { return NewSHA256(WithCost(IterationsMin)) }}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + hasher, err := tc.new() + require.NoError(t, err) + require.NoError(t, hasher.Validate()) + + 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.Equal(t, encoded, decoded.Encode()) + assert.True(t, decoded.Match("password")) + assert.False(t, decoded.Match("incorrect")) + }) + } +} + +func TestSHA256VariantIsNotLimitedTo72Bytes(t *testing.T) { + long := make([]byte, 200) + + for i := range long { + long[i] = byte('a' + i%26) + } + + hasher, err := NewSHA256(WithCost(IterationsMin)) + require.NoError(t, err) + + digest, err := hasher.Hash(string(long)) + require.NoError(t, err) + + assert.True(t, digest.MatchBytes(long)) + + truncated := append(append([]byte{}, long[:72]...), []byte("different")...) + + assert.False(t, digest.MatchBytes(truncated)) +} + +const ( + validStandardKey = "3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm" + validSHA256Salt = "3XCpXfcQBjcbXFHTLcbFju" + validSHA256Key = "AXNZ1B7NPTf7XyCqUKcvIUOB5eKKZ4C" +) diff --git a/algorithm/md5crypt/const.go b/algorithm/md5crypt/const.go index 70993ba..7839bf3 100644 --- a/algorithm/md5crypt/const.go +++ b/algorithm/md5crypt/const.go @@ -12,7 +12,7 @@ const ( EncodingFmtSun = "$md5$%s$$%s" // EncodingFmtSunIterations is the encoding format for this algorithm when using md5crypt.VariantSun and iterations more than 0. - EncodingFmtSunIterations = "$md5,iterations=%d$%s$$%s" + EncodingFmtSunIterations = "$md5,rounds=%d$%s$$%s" // AlgName is the name for this algorithm. AlgName = "md5crypt" @@ -23,6 +23,13 @@ const ( // AlgIdentifierVariantSun is the identifier used in this algorithm when using md5crypt.VariantSun. AlgIdentifierVariantSun = "md5" + // ParameterRounds is the parameter name used by the Sun variant of this algorithm to carry the iteration count. + ParameterRounds = "rounds" + + // ParameterIterations is a non standard parameter name for the iteration count which earlier versions of this + // library emitted. It is accepted when decoding so those digests remain readable. + ParameterIterations = "iterations" + // VariantNameStandard is the md5crypt.Variant name for md5crypt.VariantStandard. VariantNameStandard = "standard" diff --git a/algorithm/md5crypt/decoder.go b/algorithm/md5crypt/decoder.go index 0f1b879..59ee72d 100644 --- a/algorithm/md5crypt/decoder.go +++ b/algorithm/md5crypt/decoder.go @@ -128,7 +128,7 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error for _, param := range params { switch param.Key { - case "rounds": + case ParameterRounds, ParameterIterations: var value uint64 if value, err = strconv.ParseUint(param.Value, 10, 32); err != nil { diff --git a/algorithm/md5crypt/regression_test.go b/algorithm/md5crypt/regression_test.go new file mode 100644 index 0000000..a9e63a8 --- /dev/null +++ b/algorithm/md5crypt/regression_test.go @@ -0,0 +1,52 @@ +package md5crypt + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSunVariantWithIterationsRoundTrips(t *testing.T) { + hasher, err := New(WithVariant(VariantSun), WithIterations(1000)) + 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.Equal(t, encoded, decoded.Encode()) + assert.True(t, decoded.Match("password")) + assert.False(t, decoded.Match("incorrect")) +} + +func TestSunVariantDecodesLegacyIterationsParameter(t *testing.T) { + hasher, err := New(WithVariant(VariantSun), WithIterations(1000)) + require.NoError(t, err) + + digest, err := hasher.Hash("password") + require.NoError(t, err) + + salt, key := string(digest.Salt()), string(digest.Key()) + + legacy := "$md5,iterations=1000$" + salt + "$$" + key + + decoded, err := Decode(legacy) + require.NoError(t, err) + + assert.True(t, decoded.Match("password")) +} + +func TestSunVariantEncodesRoundsParameter(t *testing.T) { + hasher, err := New(WithVariant(VariantSun), WithIterations(1000)) + require.NoError(t, err) + + digest, err := hasher.Hash("password") + require.NoError(t, err) + + assert.Contains(t, digest.Encode(), "$md5,rounds=1000$") +} diff --git a/algorithm/pbkdf2/decoder.go b/algorithm/pbkdf2/decoder.go index 7a21ecf..4df8bc9 100644 --- a/algorithm/pbkdf2/decoder.go +++ b/algorithm/pbkdf2/decoder.go @@ -141,6 +141,10 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return nil, fmt.Errorf("%w: iterations could not be parsed: %v", algorithm.ErrEncodedHashInvalidOptionValue, err) } + if decoded.iterations < 1 { + return nil, fmt.Errorf("%w: iterations must be at least 1 but is set to '%d'", algorithm.ErrEncodedHashInvalidOptionValue, decoded.iterations) + } + if decoded.salt, err = encoding.Base64RawAdaptedEncoding.DecodeString(parts[1]); err != nil { return nil, fmt.Errorf("%w: %v", algorithm.ErrEncodedHashSaltEncoding, err) } diff --git a/algorithm/pbkdf2/hasher.go b/algorithm/pbkdf2/hasher.go index e94c197..ecdce91 100644 --- a/algorithm/pbkdf2/hasher.go +++ b/algorithm/pbkdf2/hasher.go @@ -26,67 +26,27 @@ func New(opts ...Opt) (hasher *Hasher, err error) { // NewSHA1 returns a SHA1 variant *pbkdf2.Hasher with the additional opts applied if any. func NewSHA1(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantSHA1)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantSHA1)}, opts...)...) } // NewSHA224 returns a SHA224 variant *pbkdf2.Hasher with the additional opts applied if any. func NewSHA224(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantSHA224)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantSHA224)}, opts...)...) } // NewSHA256 returns a SHA256 variant *pbkdf2.Hasher with the additional opts applied if any. func NewSHA256(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantSHA256)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantSHA256)}, opts...)...) } // NewSHA384 returns a SHA384 variant *pbkdf2.Hasher with the additional opts applied if any. func NewSHA384(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantSHA384)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantSHA384)}, opts...)...) } // NewSHA512 returns a SHA512 variant *pbkdf2.Hasher with the additional opts applied if any. func NewSHA512(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantSHA512)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantSHA512)}, opts...)...) } // Hasher is a crypt.Hash for PBKDF2 which can be initialized via pbkdf2.New using a functional options pattern. diff --git a/algorithm/pbkdf2/regression_test.go b/algorithm/pbkdf2/regression_test.go new file mode 100644 index 0000000..6059aae --- /dev/null +++ b/algorithm/pbkdf2/regression_test.go @@ -0,0 +1,75 @@ +package pbkdf2 + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestVariantConstructorsDeriveKeyLengthFromTheirVariant(t *testing.T) { + testCases := []struct { + name string + new func(opts ...Opt) (*Hasher, error) + variant Variant + }{ + {"SHA1", NewSHA1, VariantSHA1}, + {"SHA224", NewSHA224, VariantSHA224}, + {"SHA256", NewSHA256, VariantSHA256}, + {"SHA384", NewSHA384, VariantSHA384}, + {"SHA512", NewSHA512, VariantSHA512}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + hasher, err := tc.new(WithIterations(IterationsMin)) + require.NoError(t, err) + + assert.NoError(t, hasher.Validate(), "a hasher returned by the constructor must validate") + + digest, err := hasher.Hash("password") + require.NoError(t, err) + + assert.Equal(t, tc.variant.HashFunc()().Size(), len(digest.Key())) + assert.True(t, digest.Match("password")) + }) + } +} + +func TestVariantConstructorsAcceptAKeyLengthOverride(t *testing.T) { + hasher, err := NewSHA512(WithIterations(IterationsMin), WithKeyLength(80)) + require.NoError(t, err) + + digest, err := hasher.Hash("password") + require.NoError(t, err) + + assert.Equal(t, 80, len(digest.Key())) +} + +func TestDecodeRejectsUnusableIterations(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"Zero", "$pbkdf2-sha256$0$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + {"Negative", "$pbkdf2-sha256$-5$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + digest, err := Decode(tc.digest) + + assert.Nil(t, digest) + assert.Error(t, err) + }) + } +} + +func TestDecodeAcceptsLegacyIterationCounts(t *testing.T) { + const encoded = "$pbkdf2-sha256$1000$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5" + + digest, err := Decode(encoded) + + require.NoError(t, err) + assert.Equal(t, encoded, digest.Encode()) +} diff --git a/algorithm/scrypt/decoder.go b/algorithm/scrypt/decoder.go index 2aff527..e7820a5 100644 --- a/algorithm/scrypt/decoder.go +++ b/algorithm/scrypt/decoder.go @@ -139,5 +139,9 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return nil, fmt.Errorf("%w: key has 0 bytes", algorithm.ErrEncodedHashKeyEncoding) } + if err = decoded.validate(); err != nil { + return nil, err + } + return decoded, nil } diff --git a/algorithm/scrypt/digest.go b/algorithm/scrypt/digest.go index 4f39dab..9bb46b3 100644 --- a/algorithm/scrypt/digest.go +++ b/algorithm/scrypt/digest.go @@ -72,11 +72,26 @@ func (d *Digest) Salt() (salt []byte) { return d.salt } -// n returns 2 to the power of log N i.e d.ln. func (d *Digest) n() (n int) { return 1 << d.ln } +func (d *Digest) validate() (err error) { + if d.ln < IterationsMin || d.ln > IterationsMax { + return fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, oLN, IterationsMin, "", IterationsMax, d.ln) + } + + if d.r < BlockSizeMin || d.r > BlockSizeMax { + return fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, oR, BlockSizeMin, "", BlockSizeMax, d.r) + } + + if d.p < ParallelismMin || d.p > ParallelismMax { + return fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, oP, ParallelismMin, "", ParallelismMax, d.p) + } + + return nil +} + func (d *Digest) defaults() { switch d.variant { case VariantScrypt, VariantYescrypt: diff --git a/algorithm/scrypt/fuzz_test.go b/algorithm/scrypt/fuzz_test.go new file mode 100644 index 0000000..fb4e1da --- /dev/null +++ b/algorithm/scrypt/fuzz_test.go @@ -0,0 +1,56 @@ +package scrypt + +import ( + "testing" +) + +func FuzzDecodeAndMatch(f *testing.F) { + seeds := []string{ + "$scrypt$ln=16,r=8,p=1$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "$scrypt$ln=1,r=1,p=1$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "$scrypt$ln=-1,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=0,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=58,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=59,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=16,r=-8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=16,r=8,p=-1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=16,r=0,p=0$c2FsdHNhbHQ$a2V5", + "$scrypt$$$", + "$scrypt$ln=x,r=8,p=1$c2FsdA$a2V5", + "$y$j9T$MnjNJEQnQ0trkgi3VmJJ.$FvSc2M9Xr4mDaIzsHpMxu3T5AQZoCfCz.3Xn8OU9pj5", + "$y$$$", + "", + "$", + } + + for _, seed := range seeds { + f.Add(seed) + } + + f.Fuzz(func(t *testing.T, encodedDigest string) { + decoded, err := Decode(encodedDigest) + + if err != nil { + return + } + + digest, ok := decoded.(*Digest) + if !ok { + t.Fatalf("Decode(%q) returned a %T rather than a *Digest", encodedDigest, decoded) + } + + if digest.ln < IterationsMin || digest.ln > IterationsMax { + t.Fatalf("Decode(%q) accepted an out of range ln of %d", encodedDigest, digest.ln) + } + + if digest.r < BlockSizeMin || digest.p < ParallelismMin { + t.Fatalf("Decode(%q) accepted an out of range r of %d or p of %d", encodedDigest, digest.r, digest.p) + } + + if digest.ln > 12 || digest.r > 16 || digest.p > 4 { + return + } + + _, _ = digest.MatchAdvanced("password") + }) +} diff --git a/algorithm/scrypt/hasher.go b/algorithm/scrypt/hasher.go index 67d6649..231cda7 100644 --- a/algorithm/scrypt/hasher.go +++ b/algorithm/scrypt/hasher.go @@ -26,27 +26,11 @@ func New(opts ...Opt) (hasher *Hasher, err error) { } func NewScrypt(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantScrypt)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantScrypt)}, opts...)...) } func NewYescrypt(opts ...Opt) (hasher *Hasher, err error) { - if hasher, err = New(opts...); err != nil { - return nil, err - } - - if err = hasher.WithOptions(WithVariant(VariantYescrypt)); err != nil { - return nil, err - } - - return hasher, nil + return New(append([]Opt{WithVariant(VariantYescrypt)}, opts...)...) } // Hasher is a crypt.Hash for scrypt which can be initialized via New using a functional options pattern. diff --git a/algorithm/scrypt/regression_test.go b/algorithm/scrypt/regression_test.go new file mode 100644 index 0000000..2485bfa --- /dev/null +++ b/algorithm/scrypt/regression_test.go @@ -0,0 +1,90 @@ +package scrypt + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDecodeRejectsParametersThatCannotBeUsed(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"NegativeLN", "$scrypt$ln=-1,r=8,p=1$c2FsdHNhbHQ$a2V5"}, + {"ZeroLN", "$scrypt$ln=0,r=8,p=1$c2FsdHNhbHQ$a2V5"}, + {"OversizedLN", "$scrypt$ln=59,r=8,p=1$c2FsdHNhbHQ$a2V5"}, + {"NegativeR", "$scrypt$ln=16,r=-8,p=1$c2FsdHNhbHQ$a2V5"}, + {"ZeroR", "$scrypt$ln=16,r=0,p=1$c2FsdHNhbHQ$a2V5"}, + {"NegativeP", "$scrypt$ln=16,r=8,p=-1$c2FsdHNhbHQ$a2V5"}, + {"ZeroP", "$scrypt$ln=16,r=8,p=0$c2FsdHNhbHQ$a2V5"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + digest, err := Decode(tc.digest) + + assert.Nil(t, digest) + require.Error(t, err) + }) + } +} + +func TestMatchingNeverPanicsForAcceptedDigests(t *testing.T) { + testCases := []string{ + "$scrypt$ln=1,r=1,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=10,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=12,r=16,p=2$c2FsdHNhbHQ$a2V5", + } + + for _, encoded := range testCases { + t.Run(encoded, func(t *testing.T) { + digest, err := Decode(encoded) + require.NoError(t, err) + + assert.NotPanics(t, func() { + _, _ = digest.MatchAdvanced("password") + }) + }) + } +} + +func TestDecodePreservesParameters(t *testing.T) { + const encoded = "$scrypt$ln=16,r=8,p=1$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU" + + digest, err := Decode(encoded) + + require.NoError(t, err) + assert.Equal(t, encoded, digest.Encode()) +} + +func TestHashedDigestsRoundTrip(t *testing.T) { + testCases := []struct { + name string + new func() (*Hasher, error) + }{ + {"Scrypt", func() (*Hasher, error) { return NewScrypt(WithLN(10)) }}, + {"Yescrypt", func() (*Hasher, error) { return NewYescrypt(WithLN(10)) }}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + hasher, err := tc.new() + require.NoError(t, err) + require.NoError(t, hasher.Validate()) + + 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.Equal(t, encoded, decoded.Encode()) + assert.True(t, decoded.Match("password")) + assert.False(t, decoded.Match("incorrect")) + }) + } +} diff --git a/algorithm/shacrypt/decoder.go b/algorithm/shacrypt/decoder.go index a0fb81a..5eb3663 100644 --- a/algorithm/shacrypt/decoder.go +++ b/algorithm/shacrypt/decoder.go @@ -124,6 +124,10 @@ func decode(variant Variant, parts []string) (digest algorithm.Digest, err error return nil, fmt.Errorf("%w: option '%s' has invalid value '%s': %v", algorithm.ErrEncodedHashInvalidOptionValue, param.Key, param.Value, err) } + if rounds == 0 || rounds > IterationsMax { + return nil, fmt.Errorf(algorithm.ErrFmtInvalidIntParameter, algorithm.ErrEncodedHashInvalidOptionValue, param.Key, 1, "", IterationsMax, rounds) + } + decoded.iterations = int(rounds) default: return nil, fmt.Errorf("%w: option '%s' with value '%s' is unknown", algorithm.ErrEncodedHashInvalidOptionKey, param.Key, param.Value) diff --git a/algorithm/shacrypt/regression_test.go b/algorithm/shacrypt/regression_test.go new file mode 100644 index 0000000..ad144f0 --- /dev/null +++ b/algorithm/shacrypt/regression_test.go @@ -0,0 +1,68 @@ +package shacrypt + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDecodeRejectsUnusableRounds(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"Zero", "$6$rounds=0$saltsalt$keykeykey"}, + {"AboveMaximum", "$6$rounds=4294967295$saltsalt$keykeykey"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + digest, err := Decode(tc.digest) + + assert.Nil(t, digest) + assert.Error(t, err) + }) + } +} + +func TestDecodeAcceptsRoundsWithinRange(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"Minimum", "$6$rounds=1000$saltsalt$keykeykey"}, + {"Maximum", "$6$rounds=999999999$saltsalt$keykeykey"}, + {"BelowSpecMinimumButUsable", "$6$rounds=100$saltsalt$keykeykey"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + digest, err := Decode(tc.digest) + + require.NoError(t, err) + assert.Equal(t, tc.digest, digest.Encode()) + }) + } +} + +func TestHashedDigestsRoundTrip(t *testing.T) { + for _, variant := range []Variant{VariantSHA256, VariantSHA512} { + t.Run(variant.String(), func(t *testing.T) { + hasher, err := New(WithVariant(variant), WithIterations(1000)) + require.NoError(t, err) + + digest, err := hasher.Hash("password") + require.NoError(t, err) + + encoded := digest.Encode() + + decoded, err := Decode(encoded) + require.NoError(t, err) + + assert.Equal(t, encoded, decoded.Encode()) + assert.True(t, decoded.Match("password")) + assert.False(t, decoded.Match("incorrect")) + }) + } +} diff --git a/const.go b/const.go index f21ff4c..48a7cf1 100644 --- a/const.go +++ b/const.go @@ -2,6 +2,7 @@ package crypt import ( "github.com/go-crypt/crypt/internal/encoding" + "errors" ) const ( @@ -16,3 +17,8 @@ const ( // StorageFormatPrefixLDAPArgon2 is a prefix used by OpenLDAP for argon2 format encoded digests. StorageFormatPrefixLDAPArgon2 = "{ARGON2}" ) + +var ( + // ErrDigestNil is returned by the crypt.Digest matcher methods when it does not wrap an algorithm.Digest. + ErrDigestNil = errors.New("crypt.Digest does not wrap an algorithm.Digest") +) diff --git a/crypt_test.go b/crypt_test.go index ed15caf..fe4a777 100644 --- a/crypt_test.go +++ b/crypt_test.go @@ -693,7 +693,7 @@ func TestNullDigestScan(t *testing.T) { { "ShouldFailInvalidType", 123, - "invalid type for crypt.Digest: can't scan int into crypt.Digest", + "invalid type for crypt.NullDigest: can't scan int into crypt.NullDigest", }, } diff --git a/decode.go b/decode.go index a2d4d94..8db948d 100644 --- a/decode.go +++ b/decode.go @@ -1,11 +1,18 @@ package crypt import ( + "sync" + "github.com/go-crypt/crypt/algorithm" ) -// The global Decoder. This is utilized by the Decode function. -var gdecoder *Decoder +// The global Decoder. This is utilized by the Decode function. It is initialized exactly once by gdecoderOnce so that +// the Decode function is safe for concurrent use. +var ( + gdecoder *Decoder + gdecoderErr error + gdecoderOnce sync.Once +) // Decode is a convenience function which wraps the Decoder functionality. It's recommended to create your own decoder // instead via NewDecoder or NewDefaultDecoder. @@ -22,10 +29,14 @@ func Decode(encodedDigest string) (digest algorithm.Digest, err error) { } func decode(encodedDigest string) (digest algorithm.Digest, err error) { - if gdecoder == nil { - if gdecoder, err = NewDefaultDecoder(); err != nil { - return nil, err + gdecoderOnce.Do(func() { + if gdecoder, gdecoderErr = NewDefaultDecoder(); gdecoderErr == nil { + gdecoder.global = true } + }) + + if gdecoderErr != nil { + return nil, gdecoderErr } return gdecoder.Decode(encodedDigest) diff --git a/decoder.go b/decoder.go index 25923bd..57ad7ef 100644 --- a/decoder.go +++ b/decoder.go @@ -88,6 +88,7 @@ func NewDecoderAll() (d *Decoder, err error) { type Decoder struct { decoders map[string]algorithm.DecodeFunc prefixes map[string]string + global bool } // RegisterDecodeFunc registers a new algorithm.DecodeFunc with this Decoder against a specific identifier. @@ -134,10 +135,8 @@ func (d *Decoder) Decode(encodedDigest string) (digest algorithm.Digest, err err } func (d *Decoder) decode(encodedDigest string) (digest algorithm.Digest, err error) { - for prefix, key := range d.prefixes { - if strings.HasPrefix(encodedDigest, prefix) { - return d.decoders[key](encodedDigest) - } + if key, ok := d.matchPrefix(encodedDigest); ok { + return d.decoders[key](encodedDigest) } encodedDigest = Normalize(encodedDigest) @@ -156,12 +155,25 @@ func (d *Decoder) decode(encodedDigest string) (digest algorithm.Digest, err err return decodeFunc(encodedDigest) } - switch d { - case gdecoder: + if d.global { return nil, fmt.Errorf("%w: the identifier '%s' is unknown to the global decoder", algorithm.ErrEncodedHashInvalidIdentifier, parts[1]) - default: - return nil, fmt.Errorf("%w: the identifier '%s' is unknown to the decoder", algorithm.ErrEncodedHashInvalidIdentifier, parts[1]) } + + return nil, fmt.Errorf("%w: the identifier '%s' is unknown to the decoder", algorithm.ErrEncodedHashInvalidIdentifier, parts[1]) +} + +func (d *Decoder) matchPrefix(encodedDigest string) (identifier string, ok bool) { + var matched string + + for prefix, key := range d.prefixes { + if !strings.HasPrefix(encodedDigest, prefix) || len(prefix) <= len(matched) && ok { + continue + } + + matched, identifier, ok = prefix, key, true + } + + return identifier, ok } func decoderProfileDefault(decoder *Decoder) (err error) { diff --git a/fuzz_test.go b/fuzz_test.go new file mode 100644 index 0000000..d0a940e --- /dev/null +++ b/fuzz_test.go @@ -0,0 +1,145 @@ +package crypt + +import ( + "testing" +) + +func FuzzDecode(f *testing.F) { + for _, seed := range corpusDecode { + f.Add(seed) + } + + f.Fuzz(func(t *testing.T, encodedDigest string) { + digest, err := Decode(encodedDigest) + + if err != nil { + if digest != nil { + t.Fatalf("Decode(%q) returned both a digest and the error %v", encodedDigest, err) + } + + return + } + + if digest == nil { + t.Fatalf("Decode(%q) returned no digest and no error", encodedDigest) + } + + _, _ = digest.Encode(), digest.String() + _, _ = digest.Key(), digest.Salt() + }) +} + +func FuzzNormalize(f *testing.F) { + for _, seed := range corpusDecode { + f.Add(seed) + } + + f.Fuzz(func(t *testing.T, encodedDigest string) { + normalized := Normalize(encodedDigest) + + if len(normalized) > len(encodedDigest) { + t.Fatalf("Normalize(%q) returned the longer value %q", encodedDigest, normalized) + } + }) +} + +func TestCheckPasswordNeverPanicsOverCorpus(t *testing.T) { + for _, encodedDigest := range corpusDecode { + t.Run(encodedDigest, func(t *testing.T) { + var ( + valid bool + err error + ) + + if !assertNotPanics(t, func() { valid, err = CheckPassword("password", encodedDigest) }) { + return + } + + if valid && err != nil { + t.Fatalf("CheckPassword reported a match alongside the error %v", err) + } + }) + } +} + +func assertNotPanics(t *testing.T, f func()) (ok bool) { + t.Helper() + + defer func() { + if r := recover(); r != nil { + t.Errorf("panic: %v", r) + + ok = false + } + }() + + f() + + return true +} + +var corpusDecode = []string{ + // Well formed digests for each supported identifier. + "$argon2id$v=19$m=2097152,t=1,p=4$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "$argon2i$v=19$m=2097152,t=1,p=4$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "$argon2d$v=19$m=2097152,t=1,p=4$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "$2b$12$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm", + "$2a$12$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm", + "$2y$12$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm", + "$bcrypt-sha256$v=2,t=2b,r=12$3XCpXfcQBjcbXFHTLcbFju$AXNZ1B7NPTf7XyCqUKcvIUOB5eKKZ4C", + "$pbkdf2-sha256$100000$YmxhaGJsYWhibGFoYmxhaA$Rlt4rHrEcdEJA9FBAOtqE21NX2NDaCyR3xr0PJmg.dU", + "$scrypt$ln=16,r=8,p=1$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "$5$rounds=1000$saltsalt$keykeykeykeykey", + "$6$rounds=1000$saltsalt$keykeykeykeykey", + "$6$saltsalt$keykeykeykeykey", + "$1$saltsalt$keykeykeykeykey", + "$md5$saltsalt$$keykeykeykeykey", + "$md5,rounds=1000$saltsalt$$keykeykeykeykey", + "$sha1$480000$saltsalt$keykeykeykeykey", + "$plaintext$password", + "$base64$cGFzc3dvcmQ", + + // LDAP style prefixes handled by Normalize. + "{CRYPT}$6$rounds=1000$saltsalt$keykeykeykeykey", + "{ARGON2}$argon2id$v=19$m=2097152,t=1,p=4$YmxhaGJsYWhibGFoYmxhaA$Vt4rHrEcdEJ+A9FBAOtqE21NX2NDaCyR3xr0PJmg+dU", + "{PBKDF2-SHA256}100000$YmxhaGJsYWhibGFoYmxhaA$Rlt4rHrEcdEJA9FBAOtqE21NX2NDaCyR3xr0PJmg.dU", + + // Parameters which are outside the range each algorithm can actually use. + "$scrypt$ln=-1,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=64,r=8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=16,r=-8,p=1$c2FsdHNhbHQ$a2V5", + "$scrypt$ln=16,r=8,p=-1$c2FsdHNhbHQ$a2V5", + "$2b$99$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm", + "$2b$-1$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm", + "$2b$00$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm", + "$pbkdf2-sha256$0$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$pbkdf2-sha256$-5$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$argon2id$v=19$m=8,t=0,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$argon2id$v=19$m=8,t=1,p=0$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$argon2id$v=19$m=0,t=1,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$argon2id$v=16$m=8,t=1,p=1$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5", + "$6$rounds=0$saltsalt$keykeykey", + "$6$rounds=4294967295$saltsalt$keykeykey", + + // Structurally malformed input. + "", + "$", + "$$", + "$$$", + "$$$$$$$$$$", + "notadigest", + "$unknown$identifier$value", + "$argon2id$", + "$argon2id$v=19$m=2097152,t=1,p=4$$", + "$2b$", + "$2b$12$", + "$scrypt$$$", + "$plaintext$", + "$base64$!!!!not-base64!!!!", + "$argon2id$v=19$m=2097152,t=1,p=4$!!!$!!!", + "$6$rounds=notanumber$saltsalt$key", + "$scrypt$ln=notanumber,r=8,p=1$c2FsdA$a2V5", + "$argon2id$v=19$m=2097152,t=1,p=4,unknown=1$c2FsdA$a2V5", + "\x00\x00\x00", + "$\x00$\x00$\x00", +} diff --git a/regression_test.go b/regression_test.go new file mode 100644 index 0000000..9b35f73 --- /dev/null +++ b/regression_test.go @@ -0,0 +1,149 @@ +package crypt + +import ( + "sync" + "testing" + + "github.com/go-crypt/crypt/algorithm" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDecodeIsSafeForConcurrentUse(t *testing.T) { + const digest = "$2b$12$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm" + + var wg sync.WaitGroup + + errs := make([]error, 64) + + for i := range errs { + wg.Add(1) + + go func() { + defer wg.Done() + + _, errs[i] = Decode(digest) + }() + } + + wg.Wait() + + for i, err := range errs { + assert.NoError(t, err, "goroutine %d", i) + } +} + +func TestCheckPasswordDoesNotPanicOnMalformedDigests(t *testing.T) { + testCases := []struct { + name string + digest string + }{ + {"ScryptNegativeLN", "$scrypt$ln=-1,r=8,p=1$c2FsdHNhbHQ$a2V5"}, + {"ScryptNegativeR", "$scrypt$ln=16,r=-8,p=1$c2FsdHNhbHQ$a2V5"}, + {"ScryptNegativeP", "$scrypt$ln=16,r=8,p=-1$c2FsdHNhbHQ$a2V5"}, + {"ScryptOversizedLN", "$scrypt$ln=64,r=8,p=1$c2FsdHNhbHQ$a2V5"}, + {"BcryptOversizedCost", "$2b$99$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm"}, + {"BcryptNegativeCost", "$2b$-1$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm"}, + {"PBKDF2NegativeIterations", "$pbkdf2-sha256$-5$c2FsdHNhbHQ$c2FsdHNhbHRzYWx0c2FsdHNhbHRzYWx0c2FsdHNhbA"}, + {"Argon2ZeroParallelism", "$argon2id$v=19$m=8,t=1,p=0$c2FsdHNhbHQ$a2V5a2V5a2V5a2V5"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + assert.NotPanics(t, func() { + valid, err := CheckPassword("password", tc.digest) + + assert.False(t, valid) + assert.Error(t, err) + }) + }) + } +} + +func TestDigestScanAcceptsByteSlices(t *testing.T) { + const encoded = "$2b$12$3XCpXfcQBjcbXFHTLcbFju0KNQ2ipfeNbcH8b7ZgIkXlbNkYbGDWm" + + t.Run("Digest", func(t *testing.T) { + var digest Digest + + require.NoError(t, digest.Scan([]byte(encoded))) + assert.Equal(t, encoded, digest.Encode()) + }) + + t.Run("NullDigest", func(t *testing.T) { + var digest NullDigest + + require.NoError(t, digest.Scan([]byte(encoded))) + assert.Equal(t, encoded, digest.Encode()) + }) + + t.Run("NullDigestEmptyBytes", func(t *testing.T) { + var digest NullDigest + + require.NoError(t, digest.Scan([]byte(nil))) + assert.Equal(t, "", digest.Encode()) + }) +} + +func TestDigestZeroValueDoesNotPanic(t *testing.T) { + var digest Digest + + assert.NotPanics(t, func() { + assert.Equal(t, "", digest.Encode()) + assert.Equal(t, "", digest.String()) + assert.Nil(t, digest.Key()) + assert.Nil(t, digest.Salt()) + assert.False(t, digest.Match("password")) + assert.False(t, digest.MatchBytes([]byte("password"))) + }) + + assert.NotPanics(t, func() { + match, err := digest.MatchAdvanced("password") + + assert.False(t, match) + assert.Error(t, err) + }) + + assert.NotPanics(t, func() { + match, err := digest.MatchBytesAdvanced([]byte("password")) + + assert.False(t, match) + assert.Error(t, err) + }) +} + +func TestDecoderPrefixMatchingIsDeterministic(t *testing.T) { + decoder := NewDecoder() + + require.NoError(t, decoder.RegisterDecodeFunc("short", newStubDecodeFunc("short"))) + require.NoError(t, decoder.RegisterDecodeFunc("long", newStubDecodeFunc("long"))) + + require.NoError(t, decoder.RegisterDecodePrefix("{X}", "short")) + require.NoError(t, decoder.RegisterDecodePrefix("{X}{Y}", "long")) + + for i := 0; i < 32; i++ { + digest, err := decoder.Decode("{X}{Y}value") + + require.NoError(t, err) + assert.Equal(t, "long", digest.Encode()) + } +} + +type stubDigest struct { + name string +} + +func (d *stubDigest) Encode() string { return d.name } +func (d *stubDigest) String() string { return d.name } +func (d *stubDigest) Key() []byte { return nil } +func (d *stubDigest) Salt() []byte { return nil } +func (d *stubDigest) Match(string) bool { return false } +func (d *stubDigest) MatchBytes([]byte) bool { return false } +func (d *stubDigest) MatchAdvanced(string) (bool, error) { return false, nil } +func (d *stubDigest) MatchBytesAdvanced([]byte) (bool, error) { return false, nil } + +func newStubDecodeFunc(name string) algorithm.DecodeFunc { + return func(encodedDigest string) (algorithm.Digest, error) { + return &stubDigest{name: name}, nil + } +} diff --git a/roundtrip_test.go b/roundtrip_test.go new file mode 100644 index 0000000..6fb15f1 --- /dev/null +++ b/roundtrip_test.go @@ -0,0 +1,236 @@ +package crypt + +import ( + "encoding" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/go-crypt/crypt/algorithm" + "github.com/go-crypt/crypt/algorithm/argon2" + "github.com/go-crypt/crypt/algorithm/bcrypt" + "github.com/go-crypt/crypt/algorithm/md5crypt" + "github.com/go-crypt/crypt/algorithm/pbkdf2" + "github.com/go-crypt/crypt/algorithm/plaintext" + "github.com/go-crypt/crypt/algorithm/scrypt" + "github.com/go-crypt/crypt/algorithm/sha1crypt" + "github.com/go-crypt/crypt/algorithm/shacrypt" +) + +// newHasherFuncs returns a hasher per supported algorithm and variant, configured with the cheapest parameters each +// one accepts so the round trip assertions stay fast. +func newHasherFuncs() map[string]func() (algorithm.Hash, error) { + return map[string]func() (algorithm.Hash, error){ + "Argon2id": func() (algorithm.Hash, error) { + return argon2.New(argon2.WithVariantID(), argon2.WithProfileRFC9106LowMemory()) + }, + "Argon2i": func() (algorithm.Hash, error) { + return argon2.New(argon2.WithVariantI(), argon2.WithProfileRFC9106LowMemory()) + }, + "Argon2d": func() (algorithm.Hash, error) { + return argon2.New(argon2.WithVariantD(), argon2.WithProfileRFC9106LowMemory()) + }, + "Bcrypt": func() (algorithm.Hash, error) { + return bcrypt.New(bcrypt.WithCost(bcrypt.IterationsMin)) + }, + "BcryptSHA256": func() (algorithm.Hash, error) { + return bcrypt.NewSHA256(bcrypt.WithCost(bcrypt.IterationsMin)) + }, + "PBKDF2SHA1": func() (algorithm.Hash, error) { + return pbkdf2.NewSHA1(pbkdf2.WithIterations(pbkdf2.IterationsMin)) + }, + "PBKDF2SHA224": func() (algorithm.Hash, error) { + return pbkdf2.NewSHA224(pbkdf2.WithIterations(pbkdf2.IterationsMin)) + }, + "PBKDF2SHA256": func() (algorithm.Hash, error) { + return pbkdf2.NewSHA256(pbkdf2.WithIterations(pbkdf2.IterationsMin)) + }, + "PBKDF2SHA384": func() (algorithm.Hash, error) { + return pbkdf2.NewSHA384(pbkdf2.WithIterations(pbkdf2.IterationsMin)) + }, + "PBKDF2SHA512": func() (algorithm.Hash, error) { + return pbkdf2.NewSHA512(pbkdf2.WithIterations(pbkdf2.IterationsMin)) + }, + "Scrypt": func() (algorithm.Hash, error) { + return scrypt.NewScrypt(scrypt.WithLN(10)) + }, + "Yescrypt": func() (algorithm.Hash, error) { + return scrypt.NewYescrypt(scrypt.WithLN(10)) + }, + "SHACryptSHA256": func() (algorithm.Hash, error) { + return shacrypt.New(shacrypt.WithSHA256(), shacrypt.WithIterations(shacrypt.IterationsMin)) + }, + "SHACryptSHA512": func() (algorithm.Hash, error) { + return shacrypt.New(shacrypt.WithSHA512(), shacrypt.WithIterations(shacrypt.IterationsMin)) + }, + "MD5CryptStandard": func() (algorithm.Hash, error) { + return md5crypt.New(md5crypt.WithVariant(md5crypt.VariantStandard)) + }, + "MD5CryptSun": func() (algorithm.Hash, error) { + return md5crypt.New(md5crypt.WithVariant(md5crypt.VariantSun), md5crypt.WithIterations(1000)) + }, + "SHA1Crypt": func() (algorithm.Hash, error) { + return sha1crypt.New(sha1crypt.WithIterations(1000)) + }, + "PlainText": func() (algorithm.Hash, error) { + return plaintext.New(plaintext.WithVariant(plaintext.VariantPlainText)) + }, + "Base64": func() (algorithm.Hash, error) { + return plaintext.New(plaintext.WithVariant(plaintext.VariantBase64)) + }, + } +} + +// TestHashEncodeDecodeRoundTrip asserts every hasher produces an encoded digest which the decoder accepts, which +// re-encodes to the identical string, and which still matches the password it was created from. A digest this library +// writes but cannot read back is unusable, so this closes the loop each algorithm depends on. +func TestHashEncodeDecodeRoundTrip(t *testing.T) { + decoder, err := NewDecoderAll() + require.NoError(t, err) + + for name, newHasher := range newHasherFuncs() { + t.Run(name, func(t *testing.T) { + hasher, err := newHasher() + require.NoError(t, err) + require.NoError(t, hasher.Validate()) + + digest, err := hasher.Hash(password) + require.NoError(t, err) + + encoded := digest.Encode() + + t.Run("DecodesToTheSameEncoding", func(t *testing.T) { + decoded, err := decoder.Decode(encoded) + require.NoError(t, err, "encoded digest %q could not be decoded", encoded) + + assert.Equal(t, encoded, decoded.Encode()) + }) + + t.Run("MatchesTheOriginalPassword", func(t *testing.T) { + decoded, err := decoder.Decode(encoded) + require.NoError(t, err) + + match, err := decoded.MatchAdvanced(password) + require.NoError(t, err) + assert.True(t, match) + }) + + t.Run("RejectsTheWrongPassword", func(t *testing.T) { + decoded, err := decoder.Decode(encoded) + require.NoError(t, err) + + match, err := decoded.MatchAdvanced(wrongPassword) + require.NoError(t, err) + assert.False(t, match) + }) + + t.Run("DecodingIsIdempotent", func(t *testing.T) { + first, err := decoder.Decode(encoded) + require.NoError(t, err) + + second, err := decoder.Decode(first.Encode()) + require.NoError(t, err) + + assert.Equal(t, first.Encode(), second.Encode()) + assert.Equal(t, first.Key(), second.Key()) + assert.Equal(t, first.Salt(), second.Salt()) + }) + }) + } +} + +// TestHashWithSaltIsDeterministic asserts hashing the same password with the same salt twice produces the same digest. +func TestHashWithSaltIsDeterministic(t *testing.T) { + // Salts are constrained differently per algorithm, so the salt of a first hash is reused rather than invented. + for name, newHasher := range newHasherFuncs() { + t.Run(name, func(t *testing.T) { + hasher, err := newHasher() + require.NoError(t, err) + + first, err := hasher.Hash(password) + require.NoError(t, err) + + salt := first.Salt() + + if len(salt) == 0 { + t.Skip("algorithm does not use a salt") + } + + second, err := hasher.HashWithSalt(password, salt) + require.NoError(t, err) + + third, err := hasher.HashWithSalt(password, salt) + require.NoError(t, err) + + assert.Equal(t, second.Encode(), third.Encode()) + assert.True(t, second.Match(password)) + }) + } +} + +// TestDigestWrapperRoundTrip asserts the crypt.Digest and crypt.NullDigest decorators preserve an encoded digest +// across every serialisation interface they implement. +func TestDigestWrapperRoundTrip(t *testing.T) { + hasher, err := argon2.New(argon2.WithProfileRFC9106LowMemory()) + require.NoError(t, err) + + algDigest, err := hasher.Hash(password) + require.NoError(t, err) + + encoded := algDigest.Encode() + + t.Run("Digest", func(t *testing.T) { + digest, err := NewDigest(algDigest) + require.NoError(t, err) + + assertSerialisationRoundTrip(t, digest, &Digest{}, encoded) + + value, err := digest.Value() + require.NoError(t, err) + assert.Equal(t, encoded, value) + }) + + t.Run("NullDigest", func(t *testing.T) { + digest := NewNullDigest(algDigest) + + assertSerialisationRoundTrip(t, digest, &NullDigest{}, encoded) + + value, err := digest.Value() + require.NoError(t, err) + assert.Equal(t, encoded, value) + }) +} + +type digestSerialiser interface { + encoding.TextMarshaler + encoding.BinaryMarshaler + Encode() string +} + +type digestDeserialiser interface { + encoding.TextUnmarshaler + encoding.BinaryUnmarshaler + Encode() string +} + +func assertSerialisationRoundTrip(t *testing.T, src digestSerialiser, dst digestDeserialiser, encoded string) { + t.Helper() + + t.Run("Text", func(t *testing.T) { + data, err := src.MarshalText() + require.NoError(t, err) + assert.Equal(t, encoded, string(data)) + + require.NoError(t, dst.UnmarshalText(data)) + assert.Equal(t, encoded, dst.Encode()) + }) + + t.Run("Binary", func(t *testing.T) { + data, err := src.MarshalBinary() + require.NoError(t, err) + + require.NoError(t, dst.UnmarshalBinary(data)) + assert.Equal(t, encoded, dst.Encode()) + }) +} diff --git a/types.go b/types.go index 7fb929d..53ff51b 100644 --- a/types.go +++ b/types.go @@ -62,41 +62,73 @@ type Digest struct { // Encode decorates the algorithm.Digest Encode function. func (d *Digest) Encode() string { + if d.digest == nil { + return "" + } + return d.digest.Encode() } // String decorates the algorithm.Digest String function. func (d *Digest) String() string { + if d.digest == nil { + return "" + } + return d.digest.String() } // MatchBytes decorates the algorithm.Digest MatchBytes function. func (d *Digest) MatchBytes(passwordBytes []byte) (match bool) { + if d.digest == nil { + return false + } + return d.digest.MatchBytes(passwordBytes) } // MatchAdvanced decorates the algorithm.Digest MatchAdvanced function. func (d *Digest) MatchAdvanced(password string) (match bool, err error) { + if d.digest == nil { + return false, ErrDigestNil + } + return d.digest.MatchAdvanced(password) } // MatchBytesAdvanced decorates the algorithm.Digest MatchBytesAdvanced function. func (d *Digest) MatchBytesAdvanced(passwordBytes []byte) (match bool, err error) { + if d.digest == nil { + return false, ErrDigestNil + } + return d.digest.MatchBytesAdvanced(passwordBytes) } // Match decorates the algorithm.Digest Match function. func (d *Digest) Match(password string) (match bool) { + if d.digest == nil { + return false + } + return d.digest.Match(password) } // Key returns the key which is the final result of this digest. func (d *Digest) Key() (key []byte) { + if d.digest == nil { + return nil + } + return d.digest.Key() } // Salt returns the salt used to generate this digest. func (d *Digest) Salt() (salt []byte) { + if d.digest == nil { + return nil + } + return d.digest.Salt() } @@ -120,7 +152,7 @@ func (d *Digest) Scan(src any) (err error) { } return nil - case byte: + case []byte: if d.digest, err = Decode(string(digest)); err != nil { return err } @@ -261,19 +293,31 @@ func (d *NullDigest) Scan(src any) (err error) { return nil case string: + if len(digest) == 0 { + d.digest = nil + + return nil + } + if d.digest, err = Decode(digest); err != nil { return err } return nil - case byte: + case []byte: + if len(digest) == 0 { + d.digest = nil + + return nil + } + if d.digest, err = Decode(string(digest)); err != nil { return err } return nil default: - return fmt.Errorf("invalid type for crypt.Digest: can't scan %T into crypt.Digest", digest) + return fmt.Errorf("invalid type for crypt.NullDigest: can't scan %T into crypt.NullDigest", digest) } }