go-ethereum/trie/zk_trie_impl_test.go
Ho c516a9e477
zktrie part2: add zktrie; allow switch trie type by config; (#113)
* induce zktrie

* refactoring zktrie

* fix crash issue in logger

* renaming JSON field

* unify hash scheme

* goimport and mod lint

* backward compatible with go 1.17

* lints

* add option on genesis file

* corrections according to the reviews

* trivial fixes: ValueKey entry, key in prove nodes

* fixing for the proof fix ...

* avoiding panic before loading stateDb in genesis setup

* revert ExtraData.StateList json annotation for compatibility

* fix goimports lint

* fix goimports lint

* better encoding for leaf node

* fix proof's printing issue, add handling on coinbase

* update genesis, and rule out snapshot in zktrie mode

* update readme and lint

* fix an issue

Co-authored-by: HAOYUatHZ <37070449+HAOYUatHZ@users.noreply.github.com>
Co-authored-by: HAOYUatHZ <haoyu@protonmail.com>
2022-06-27 11:17:02 +08:00

184 lines
6.2 KiB
Go

package trie
import (
"math/big"
"testing"
"github.com/iden3/go-iden3-crypto/constants"
cryptoUtils "github.com/iden3/go-iden3-crypto/utils"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/scroll-tech/go-ethereum/common"
"github.com/scroll-tech/go-ethereum/core/types"
zkt "github.com/scroll-tech/go-ethereum/core/types/zktrie"
"github.com/scroll-tech/go-ethereum/ethdb/memorydb"
)
type Fatalable interface {
Fatal(args ...interface{})
}
func newTestingMerkle(f Fatalable, numLevels int) *ZkTrieImpl {
mt, err := NewZkTrieImpl(NewZktrieDatabase((memorydb.New())), numLevels)
if err != nil {
f.Fatal(err)
return nil
}
return mt
}
func TestHashParsers(t *testing.T) {
h0 := zkt.NewHashFromBigInt(big.NewInt(0))
assert.Equal(t, "0", h0.String())
h1 := zkt.NewHashFromBigInt(big.NewInt(1))
assert.Equal(t, "1", h1.String())
h10 := zkt.NewHashFromBigInt(big.NewInt(10))
assert.Equal(t, "10", h10.String())
h7l := zkt.NewHashFromBigInt(big.NewInt(1234567))
assert.Equal(t, "1234567", h7l.String())
h8l := zkt.NewHashFromBigInt(big.NewInt(12345678))
assert.Equal(t, "12345678...", h8l.String())
b, ok := new(big.Int).SetString("4932297968297298434239270129193057052722409868268166443802652458940273154854", 10) //nolint:lll
assert.True(t, ok)
h := zkt.NewHashFromBigInt(b)
assert.Equal(t, "4932297968297298434239270129193057052722409868268166443802652458940273154854", h.BigInt().String()) //nolint:lll
assert.Equal(t, "49322979...", h.String())
assert.Equal(t, "265baaf161e875c372d08e50f52abddc01d32efc93e90290bb8b3d9ceb94e70a", h.Hex())
b1, err := zkt.NewBigIntFromHashBytes(b.Bytes())
assert.Nil(t, err)
assert.Equal(t, new(big.Int).SetBytes(b.Bytes()).String(), b1.String())
b2, err := zkt.NewHashFromBytes(b.Bytes())
assert.Nil(t, err)
assert.Equal(t, b.String(), b2.BigInt().String())
h2, err := zkt.NewHashFromHex(h.Hex())
assert.Nil(t, err)
assert.Equal(t, h, h2)
_, err = zkt.NewHashFromHex("0x12")
assert.NotNil(t, err)
// check limits
a := new(big.Int).Sub(constants.Q, big.NewInt(1))
testHashParsers(t, a)
a = big.NewInt(int64(1))
testHashParsers(t, a)
}
func testHashParsers(t *testing.T, a *big.Int) {
require.True(t, cryptoUtils.CheckBigIntInField(a))
h := zkt.NewHashFromBigInt(a)
assert.Equal(t, a, h.BigInt())
hFromBytes, err := zkt.NewHashFromBytes(h.Bytes())
assert.Nil(t, err)
assert.Equal(t, h, hFromBytes)
assert.Equal(t, a, hFromBytes.BigInt())
assert.Equal(t, a.String(), hFromBytes.BigInt().String())
hFromHex, err := zkt.NewHashFromHex(h.Hex())
assert.Nil(t, err)
assert.Equal(t, h, hFromHex)
aBIFromHBytes, err := zkt.NewBigIntFromHashBytes(h.Bytes())
assert.Nil(t, err)
assert.Equal(t, a, aBIFromHBytes)
assert.Equal(t, new(big.Int).SetBytes(a.Bytes()).String(), aBIFromHBytes.String())
}
func TestMerkleTree_AddUpdateGetWord(t *testing.T) {
mt := newTestingMerkle(t, 10)
err := mt.AddWord(&zkt.Byte32{1}, &zkt.Byte32{2})
assert.Nil(t, err)
err = mt.AddWord(&zkt.Byte32{3}, &zkt.Byte32{4})
assert.Nil(t, err)
err = mt.AddWord(&zkt.Byte32{5}, &zkt.Byte32{6})
assert.Nil(t, err)
err = mt.AddWord(&zkt.Byte32{5}, &zkt.Byte32{7})
assert.Equal(t, ErrEntryIndexAlreadyExists, err)
node, err := mt.GetLeafNodeByWord(&zkt.Byte32{1})
assert.Nil(t, err)
assert.Equal(t, len(node.ValuePreimage), 1)
assert.Equal(t, (&zkt.Byte32{2})[:], node.ValuePreimage[0][:])
node, err = mt.GetLeafNodeByWord(&zkt.Byte32{3})
assert.Nil(t, err)
assert.Equal(t, len(node.ValuePreimage), 1)
assert.Equal(t, (&zkt.Byte32{4})[:], node.ValuePreimage[0][:])
node, err = mt.GetLeafNodeByWord(&zkt.Byte32{5})
assert.Nil(t, err)
assert.Equal(t, len(node.ValuePreimage), 1)
assert.Equal(t, (&zkt.Byte32{6})[:], node.ValuePreimage[0][:])
err = mt.UpdateWord(&zkt.Byte32{1}, &zkt.Byte32{7})
assert.Nil(t, err)
err = mt.UpdateWord(&zkt.Byte32{3}, &zkt.Byte32{8})
assert.Nil(t, err)
err = mt.UpdateWord(&zkt.Byte32{5}, &zkt.Byte32{9})
assert.Nil(t, err)
node, err = mt.GetLeafNodeByWord(&zkt.Byte32{1})
assert.Nil(t, err)
assert.Equal(t, len(node.ValuePreimage), 1)
assert.Equal(t, (&zkt.Byte32{7})[:], node.ValuePreimage[0][:])
node, err = mt.GetLeafNodeByWord(&zkt.Byte32{3})
assert.Nil(t, err)
assert.Equal(t, len(node.ValuePreimage), 1)
assert.Equal(t, (&zkt.Byte32{8})[:], node.ValuePreimage[0][:])
node, err = mt.GetLeafNodeByWord(&zkt.Byte32{5})
assert.Nil(t, err)
assert.Equal(t, len(node.ValuePreimage), 1)
assert.Equal(t, (&zkt.Byte32{9})[:], node.ValuePreimage[0][:])
_, err = mt.GetLeafNodeByWord(&zkt.Byte32{100})
assert.Equal(t, ErrKeyNotFound, err)
}
func TestMerkleTree_UpdateAccount(t *testing.T) {
mt := newTestingMerkle(t, 10)
acc1 := &types.StateAccount{
Nonce: 1,
Balance: big.NewInt(10000000),
Root: common.HexToHash("22fb59aa5410ed465267023713ab42554c250f394901455a3366e223d5f7d147"),
CodeHash: common.HexToHash("cc0a77f6e063b4b62eb7d9ed6f427cf687d8d0071d751850cfe5d136bc60d3ab").Bytes(),
}
err := mt.TryUpdateAccount(common.HexToAddress("0x05fDbDfaE180345C6Cff5316c286727CF1a43327").Bytes(), acc1)
assert.Nil(t, err)
acc2 := &types.StateAccount{
Nonce: 5,
Balance: big.NewInt(50000000),
Root: common.HexToHash("0"),
CodeHash: common.HexToHash("c5d2460186f7233c927e7db2dcc703c0e500b653ca82273b7bfad8045d85a470").Bytes(),
}
err = mt.TryUpdateAccount(common.HexToAddress("0x4cb1aB63aF5D8931Ce09673EbD8ae2ce16fD6571").Bytes(), acc2)
assert.Nil(t, err)
bt, err := mt.TryGet(common.HexToAddress("0x05fDbDfaE180345C6Cff5316c286727CF1a43327").Bytes())
assert.Nil(t, err)
acc, err := types.UnmarshalStateAccount(bt)
assert.Nil(t, err)
assert.Equal(t, acc1.Nonce, acc.Nonce)
assert.Equal(t, acc1.Balance.Uint64(), acc.Balance.Uint64())
assert.Equal(t, acc1.Root.Bytes(), acc.Root.Bytes())
assert.Equal(t, acc1.CodeHash, acc.CodeHash)
bt, err = mt.TryGet(common.HexToAddress("0x4cb1aB63aF5D8931Ce09673EbD8ae2ce16fD6571").Bytes())
assert.Nil(t, err)
acc, err = types.UnmarshalStateAccount(bt)
assert.Nil(t, err)
assert.Equal(t, acc2.Nonce, acc.Nonce)
assert.Equal(t, acc2.Balance.Uint64(), acc.Balance.Uint64())
assert.Equal(t, acc2.Root.Bytes(), acc.Root.Bytes())
assert.Equal(t, acc2.CodeHash, acc.CodeHash)
bt, err = mt.TryGet(common.HexToAddress("0x8dE13967F19410A7991D63c2c0179feBFDA0c261").Bytes())
assert.Nil(t, err)
assert.Nil(t, bt)
}