diff --git a/common/yerror/basic_error.go b/common/yerror/basic_error.go index fb7b217..76c0972 100644 --- a/common/yerror/basic_error.go +++ b/common/yerror/basic_error.go @@ -20,10 +20,11 @@ var NoKvdbType = errors.New("no kvdb type") var NoSqlDbType = errors.New("no sqlDB type") var ( - PoolOverflow error = errors.New("pool size is full") - TxnTimeoutErr error = errors.New("Txn time out") - TxnTooLarge error = errors.New("the size of txn is too large") - TxnDuplicated error = errors.New("Transaction duplicated") + PoolOverflow error = errors.New("pool size is full") + TxnTimeoutErr error = errors.New("Txn time out") + TxnTooLarge error = errors.New("the size of txn is too large") + TxnDuplicated error = errors.New("Transaction duplicated") + ChainIDIllegal error = errors.New("chain id illegal") ) var ErrBlockNotFound error = errors.New("block not found") diff --git a/core/kernel/handle_input.go b/core/kernel/handle_input.go index c487382..3cd08b4 100644 --- a/core/kernel/handle_input.go +++ b/core/kernel/handle_input.go @@ -70,10 +70,11 @@ func (k *Kernel) handleTxnLocally(stxn *SignedTxn, topic string) error { return err } } - if k.CheckReplayAttack(stxn) { - return yerror.TxnDuplicated + err := k.CheckReplayAttack(stxn) + if err != nil { + return err } - err := k.Pool.CheckTxn(stxn) + err = k.Pool.CheckTxn(stxn) if err != nil { return err } @@ -97,14 +98,18 @@ func (k *Kernel) HandleReading(rdCall *common.RdCall) (*context.ResponseData, er return ctx.Response(), nil } -func (k *Kernel) CheckReplayAttack(txn *SignedTxn) bool { +func (k *Kernel) CheckReplayAttack(txn *SignedTxn) error { + if k.Chain.ChainID() != txn.ChainID() { + return yerror.ChainIDIllegal + } if k.Pool.Exist(txn.TxnHash) { - return true + return yerror.TxnDuplicated } - if k.Chain.ChainID() != txn.ChainID() { - return true + + if k.TxDB.ExistTxn(txn.TxnHash) { + return yerror.TxnDuplicated } - return k.TxDB.ExistTxn(txn.TxnHash) + return nil } //func getRdFromHttp(req *http.Request, params string) (rdCall *RdCall, err error) { diff --git a/core/kernel/kernel.go b/core/kernel/kernel.go index 8d216b7..2fab2c5 100644 --- a/core/kernel/kernel.go +++ b/core/kernel/kernel.go @@ -101,7 +101,7 @@ func (k *Kernel) AcceptUnpackedTxns() error { } for _, txn := range writings { - if k.CheckReplayAttack(txn) { + if err := k.CheckReplayAttack(txn); err != nil { continue } txn.FromP2P = true @@ -130,7 +130,7 @@ func (k *Kernel) AcceptUnpackedTxns() error { if txn == nil { continue } - if k.CheckReplayAttack(txn) { + if err := k.CheckReplayAttack(txn); err != nil { continue } txn.FromP2P = true