diff --git a/consensus/bor/snapshot_test.go b/consensus/bor/snapshot_test.go index de576c18d6..ee15add259 100644 --- a/consensus/bor/snapshot_test.go +++ b/consensus/bor/snapshot_test.go @@ -1,10 +1,10 @@ package bor import ( - "math/rand" + "crypto/rand" + "math/big" "sort" "testing" - "time" "github.com/stretchr/testify/require" "pgregory.net/rapid" @@ -28,8 +28,8 @@ func TestGetSignerSuccessionNumber_ProposerIsSigner(t *testing.T) { } // proposer is signer - signer := validatorSet.Proposer.Address - successionNumber, err := snap.GetSignerSuccessionNumber(signer) + signerTest := validatorSet.Proposer.Address + successionNumber, err := snap.GetSignerSuccessionNumber(signerTest) if err != nil { t.Fatalf("%s", err) } @@ -54,8 +54,8 @@ func TestGetSignerSuccessionNumber_SignerIndexIsLarger(t *testing.T) { } // choose a signer at an index greater than proposer index - signer := snap.ValidatorSet.Validators[signerIndex].Address - successionNumber, err := snap.GetSignerSuccessionNumber(signer) + signerTest := snap.ValidatorSet.Validators[signerIndex].Address + successionNumber, err := snap.GetSignerSuccessionNumber(signerTest) if err != nil { t.Fatalf("%s", err) } @@ -76,8 +76,8 @@ func TestGetSignerSuccessionNumber_SignerIndexIsSmaller(t *testing.T) { } // choose a signer at an index greater than proposer index - signer := snap.ValidatorSet.Validators[signerIndex].Address - successionNumber, err := snap.GetSignerSuccessionNumber(signer) + signerTest := snap.ValidatorSet.Validators[signerIndex].Address + successionNumber, err := snap.GetSignerSuccessionNumber(signerTest) if err != nil { t.Fatalf("%s", err) } @@ -95,13 +95,13 @@ func TestGetSignerSuccessionNumber_ProposerNotFound(t *testing.T) { require.Len(t, snap.ValidatorSet.Validators, numVals) - dummyProposerAddress := randomAddress() + dummyProposerAddress := randomAddress(toAddresses(validators)...) snap.ValidatorSet.Proposer = &valset.Validator{Address: dummyProposerAddress} // choose any signer - signer := snap.ValidatorSet.Validators[3].Address + signerTest := snap.ValidatorSet.Validators[3].Address - _, err := snap.GetSignerSuccessionNumber(signer) + _, err := snap.GetSignerSuccessionNumber(signerTest) require.NotNil(t, err) e, ok := err.(*UnauthorizedProposerError) @@ -116,26 +116,30 @@ func TestGetSignerSuccessionNumber_SignerNotFound(t *testing.T) { snap := Snapshot{ ValidatorSet: valset.NewValidatorSet(validators), } - dummySignerAddress := randomAddress() + + dummySignerAddress := randomAddress(toAddresses(validators)...) _, err := snap.GetSignerSuccessionNumber(dummySignerAddress) require.NotNil(t, err) + e, ok := err.(*UnauthorizedSignerError) require.True(t, ok) + require.Equal(t, dummySignerAddress.Bytes(), e.Signer) } // nolint: unparam func buildRandomValidatorSet(numVals int) []*valset.Validator { - rand.Seed(time.Now().Unix()) - validators := make([]*valset.Validator, numVals) valAddrs := randomAddresses(numVals) for i := 0; i < numVals; i++ { + power, _ := rand.Int(nil, big.NewInt(99)) + powerN := power.Int64() + 1 + validators[i] = &valset.Validator{ Address: valAddrs[i], // cannot process validators with voting power 0, hence +1 - VotingPower: int64(rand.Intn(99) + 1), + VotingPower: powerN, } } @@ -145,11 +149,27 @@ func buildRandomValidatorSet(numVals int) []*valset.Validator { return validators } -func randomAddress() common.Address { - bytes := make([]byte, 32) - rand.Read(bytes) +func randomAddress(exclude ...common.Address) common.Address { + excl := make(map[common.Address]struct{}, len(exclude)) - return common.BytesToAddress(bytes) + for _, addr := range exclude { + excl[addr] = struct{}{} + } + + bytes := make([]byte, 32) + + var addr common.Address + + for { + _, _ = rand.Read(bytes) + addr = common.BytesToAddress(bytes) + + if _, ok := excl[addr]; ok { + continue + } + + return addr + } } func randomAddresses(n int) []common.Address { @@ -168,7 +188,7 @@ func randomAddresses(n int) []common.Address { bytes := make([]byte, 32) for { - rand.Read(bytes) + _, _ = rand.Read(bytes) addr = common.BytesToAddress(bytes) @@ -189,7 +209,7 @@ func TestRandomAddresses(t *testing.T) { t.Parallel() rapid.Check(t, func(t *rapid.T) { - length := rapid.IntMax(100).Draw(t, "length").(int) + length := rapid.IntMax(300).Draw(t, "length").(int) addrs := randomAddresses(length) addressSet := unique.New(addrs) @@ -199,3 +219,13 @@ func TestRandomAddresses(t *testing.T) { } }) } + +func toAddresses(vals []*valset.Validator) []common.Address { + addrs := make([]common.Address, len(vals)) + + for i, val := range vals { + addrs[i] = val.Address + } + + return addrs +}