Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion cmd/shisui/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -277,7 +277,8 @@ func initState(config Config, server *rpc.Server, conn discover.UDPConn, localNo
if err != nil {
return err
}
historyNetwork := state.NewStateNetwork(protocol)
client := rpc.DialInProc(server)
historyNetwork := state.NewStateNetwork(protocol, client)
return historyNetwork.Start()
}

Expand Down
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ require (
golang.org/x/tools v0.20.0
google.golang.org/protobuf v1.34.2
gopkg.in/natefinch/lumberjack.v2 v2.2.1
gopkg.in/yaml.v2 v2.4.0
gopkg.in/yaml.v3 v3.0.1
)

Expand Down Expand Up @@ -148,7 +149,6 @@ require (
go.uber.org/multierr v1.11.0 // indirect
golang.org/x/mod v0.17.0 // indirect
golang.org/x/net v0.24.0 // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
rsc.io/tmplfunc v0.0.3 // indirect
)

Expand Down
19 changes: 17 additions & 2 deletions portalnetwork/history/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,8 +58,23 @@ func (p *BlockHeaderProof) MarshalSSZ() ([]byte, error) {
return ssz.MarshalSSZ(p)
}

func (p *BlockHeaderProof) MarshalSSZTo(_ []byte) (dst []byte, err error) {
return ssz.MarshalSSZ(p)
func (p *BlockHeaderProof) MarshalSSZTo(buf []byte) (dst []byte, err error) {
dst = buf
dst = append(dst, byte(p.Selector))
if p.Selector != none {
if len(p.Proof) != 15 {
err = ssz.ErrBytesLengthFn("proofs size should be", len(p.Proof), 15)
return
}
for _, item := range p.Proof {
if len(item) != 32 {
err = ssz.ErrBytesLengthFn("single proof size should be", len(item), 32)
return
}
dst = append(dst, item...)
}
}
return
}

func (p *BlockHeaderProof) UnmarshalSSZ(buf []byte) (err error) {
Expand Down
161 changes: 157 additions & 4 deletions portalnetwork/state/network.go
Original file line number Diff line number Diff line change
@@ -1,27 +1,43 @@
package state

import (
"bytes"
"context"
"errors"
"fmt"
"time"

"github.com/ethereum/go-ethereum/common/hexutil"
"github.com/ethereum/go-ethereum/core/types"
"github.com/ethereum/go-ethereum/log"
"github.com/ethereum/go-ethereum/p2p/discover"
"github.com/ethereum/go-ethereum/portalnetwork/history"
"github.com/ethereum/go-ethereum/rlp"
"github.com/ethereum/go-ethereum/rpc"
"github.com/ethereum/go-ethereum/trie"
"github.com/protolambda/zrnt/eth2/beacon/common"
"github.com/protolambda/zrnt/eth2/configs"
"github.com/protolambda/ztyp/codec"
)

type StateNetwork struct {
portalProtocol *discover.PortalProtocol
closeCtx context.Context
closeFunc context.CancelFunc
log log.Logger
spec *common.Spec
client *rpc.Client
}

func NewStateNetwork(portalProtocol *discover.PortalProtocol) *StateNetwork {
func NewStateNetwork(portalProtocol *discover.PortalProtocol, client *rpc.Client) *StateNetwork {
ctx, cancel := context.WithCancel(context.Background())

return &StateNetwork{
portalProtocol: portalProtocol,
closeCtx: ctx,
closeFunc: cancel,
log: log.New("sub-protocol", "state"),
spec: configs.Mainnet,
client: client,
}
}

Expand Down Expand Up @@ -72,6 +88,143 @@ func (h *StateNetwork) processContentLoop(ctx context.Context) {
}

func (h *StateNetwork) validateContents(contentKeys [][]byte, contents [][]byte) error {
// TODO
panic("implement me")
for i, content := range contents {
contentKey := contentKeys[i]
err := h.validateContent(contentKey, content)
if err != nil {
h.log.Error("content validate failed", "contentKey", hexutil.Encode(contentKey), "content", hexutil.Encode(content), "err", err)
return fmt.Errorf("content validate failed with content key %x and content %x", contentKey, content)
}
contentId := h.portalProtocol.ToContentId(contentKey)
_ = h.portalProtocol.Put(contentKey, contentId, content)
}
return nil
}

func (h *StateNetwork) validateContent(contentKey []byte, content []byte) error {
keyType := contentKey[0]
switch keyType {
case AccountTrieNodeType:
return h.validateAccountTrieNode(contentKey[1:], content)
case ContractStorageTrieNodeType:
return validateContractStorageTrieNode(h.spec, contentKey[1:], content)
case ContractByteCodeType:
return validateContractByteCode(h.spec, contentKey[1:], content)
}
return errors.New("unknown content type")
}

func (h *StateNetwork) validateAccountTrieNode(contentKey []byte, content []byte) error {
accountKey := &AccountTrieNodeKey{}
err := accountKey.Deserialize(codec.NewDecodingReader(bytes.NewReader(contentKey), uint64(len(contentKey))))
if err != nil {
return err
}
accountData := &AccountTrieNodeWithProof{}
err = accountData.Deserialize(codec.NewDecodingReader(bytes.NewReader(content), uint64(len(content))))
if err != nil {
return err
}
// get HeaderWithProof in history network
stateRoot, err := h.getStateRoot(accountData.BlockHash)

if err != nil {
return err
}
err = validateNodeTrieProof(stateRoot, accountKey.NodeHash, &accountKey.Path, &accountData.Proof)
return err
}

func validateContractStorageTrieNode(spec *common.Spec, contentKey []byte, content []byte) error {
return nil
}

func validateContractByteCode(spec *common.Spec, contentKey []byte, content []byte) error {
return nil
}

func (h *StateNetwork) getStateRoot(blockHash common.Bytes32) (common.Bytes32, error) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second*2)
defer cancel()
contentKey := make([]byte, 0)
contentKey = append(contentKey, byte(history.BlockHeaderType))
contentKey = append(contentKey, blockHash[:]...)

arg := hexutil.Encode(contentKey)
res := &discover.ContentInfo{}
err := h.client.CallContext(ctx, res, "portal_historyRecursiveFindContent", arg)
if err != nil {
return common.Bytes32{}, err
}
data, err := hexutil.Decode(res.Content)
if err != nil {
return common.Bytes32{}, err
}
headerWithProof, err := history.DecodeBlockHeaderWithProof(data)
if err != nil {
return common.Bytes32{}, err
}
header := new(types.Header)
err = rlp.DecodeBytes(headerWithProof.Header, header)
if err != nil {
return common.Bytes32{}, err
}
return common.Bytes32(header.Root), nil
}

func validateNodeTrieProof(rootHash common.Bytes32, nodeHash common.Bytes32, path *Nibbles, proof *TrieProof) error {
lastNode, p, err := validateTrieProof(rootHash, path.Nibbles, proof)
if err != nil {
return err
}
if len(p) != 0 {
return errors.New("path is too long")
}
err = checkNodeHash(&lastNode, nodeHash[:])
if err != nil {
return err
}
return nil
}

func validateTrieProof(rootHash common.Bytes32, path []byte, proof *TrieProof) (EncodedTrieNode, []byte, error) {
if len(*proof) == 0 {
return nil, nil, errors.New("proof should be empty")
}
firstNode := []EncodedTrieNode(*proof)[0]
err := checkNodeHash(&firstNode, rootHash[:])
if err != nil {
return nil, nil, err
}

node := firstNode
remainingPath := path

for _, nextNode := range []EncodedTrieNode(*proof)[1:] {
n, err := trie.DecodeTrieNode(nil, node)
if err != nil {
return nil, nil, err
}
hashNode, p, err := trie.TraverseTrieNode(n, remainingPath)
if err != nil {
return nil, nil, err
}
err = checkNodeHash(&nextNode, hashNode)

if err != nil {
return nil, nil, err
}

node = nextNode
remainingPath = p
}
return node, remainingPath, nil
}

func checkNodeHash(node *EncodedTrieNode, hash []byte) error {
nodeHash := node.NodeHash()
if !bytes.Equal(nodeHash[:], hash[:]) {
return fmt.Errorf("node hash is not equal, expect: %v, actual: %v", hash, nodeHash)
}
return nil
}
73 changes: 73 additions & 0 deletions portalnetwork/state/network_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
package state

import (
"fmt"
"os"
"testing"

"github.com/ethereum/go-ethereum/common/hexutil"
"github.com/ethereum/go-ethereum/p2p/discover"
"github.com/ethereum/go-ethereum/portalnetwork/history"
"github.com/ethereum/go-ethereum/rpc"
"github.com/stretchr/testify/require"
"gopkg.in/yaml.v2"
)

type TestCase struct {
BlockHeader string `yaml:"block_header"`
ContentKey string `yaml:"content_key"`
ContentValueOffer string `yaml:"content_value_offer"`
ContentValueRetrieval string `yaml:"content_value_retrieval"`
}

func getTestCases(filename string) ([]TestCase, error) {
file, err := os.ReadFile(fmt.Sprintf("./testdata/%s", filename))
if err != nil {
return nil, err
}
res := make([]TestCase, 0)
err = yaml.Unmarshal(file, &res)
if err != nil {
return nil, err
}
return res, nil
}

type MockAPI struct {
header string
}

func (p *MockAPI) HistoryRecursiveFindContent(contentKeyHex string) (*discover.ContentInfo, error) {
headerWithProof := &history.BlockHeaderWithProof{
Header: hexutil.MustDecode(p.header),
Proof: &history.BlockHeaderProof{
Selector: 0,
Proof: [][]byte{},
},
}
data, err := headerWithProof.MarshalSSZ()
if err != nil {
return nil, err
}
return &discover.ContentInfo{
Content: hexutil.Encode(data),
UtpTransfer: false,
}, nil
}

func TestValidateAccountTrieNode(t *testing.T) {
cases, err := getTestCases("account_trie_node.yaml")
require.NoError(t, err)

for _, tt := range cases {
server := rpc.NewServer()
api := &MockAPI{
header: tt.BlockHeader,
}
server.RegisterName("portal", api)
client := rpc.DialInProc(server)
bn := NewStateNetwork(nil, client)
err = bn.validateContent(hexutil.MustDecode(tt.ContentKey), hexutil.MustDecode(tt.ContentValueOffer))
require.NoError(t, err)
}
}
43 changes: 43 additions & 0 deletions portalnetwork/state/testdata/account_trie_node.yaml

Large diffs are not rendered by default.

Loading