diff --git a/utils/chainhashutil/strict.go b/utils/chainhashutil/strict.go new file mode 100644 index 00000000..20260179 --- /dev/null +++ b/utils/chainhashutil/strict.go @@ -0,0 +1,26 @@ +package chainhashutil + +import ( + "fmt" + + "github.com/btcsuite/btcd/chaincfg/chainhash" +) + +// NewHashFromStrExact parses a chainhash string that must be fully specified. +func NewHashFromStrExact(hash string) (chainhash.Hash, error) { + if len(hash) != chainhash.MaxHashStringSize { + return chainhash.Hash{}, fmt.Errorf( + "invalid hash string length of %v, want %v", + len(hash), chainhash.MaxHashStringSize) + } + + parsed, err := chainhash.NewHashFromStr(hash) + if err != nil { + return chainhash.Hash{}, err + } + + // chainhash.NewHashFromStr uses a pointer return, but on success it + // returns a populated hash, not (nil, nil), so dereferencing here is + // safe. + return *parsed, nil +} diff --git a/utils/chainhashutil/strict_test.go b/utils/chainhashutil/strict_test.go new file mode 100644 index 00000000..91696402 --- /dev/null +++ b/utils/chainhashutil/strict_test.go @@ -0,0 +1,64 @@ +package chainhashutil + +import ( + "strings" + "testing" + + "github.com/btcsuite/btcd/chaincfg/chainhash" + "github.com/stretchr/testify/require" +) + +// TestNewHashFromStrExact verifies that strict hash parsing rejects any +// non-fully-specified chainhash string. +func TestNewHashFromStrExact(t *testing.T) { + t.Parallel() + + validHash := strings.Repeat("01", 32) + + testCases := []struct { + name string + hash string + wantErr string + }{ + { + name: "valid", + hash: validHash, + wantErr: "", + }, + { + name: "short", + hash: validHash[:62], + wantErr: "invalid hash string length", + }, + { + name: "odd length", + hash: validHash[:63], + wantErr: "invalid hash string length", + }, + { + name: "empty", + hash: "", + wantErr: "invalid hash string length", + }, + { + name: "non hex", + hash: strings.Repeat("z", 64), + wantErr: "invalid byte", + }, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + hash, err := NewHashFromStrExact(testCase.hash) + if testCase.wantErr != "" { + require.ErrorContains(t, err, testCase.wantErr) + require.Equal(t, chainhash.Hash{}, hash) + } else { + require.NoError(t, err) + require.Equal(t, testCase.hash, hash.String()) + } + }) + } +}