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 EpocBlockOpening = 850
EpocBlockRandomize = 900 EpocBlockRandomize = 900
MaxMasternodes = 150 MaxMasternodes = 150
LimitPenaltyEpoch = 4
) )

View file

@ -243,3 +243,38 @@ func (a *UnprefixedAddress) UnmarshalText(input []byte) error {
func (a UnprefixedAddress) MarshalText() ([]byte, error) { func (a UnprefixedAddress) MarshalText() ([]byte, error) {
return []byte(hex.EncodeToString(a[:])), nil 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() 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). // ones).
errInvalidCheckpointSigners = errors.New("invalid signer list on checkpoint block") 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 is returned if a block's mix digest is non-zero.
errInvalidMixDigest = errors.New("non-zero mix digest") 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 the block is a checkpoint block, verify the signer list
if number%c.config.Epoch == 0 { if number%c.config.Epoch == 0 {
signers := make([]byte, len(snap.Signers)*common.AddressLength) penPenalties := []common.Address{}
for i, signer := range snap.signers() { if c.HookPenalty != nil {
copy(signers[i*common.AddressLength:], signer[:]) 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 extraSuffix := len(header.Extra) - extraSeal
if !bytes.Equal(header.Extra[extraVanity:extraSuffix], signers) { if !bytes.Equal(header.Extra[extraVanity:extraSuffix], byteMasterNodes) {
return errInvalidCheckpointSigners 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 = append(header.Extra, bytes.Repeat([]byte{0x00}, extraVanity-len(header.Extra))...)
} }
header.Extra = header.Extra[:extraVanity] header.Extra = header.Extra[:extraVanity]
signers := snap.signers()
if number%c.config.Epoch == 0 { if number%c.config.Epoch == 0 {
signers := snap.signers()
if c.HookPenalty != nil { if c.HookPenalty != nil {
penSigners, _ := c.HookPenalty(chain, number) penSigners, err := c.HookPenalty(chain, number)
if err != nil {
return err
}
if len(penSigners) > 0 { if len(penSigners) > 0 {
// Keep remove penalty signer out of signer list. // Keep remove penalty signer out of signer list.
for i, signer := range signers { signers = common.RemoveItemFromArray(signers, penSigners)
for _, penSigner := range penSigners { for _, address := range penSigners {
if signer == penSigner { log.Debug("Penalty Info", "address", address, "number", number)
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[:]...)
} }
header.Penalties = common.ExtractAddressToBytes(penSigners)
} }
} }
// Prevent penaltied signer in 4 epocs ago jump into signer list. // Prevent penaltied signer in 4 epocs ago jump into signer list.
for i := 1; i <= 4; i++ { for i := 1; i <= common.LimitPenaltyEpoch; i++ {
checkEpoc := uint64(i) * c.config.Epoch if number > uint64(i)*c.config.Epoch {
if number > checkEpoc { signers = RemovePenaltiesFromBlock(chain, signers, number-uint64(i)*c.config.Epoch)
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 _, signer := range signers { for _, signer := range signers {
header.Extra = append(header.Extra, signer[:]...) 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() { if header.Time.Int64() < time.Now().Unix() {
header.Time = big.NewInt(time.Now().Unix()) header.Time = big.NewInt(time.Now().Unix())
} }
if c.HookPrepare != nil { if c.HookPrepare != nil {
signers := snap.signers()
c.HookPrepare(header, signers) c.HookPrepare(header, signers)
} }
return nil return nil
} }
@ -897,10 +890,16 @@ func (c *Posv) GetMasternodesFromCheckpointHeader(preCheckpointHeader *types.Hea
} }
// Extract validators from byte array. // Extract validators from byte array.
func ExtractPenaltiesFromBytes(bytePenalties []byte) []common.Address { func RemovePenaltiesFromBlock(chain consensus.ChainReader, signers []common.Address, epochNumber uint64) []common.Address {
penalties := make([]common.Address, len(bytePenalties)/common.AddressLength) if epochNumber <= 0 {
for i := 0; i < len(penalties); i++ { return signers
copy(penalties[i][:], bytePenalties[i*common.AddressLength:])
} }
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 { if len(m2) > 0 {
header.Validators = contracts.BuildValidatorFromM2(m2) 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) { c.HookPenalty = func(chain consensus.ChainReader, blockNumberEpoc uint64) ([]common.Address, error) {
client, err := eth.blockchain.GetClient() client, err := eth.blockchain.GetClient()
if err != nil { if err != nil {
log.Error("Fail to connect IPC client for blockSigner", "error", err) return nil, err
} }
prevEpoc := blockNumberEpoc - chain.Config().Posv.Epoch prevEpoc := blockNumberEpoc - chain.Config().Posv.Epoch
if prevEpoc > 0 { if prevEpoc >= 0 {
prevHeader := chain.GetHeaderByNumber(prevEpoc) prevHeader := chain.GetHeaderByNumber(prevEpoc)
penSigners := c.GetMasternodes(chain, prevHeader) penSigners := c.GetMasternodes(chain, prevHeader)
if len(penSigners) > 0 { if len(penSigners) > 0 {
@ -259,7 +260,10 @@ func New(ctx *node.ServiceContext, config *Config) (*Ethereum, error) {
for i := prevEpoc; i <= blockNumberEpoc; i++ { for i := prevEpoc; i <= blockNumberEpoc; i++ {
blockHeader := chain.GetHeaderByNumber(i) blockHeader := chain.GetHeaderByNumber(i)
if len(penSigners) > 0 { 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 { if len(signedMasternodes) > 0 {
// Check signer signed? // Check signer signed?
for _, signed := range signedMasternodes { for _, signed := range signedMasternodes {

View file

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