verify penalty info in header check point

This commit is contained in:
Nguyen Ba Tam 2018-10-11 16:14:08 +07:00
parent 2bf1906342
commit ad4b0f223c
7 changed files with 102 additions and 54 deletions

View file

@ -11,4 +11,5 @@ const (
EpocBlockOpening = 850
EpocBlockRandomize = 900
MaxMasternodes = 150
LimitPenaltyEpoch = 4
)

View file

@ -243,3 +243,38 @@ func (a *UnprefixedAddress) UnmarshalText(input []byte) error {
func (a UnprefixedAddress) MarshalText() ([]byte, error) {
return []byte(hex.EncodeToString(a[:])), nil
}
// Extract validators from byte array.
func RemoveItemFromArray(array []Address, items []Address) []Address {
if items == nil || len(items) == 0 {
return array
}
for i, value := range array {
for _, item := range items {
if value == item {
array = append(array[:i], array[i+1:]...)
}
}
}
return array
}
// Extract validators from byte array.
func ExtractAddressToBytes(penalties []Address) []byte {
data := []byte{}
for _, signer := range penalties {
data = append(data, signer[:]...)
}
return data
}
func ExtractAddressFromBytes(bytePenalties []byte) []Address {
if bytePenalties != nil && len(bytePenalties) < AddressLength {
return []Address{}
}
penalties := make([]Address, len(bytePenalties)/AddressLength)
for i := 0; i < len(penalties); i++ {
copy(penalties[i][:], bytePenalties[i*AddressLength:])
}
return penalties
}

View file

@ -149,3 +149,12 @@ func BenchmarkAddressHex(b *testing.B) {
testAddr.Hex()
}
}
func TestRemoveItemInArray(t *testing.T) {
array := []Address{HexToAddress("0x0000000"), HexToAddress("0x0000001"), HexToAddress("0x0000002")}
remove := []Address{HexToAddress("0x0000000"), HexToAddress("0x0000004"), HexToAddress("0x0000003")}
array = RemoveItemFromArray(array, remove)
if len(array) != 2 {
t.Error("fail remove item from array addres ")
}
}

View file

