diff --git a/les/odr_requests.go b/les/odr_requests.go index 7a5c7c26fb..6da790e5af 100644 --- a/les/odr_requests.go +++ b/les/odr_requests.go @@ -370,19 +370,20 @@ func (r *ChtRequest) CanSend(peer *peer, config *light.IndexerConfig) bool { // Request sends an ODR request to the LES network (implementation of LesOdrRequest) func (r *ChtRequest) Request(reqID uint64, peer *peer, config *light.IndexerConfig) error { - peer.Log().Debug("Requesting CHT", "cht", r.ChtNum, "block", r.Numbers) + peer.Log().Debug("Requesting CHT", "cht", r.ChtNum, "blocks", r.Numbers) + var ( encNum [8]byte - reqs = make([]HelperTrieReq, 0, len(r.Numbers)) + reqs = make([]HelperTrieReq, len(r.Numbers)) ) - for _, num := range r.Numbers { + for i, num := range r.Numbers { binary.BigEndian.PutUint64(encNum[:], num) - reqs = append(reqs, HelperTrieReq{ + reqs[i] = HelperTrieReq{ Type: htCanonical, TrieIdx: r.ChtNum, Key: common.CopyBytes(encNum[:]), AuxReq: auxHeader, - }) + } } switch peer.version { case lpv1: @@ -426,7 +427,6 @@ func (r *ChtRequest) Validate(db ethdb.Database, msg *Msg) error { binary.BigEndian.PutUint64(encNumber[:], num) value, _, err := trie.VerifyProof(r.ChtRoot, encNumber[:], light.NodeList(resp.Proof).NodeSet()) if err != nil { - fmt.Println(err) return err } var node light.ChtNode diff --git a/les/odr_test.go b/les/odr_test.go index ee1f7d5a4f..2eb400f1d0 100644 --- a/les/odr_test.go +++ b/les/odr_test.go @@ -36,13 +36,17 @@ import ( "github.com/ethereum/go-ethereum/rlp" ) -type chtTestFn func(ctx context.Context, bc *core.BlockChain, lc *light.LightChain, number uint64) []byte +type chtTestFn func(ctx context.Context, bc *core.BlockChain, lc *light.LightChain, numbers []uint64) []byte -func TestChtGetHeadersLes1(t *testing.T) { testCht(t, 1, chtGetHeader) } +func TestChtGetHeaderLes1(t *testing.T) { testCht(t, 1, 1, chtGetHeader) } -func TestChtGetHeadersLes2(t *testing.T) { testCht(t, 2, chtGetHeader) } +func TestChtGetHeaderLes2(t *testing.T) { testCht(t, 2, 1, chtGetHeader) } -func chtGetHeader(ctx context.Context, bc *core.BlockChain, lc *light.LightChain, number uint64) []byte { +func chtGetHeader(ctx context.Context, bc *core.BlockChain, lc *light.LightChain, numbers []uint64) []byte { + if len(numbers) != 1 { + return nil + } + number := numbers[0] var header *types.Header if bc != nil { header = bc.GetHeaderByNumber(number) @@ -56,6 +60,26 @@ func chtGetHeader(ctx context.Context, bc *core.BlockChain, lc *light.LightChain return rlp } +func TestChtGetHeadersLes1(t *testing.T) { testCht(t, 1, MaxHelperTrieProofsFetch, chtGetHeaders) } + +func TestChtGetHeadersLes2(t *testing.T) { testCht(t, 2, MaxHelperTrieProofsFetch, chtGetHeaders) } + +func chtGetHeaders(ctx context.Context, bc *core.BlockChain, lc *light.LightChain, numbers []uint64) []byte { + var headers []*types.Header + if bc != nil { + for _, number := range numbers { + headers = append(headers, bc.GetHeaderByNumber(number)) + } + } else { + headers, _ = lc.GetHeadersByNumberOdr(ctx, numbers) + } + if headers == nil { + return nil + } + rlp, _ := rlp.EncodeToBytes(headers) + return rlp +} + type odrTestFn func(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte func TestOdrGetBlockLes1(t *testing.T) { testOdr(t, 1, 1, odrGetBlock) } @@ -180,7 +204,7 @@ func odrContractCall(ctx context.Context, db ethdb.Database, config *params.Chai } // testCht tests cht requests whose validation guaranteed by calculated cht root. -func testCht(t *testing.T, protocol int, fn chtTestFn) { +func testCht(t *testing.T, protocol int, maxFetch int, fn chtTestFn) { // Assemble the test environment config := light.TestServerIndexerConfig waitIndexers := func(cIndexer, bIndexer, btIndexer *core.ChainIndexer) { @@ -188,23 +212,20 @@ func testCht(t *testing.T, protocol int, fn chtTestFn) { cs, _, _ := cIndexer.Sections() bs, _, _ := bIndexer.Sections() bts, _, _ := btIndexer.Sections() - if cs >= 8 && bs >= 8 && bts >= 1 { + if cs >= config.PairChtSize/config.ChtSize && bs >= config.PairChtSize/config.BloomSize && + bts >= config.PairChtSize/config.BloomTrieSize { break } time.Sleep(10 * time.Millisecond) } } - server, client, tearDown := newClientServerEnv(t, int(config.ChtSize*8+config.ChtConfirm), protocol, waitIndexers, false) - defer func() { - if tearDown != nil { - tearDown() - } - }() + server, client, tearDown := newClientServerEnv(t, int(config.PairChtSize+config.ChtConfirm), protocol, waitIndexers, false) + defer tearDown() // Add trusted checkpoint for client side indexers. cs, _, head := server.chtIndexer.Sections() - light.StoreChtRoot(client.db, cs/8-1, head, light.GetChtRoot(server.db, cs-1, head)) - client.chtIndexer.AddKnownSectionHead(cs/8-1, head) + light.StoreChtRoot(client.db, cs*config.ChtSize/config.PairChtSize-1, head, light.GetChtRoot(server.db, cs-1, head)) + client.chtIndexer.AddKnownSectionHead(cs*config.ChtSize/config.PairChtSize-1, head) bts, _, head := server.bloomTrieIndexer.Sections() light.StoreBloomTrieRoot(client.db, bts-1, head, light.GetBloomTrieRoot(server.db, bts-1, head)) client.bloomTrieIndexer.AddKnownSectionHead(bts-1, head) @@ -220,18 +241,23 @@ func testCht(t *testing.T, protocol int, fn chtTestFn) { } server.rPeer, client.rPeer = peer, lPeer - test := func() { - for i := uint64(0); i <= config.ChtSize*8-1; i++ { - h1 := fn(light.NoOdr, server.pm.blockchain.(*core.BlockChain), nil, i) - ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) - h2 := fn(ctx, nil, client.pm.blockchain.(*light.LightChain), i) - if !bytes.Equal(h1, h2) { - t.Error("cht mismatch") - } - cancel() + i := uint64(0) + for { + var numbers []uint64 + for ; i <= config.PairChtSize-1 && len(numbers) < maxFetch; i += 1 { + numbers = append(numbers, i) } + if len(numbers) == 0 { + break + } + h1 := fn(light.NoOdr, server.pm.blockchain.(*core.BlockChain), nil, numbers) + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + h2 := fn(ctx, nil, client.pm.blockchain.(*light.LightChain), numbers) + if !bytes.Equal(h1, h2) { + t.Error("cht mismatch") + } + cancel() } - test() } // testOdr tests odr requests whose validation guaranteed by block headers. diff --git a/light/lightchain.go b/light/lightchain.go index 5cd161fd69..0a021fc3ef 100644 --- a/light/lightchain.go +++ b/light/lightchain.go @@ -462,6 +462,13 @@ func (self *LightChain) GetHeaderByNumberOdr(ctx context.Context, number uint64) return GetHeaderByNumber(ctx, self.odr, number) } +// GetHeadersByNumberOdr retrieves a batch of block headers from the database or network +// by number, caching it (associated with its hash) if found. +func (self *LightChain) GetHeadersByNumberOdr(ctx context.Context, numbers []uint64) ([]*types.Header, error) { + // TODO(rjl4935456442) Caching the read block headers. + return GetHeadersByNumber(ctx, self.odr, numbers) +} + // Config retrieves the header chain's chain configuration. func (self *LightChain) Config() *params.ChainConfig { return self.hc.Config() } diff --git a/light/odr_util.go b/light/odr_util.go index 659bc4c8d1..463be1a9d8 100644 --- a/light/odr_util.go +++ b/light/odr_util.go @@ -70,6 +70,63 @@ func GetHeaderByNumber(ctx context.Context, odr OdrBackend, number uint64) (*typ return r.Headers[0], nil } +func GetHeadersByNumber(ctx context.Context, odr OdrBackend, numbers []uint64) ([]*types.Header, error) { + var ( + db = odr.Database() + headers = make([]*types.Header, len(numbers)) + missing []int + ) + + for i, number := range numbers { + hash := rawdb.ReadCanonicalHash(db, number) + if (hash != common.Hash{}) { + // if there is a canonical hash, there is a header too + header := rawdb.ReadHeader(db, hash, number) + if header == nil { + panic("Canonical hash present but header not found") + } + headers[i] = header + } else { + missing = append(missing, i) + } + } + + var ( + chtCount, sectionHeadNum uint64 + sectionHead common.Hash + ) + if odr.ChtIndexer() != nil { + chtCount, sectionHeadNum, sectionHead = odr.ChtIndexer().Sections() + canonicalHash := rawdb.ReadCanonicalHash(db, sectionHeadNum) + // if the CHT was injected as a trusted checkpoint, we have no canonical hash yet so we accept zero hash too + for chtCount > 0 && canonicalHash != sectionHead && canonicalHash != (common.Hash{}) { + chtCount-- + if chtCount > 0 { + sectionHeadNum = chtCount*odr.IndexerConfig().ChtSize - 1 + sectionHead = odr.ChtIndexer().SectionHead(chtCount - 1) + canonicalHash = rawdb.ReadCanonicalHash(db, sectionHeadNum) + } + } + } + + reqs := make([]uint64, 0, len(missing)) + for _, index := range missing { + if numbers[index] >= chtCount*odr.IndexerConfig().ChtSize { + return nil, ErrNoTrustedCht + } + reqs = append(reqs, numbers[index]) + } + r := &ChtRequest{ChtRoot: GetChtRoot(db, chtCount-1, sectionHead), ChtNum: chtCount - 1, Numbers: reqs} + if err := odr.Retrieve(ctx, r); err != nil { + return nil, err + } + // Assemble the final result + for i, index := range missing { + headers[index] = r.Headers[i] + } + return headers, nil +} + func GetCanonicalHash(ctx context.Context, odr OdrBackend, number uint64) (common.Hash, error) { hash := rawdb.ReadCanonicalHash(odr.Database(), number) if (hash != common.Hash{}) {