From 152b564ef7afd40b1c82c10497e46591594c2367 Mon Sep 17 00:00:00 2001 From: NguyenNguyen Date: Thu, 4 Apr 2019 14:20:11 +0700 Subject: [PATCH] Refactoring and adding unit test of compareSignersLists --- consensus/posv/posv.go | 21 +++++++++++++-------- consensus/posv/posv_test.go | 26 ++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/consensus/posv/posv.go b/consensus/posv/posv.go index 35177163ac..24cb26d059 100644 --- a/consensus/posv/posv.go +++ b/consensus/posv/posv.go @@ -428,14 +428,7 @@ func (c *Posv) verifyCascadingFields(chain consensus.ChainReader, header *types. } extraSuffix := len(header.Extra) - extraSeal masternodesFromCheckpointHeader := common.ExtractAddressFromBytes(header.Extra[extraVanity:extraSuffix]) - validSigners := true - sort.Slice(masternodesFromCheckpointHeader, func(i, j int) bool { - return masternodesFromCheckpointHeader[i].String() <= masternodesFromCheckpointHeader[j].String() - }) - sort.Slice(signers, func(i, j int) bool { - return signers[i].String() <= signers[j].String() - }) - validSigners = reflect.DeepEqual(masternodesFromCheckpointHeader, signers) + validSigners := compareSignersLists(masternodesFromCheckpointHeader, signers) if !validSigners { log.Error("Masternodes lists are different in checkpoint header and snapshot", "number", number, "masternodes_from_checkpoint_header", masternodesFromCheckpointHeader, "masternodes_in_snapshot", signers, "penList", penPenalties) return errInvalidCheckpointSigners @@ -451,6 +444,18 @@ func (c *Posv) verifyCascadingFields(chain consensus.ChainReader, header *types. return c.verifySeal(chain, header, parents, fullVerify) } +// compare 2 signers lists +// return true if they are same elements, otherwise return false +func compareSignersLists(list1 []common.Address, list2 []common.Address) bool { + sort.Slice(list1, func(i, j int) bool { + return list1[i].String() <= list1[j].String() + }) + sort.Slice(list2, func(i, j int) bool { + return list2[i].String() <= list2[j].String() + }) + return reflect.DeepEqual(list1, list2) +} + func (c *Posv) GetSnapshot(chain consensus.ChainReader, header *types.Header) (*Snapshot, error) { number := header.Number.Uint64() log.Trace("take snapshot", "number", number, "hash", header.Hash()) diff --git a/consensus/posv/posv_test.go b/consensus/posv/posv_test.go index 1c0a5baeb1..1ccf9f7860 100644 --- a/consensus/posv/posv_test.go +++ b/consensus/posv/posv_test.go @@ -47,3 +47,29 @@ func TestGetM1M2FromCheckpointHeader(t *testing.T) { } } } + +func TestCompareSignersLists(t *testing.T) { + list1 := []common.Address{ + common.StringToAddress("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"), + common.StringToAddress("bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"), + common.StringToAddress("cccccccccccccccccccccccccccccccccccccccc"), + common.StringToAddress("dddddddddddddddddddddddddddddddddddddddd"), + } + list2 := []common.Address{ + common.StringToAddress("cccccccccccccccccccccccccccccccccccccccc"), + common.StringToAddress("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"), + common.StringToAddress("dddddddddddddddddddddddddddddddddddddddd"), + common.StringToAddress("bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"), + } + list3 := []common.Address{ + common.StringToAddress("cccccccccccccccccccccccccccccccccccccccc"), + common.StringToAddress("dddddddddddddddddddddddddddddddddddddddd"), + common.StringToAddress("bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"), + } + if !compareSignersLists(list1, list2) { + t.Error("list1 should be equal to list2", "list1", list1, "list2", list2) + } + if compareSignersLists(list1, list3) { + t.Error("list1 and list3 should not be same", "list1", list1, "list3", list3) + } +}