@ -109,6 +109,8 @@ var (
// ones).
errInvalidCheckpointSigners = errors.New("invalid signer list on checkpoint block")
errInvalidCheckpointPenalties = errors.New("invalid penalty list on checkpoint block")
// errInvalidMixDigest is returned if a block's mix digest is non-zero.
errInvalidMixDigest = errors.New("non-zero mix digest")
@ -363,12 +365,30 @@ func (c *Posv) verifyCascadingFields(chain consensus.ChainReader, header *types.
}
// If the block is a checkpoint block, verify the signer list
if number%c.config.Epoch == 0 {
signers := make([]byte, len(snap.Signers)*common.AddressLength)
for i, signer := range snap.signers() {
copy(signers[i*common.AddressLength:], signer[:])
penPenalties := []common.Address{}
if c.HookPenalty != nil {
penPenalties, err = c.HookPenalty(chain, number)
if err != nil {
return err
}
for _, address := range penPenalties {
log.Debug("Penalty Info", "address", address, "number", number)
}
bytePenalties := common.ExtractAddressToBytes(penPenalties)
if !bytes.Equal(header.Penalties, bytePenalties) {
return errInvalidCheckpointPenalties
}
}
signers := snap.signers()
signers = common.RemoveItemFromArray(signers, penPenalties)
for i := 1; i <= common.LimitPenaltyEpoch; i++ {
if number > uint64(i)*c.config.Epoch {
signers = RemovePenaltiesFromBlock(chain, signers, number-uint64(i)*c.config.Epoch)
}
}
byteMasterNodes := common.ExtractAddressToBytes(signers)
extraSuffix := len(header.Extra) - extraSeal
if !bytes.Equal(header.Extra[extraVanity:extraSuffix], signers) {
if !bytes.Equal(header.Extra[extraVanity:extraSuffix], byteMasterNodes) {
return errInvalidCheckpointSigners
}
}
@ -636,52 +656,28 @@ func (c *Posv) Prepare(chain consensus.ChainReader, header *types.Header) error
header.Extra = append(header.Extra, bytes.Repeat([]byte{0x00}, extraVanity-len(header.Extra))...)
}
header.Extra = header.Extra[:extraVanity]
signers := snap.signers()
if number%c.config.Epoch == 0 {
signers := snap.signers()
if c.HookPenalty != nil {
penSigners, _ := c.HookPenalty(chain, number)
penSigners, err := c.HookPenalty(chain, number)
if err != nil {
return err
}
if len(penSigners) > 0 {
// Keep remove penalty signer out of signer list.
for i, signer := range signers {
for _, penSigner := range penSigners {
if signer == penSigner {
signers = append(signers[:i], signers[i+1:]...)
}
}
}
log.Debug("Penalty Info", "signers", penSigners, "number", number)
for _, penSigner := range penSigners {
header.Penalties = append(header.Penalties, penSigner[:]...)
signers = common.RemoveItemFromArray(signers, penSigners)
for _, address := range penSigners {
log.Debug("Penalty Info", "address", address, "number", number)
}
header.Penalties = common.ExtractAddressToBytes(penSigners)
}
}
// Prevent penaltied signer in 4 epocs ago jump into signer list.
for i := 1; i <= 4; i++ {
checkEpoc := uint64(i) * c.config.Epoch
if number > checkEpoc {
prevEpoc := number - checkEpoc
prevHeader := chain.GetHeaderByNumber(prevEpoc)
prevEpocBlock := chain.GetBlock(prevHeader.Hash(), prevEpoc)
penalties := prevEpocBlock.Penalties()
if penalties != nil {
prevSigners := ExtractPenaltiesFromBytes(penalties)
if len(prevSigners) > 0 {
for i, signer := range signers {
for _, preventSigner := range prevSigners {
if signer == preventSigner {
signers = append(signers[:i], signers[i+1:]...)
}
}
}
}
}
for i := 1; i <= common.LimitPenaltyEpoch; i++ {
if number > uint64(i)*c.config.Epoch {
signers = RemovePenaltiesFromBlock(chain, signers, number-uint64(i)*c.config.Epoch)
}
}
for _, signer := range signers {
header.Extra = append(header.Extra, signer[:]...)
}
@ -700,12 +696,9 @@ func (c *Posv) Prepare(chain consensus.ChainReader, header *types.Header) error
if header.Time.Int64() < time.Now().Unix() {
header.Time = big.NewInt(time.Now().Unix())
}
if c.HookPrepare != nil {
signers := snap.signers()
c.HookPrepare(header, signers)
}
return nil
}
@ -897,10 +890,16 @@ func (c *Posv) GetMasternodesFromCheckpointHeader(preCheckpointHeader *types.Hea
}
// Extract validators from byte array.
func ExtractPenaltiesFromBytes(bytePenalties []byte) []common.Address {
penalties := make([]common.Address, len(bytePenalties)/common.AddressLength)
for i := 0; i < len(penalties); i++ {
copy(penalties[i][:], bytePenalties[i*common.AddressLength:])
func RemovePenaltiesFromBlock(chain consensus.ChainReader, signers []common.Address, epochNumber uint64) []common.Address {
if epochNumber <= 0 {
return signers
}
return penalties
header := chain.GetHeaderByNumber(epochNumber)
block := chain.GetBlock(header.Hash(), epochNumber)
penalties := block.Penalties()
if penalties != nil {
prevPenalties := common.ExtractAddressFromBytes(penalties)
signers = common.RemoveItemFromArray(signers, prevPenalties)
}
return signers
}

View file

@ -238,6 +238,7 @@ func New(ctx *node.ServiceContext, config *Config) (*Ethereum, error) {
}
if len(m2) > 0 {
header.Validators = contracts.BuildValidatorFromM2(m2)
log.Debug("New set Validators", "m2", m2, "number", header.Number.Uint64())
}
}
}
@ -247,10 +248,10 @@ func New(ctx *node.ServiceContext, config *Config) (*Ethereum, error) {
c.HookPenalty = func(chain consensus.ChainReader, blockNumberEpoc uint64) ([]common.Address, error) {
client, err := eth.blockchain.GetClient()
if err != nil {
log.Error("Fail to connect IPC client for blockSigner", "error", err)
return nil, err
}
prevEpoc := blockNumberEpoc - chain.Config().Posv.Epoch
if prevEpoc > 0 {
if prevEpoc >= 0 {
prevHeader := chain.GetHeaderByNumber(prevEpoc)
penSigners := c.GetMasternodes(chain, prevHeader)
if len(penSigners) > 0 {
@ -259,7 +260,10 @@ func New(ctx *node.ServiceContext, config *Config) (*Ethereum, error) {
for i := prevEpoc; i <= blockNumberEpoc; i++ {
blockHeader := chain.GetHeaderByNumber(i)
if len(penSigners) > 0 {
signedMasternodes, _ := contracts.GetSignersFromContract(blockSignerAddr, client, blockHeader.Hash())
signedMasternodes, err := contracts.GetSignersFromContract(blockSignerAddr, client, blockHeader.Hash())
if err != nil {
return nil, err
}
if len(signedMasternodes) > 0 {
// Check signer signed?
for _, signed := range signedMasternodes {

View file

@ -63,8 +63,8 @@ func testStatusMsgErrors(t *testing.T, protocol int) {
wantError: errResp(ErrProtocolVersionMismatch, "10 (!= %d)", protocol),
},
{
code: StatusMsg, data: statusData{uint32(protocol), 89, td, head.Hash(), genesis.Hash()},
wantError: errResp(ErrNetworkIdMismatch, "89 (!= 1)"),
code: StatusMsg, data: statusData{uint32(protocol), 999, td, head.Hash(), genesis.Hash()},
wantError: errResp(ErrNetworkIdMismatch, "999 (!= 89)"),
},
{
code: StatusMsg, data: statusData{uint32(protocol), DefaultConfig.NetworkId, td, head.Hash(), common.Hash{3}},