diff --git a/lib/xchain/merkle.go b/lib/xchain/merkle.go index c1cba5c8..1ba5f077 100644 --- a/lib/xchain/merkle.go +++ b/lib/xchain/merkle.go @@ -102,6 +102,15 @@ func msgLeaf(msg Msg) ([32]byte, error) { return merkle.StdLeafHash(DSTMessage, bz), nil } +func MsgLeaf(msg Msg) ([32]byte, error) { + bz, err := encodeMsg(msg) + if err != nil { + return [32]byte{}, errors.Wrap(err, "encode message") + } + + return merkle.StdLeafHash(DSTMessage, bz), nil +} + func submissionHeaderLeaf(attHeader AttestHeader, blockHeader BlockHeader) ([32]byte, error) { bz, err := encodeSubmissionHeader(attHeader, blockHeader) if err != nil { diff --git a/relayer/app/creator_test.go b/relayer/app/creator_test.go index c7978a3b..2b99b911 100644 --- a/relayer/app/creator_test.go +++ b/relayer/app/creator_test.go @@ -4,6 +4,7 @@ import ( "math/rand" "testing" "time" + "encoding/hex" "github.com/omni-network/omni/halo/attest/voter" "github.com/omni-network/omni/lib/cchain" @@ -20,6 +21,100 @@ import ( "github.com/stretchr/testify/require" ) +func TestMsgTreeOrdering(t *testing.T) { + t.Parallel() + + var ( + SourceChainID = uint64(1) + DestChainID = uint64(2) + ) + + privKey := k1.GenPrivKey() + addr, err := k1util.PubKeyToAddress(privKey.PubKey()) + require.NoError(t, err) + + fuzzer := fuzz.New().NilChance(0).Funcs( + func(e *xchain.Msg, c fuzz.Continue) { + e.DestChainID = DestChainID + e.SourceChainID = SourceChainID + e.DestAddress = common.Address(crypto.CRandBytes(20)) + e.SourceMsgSender = common.Address(crypto.CRandBytes(20)) + e.Data = crypto.CRandBytes(100) + }, + ) + + var block xchain.Block + fuzzer.NilChance(0).NumElements(4, 4).Fuzz(&block) + require.NotEmpty(t, block.Msgs) + // Ensure msg.LogIndex is increasing + for i := 1; i < len(block.Msgs); i++ { + block.Msgs[i].LogIndex = block.Msgs[i-1].LogIndex + 1 + uint64(rand.Intn(1000)) + } + + var attestHeader xchain.AttestHeader + fuzzer.Fuzz(&attestHeader) + attestHeader.ChainVersion.ID = block.ChainID // Align headers + + var valSetID uint64 + fuzzer.Fuzz(&valSetID) + + // make all msg offset sequential + for i := range block.Msgs { + block.Msgs[i].StreamOffset = uint64(i) + } + + vote, err := voter.CreateVote(privKey, attestHeader, block) + require.NoError(t, err) + header, err := vote.BlockHeader.ToXChain() + require.NoError(t, err) + require.Equal(t, block.BlockHeader, header) + require.Equal(t, addr, common.Address(vote.Signature.ValidatorAddress)) + + tree, err := xchain.NewMsgTree(block.Msgs) + require.NoError(t, err) + treeRoot := tree.MsgRoot() + t.Log("***** MsgTree root:", hex.EncodeToString(treeRoot[:])) + + // create new Msgs + cMsgs := make([]xchain.Msg, len(block.Msgs)) + copy(cMsgs, block.Msgs) + + s := len(cMsgs)-1 + a, b := cMsgs[s], cMsgs[s-1] + + // swap last two messages + cMsgs[s], cMsgs[s-1] = b, a + // swap LogIndices + cMsgs[s].LogIndex = a.LogIndex + cMsgs[s-1].LogIndex = b.LogIndex + lastLeafOrig, err := xchain.MsgLeaf(block.Msgs[s]) + require.NoError(t, err) + lastLeafNew, err := xchain.MsgLeaf(cMsgs[s]) + require.NoError(t, err) + require.NotEqual(t, lastLeafOrig, lastLeafNew) + + // recalculate the tree + newTree, err := xchain.NewMsgTree(cMsgs) + require.NoError(t, err) + newTreeRoot := newTree.MsgRoot() + t.Log("***** MsgTree root:", hex.EncodeToString(newTreeRoot[:])) + + blockHeader, err := vote.BlockHeader.ToXChain() + require.NoError(t, err) + + sig, err := vote.Signature.ToXChain() + require.NoError(t, err) + + att := xchain.Attestation{ + AttestHeader: vote.AttestHeader.ToXChain(), + BlockHeader: blockHeader, + ValidatorSetID: valSetID, + MsgRoot: [32]byte(vote.MsgRoot), + Signatures: []xchain.SigTuple{sig}, + } + t.Log("***** Attestation MsgTree root:", hex.EncodeToString(att.MsgRoot[:])) +} + func TestCreatorService_CreateSubmissions(t *testing.T) { t.Parallel()