From 8cad8737b6e7ee10d34c30130b11308f35216a2f Mon Sep 17 00:00:00 2001 From: Felix Lange Date: Sat, 3 Aug 2019 19:09:31 +0200 Subject: [PATCH] p2p/discover: WIP --- p2p/discover/common.go | 172 +++++++++++++++++------------------- p2p/discover/common_test.go | 41 +-------- p2p/discover/v4_udp.go | 14 +-- 3 files changed, 93 insertions(+), 134 deletions(-) diff --git a/p2p/discover/common.go b/p2p/discover/common.go index 06bcc9f9a7..5ec323809b 100644 --- a/p2p/discover/common.go +++ b/p2p/discover/common.go @@ -18,10 +18,9 @@ package discover import ( "crypto/ecdsa" - "math/rand" + "fmt" "net" "sync" - "sync/atomic" "github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/p2p/enode" @@ -60,31 +59,27 @@ type ReadPacket struct { Addr *net.UDPAddr } -type lookupFunc func(func(*enode.Node)) +type lookupFunc func(cancel <-chan struct{}, seenNode func(*enode.Node)) // lookupWalker performs recursive lookups, walking the DHT. // It manages a set iterators which receive lookup results as they are found. type lookupWalker struct { - newIterCh chan *lookupIterator - delIterCh chan *lookupIterator - triggerCh chan struct{} - closeCh chan struct{} - wg sync.WaitGroup + lookup lookupFunc + closeCh chan struct{} - lookup lookupFunc - lookupDone chan struct{} - liveItersVal atomic.Value // []*lookupIterator + mu sync.Mutex + cond *sync.Cond + wg sync.WaitGroup + iters map[*lookupIterator]struct{} } func newLookupWalker(fn lookupFunc) *lookupWalker { w := &lookupWalker{ - lookup: fn, - newIterCh: make(chan *lookupIterator), - delIterCh: make(chan *lookupIterator), - triggerCh: make(chan struct{}), - closeCh: make(chan struct{}), + lookup: fn, + closeCh: make(chan struct{}), + iters: make(map[*lookupIterator]struct{}), } - w.setLiveIters(nil) + w.cond = sync.NewCond(&w.mu) w.wg.Add(1) go w.loop() return w @@ -98,41 +93,74 @@ func (w *lookupWalker) close() { // loop schedules lookups. It ensures a lookup is running while // any live iterator needs more nodes. func (w *lookupWalker) loop() { - var lookupDone chan struct{} - iters := make(map[*lookupIterator]struct{}) - + var ( + done = make(chan struct{}) + cancel = make(chan struct{}) + running bool + ) for { - if lookupDone == nil && anyIterNeedsNodes(iters) { - lookupDone = make(chan struct{}) - go w.runLookup(lookupDone) + if !running { + go w.runLookup(cancel, done) } - select { - case it := <-w.newIterCh: - iters[it] = struct{}{} - w.setLiveIters(iters) - - case it := <-w.delIterCh: - delete(iters, it) - w.setLiveIters(iters) - - case <-w.triggerCh: - - case <-lookupDone: - lookupDone = nil - + case <-done: case <-w.closeCh: - w.setLiveIters(nil) - for it := range iters { - it.close() + if running { + close(cancel) + <-done } - if lookupDone != nil { - <-lookupDone - } - w.wg.Done() - return + goto shutdown } } + +shutdown: + w.mu.Lock() + defer w.mu.Unlock() + for it := range w.iters { + it.close() + } + w.wg.Done() +} + +func (w *lookupWalker) runLookup(cancel, done chan struct{}) { + w.lookup(cancel, w.foundNode) + done <- struct{}{} +} + +func (w *lookupWalker) foundNode(n *enode.Node) { + w.mu.Lock() + defer w.mu.Unlock() + for it := range w.iters { + it.deliver(n) + fmt.Println("delivered", len(it.buf), it.needsNodes()) + } + for !anyIterNeedsNodes(w.iters) { + w.cond.Wait() + } +} + +func (w *lookupWalker) newIterator(filter filterFunc) *lookupIterator { + it := newLookupIterator(w, filter) + it.walker.add(it) + return it +} + +func (w *lookupWalker) add(it *lookupIterator) { + w.mu.Lock() + defer w.mu.Unlock() + w.iters[it] = struct{}{} + w.unblockLookup() +} + +func (w *lookupWalker) remove(it *lookupIterator) { + w.mu.Lock() + defer w.mu.Unlock() + delete(w.iters, it) + w.unblockLookup() +} + +func (w *lookupWalker) unblockLookup() { + w.cond.Signal() } func anyIterNeedsNodes(iters map[*lookupIterator]struct{}) bool { @@ -144,27 +172,6 @@ func anyIterNeedsNodes(iters map[*lookupIterator]struct{}) bool { return false } -func (w *lookupWalker) runLookup(done chan struct{}) { - w.lookup(func(n *enode.Node) { - for _, it := range w.liveIters() { - it.deliver(n) - } - }) - close(done) -} - -func (w *lookupWalker) setLiveIters(iters map[*lookupIterator]struct{}) { - s := make([]*lookupIterator, 0, len(iters)) - for it := range iters { - s = append(s, it) - } - w.liveItersVal.Store(s) -} - -func (w *lookupWalker) liveIters() []*lookupIterator { - return w.liveItersVal.Load().([]*lookupIterator) -} - // lookupIterator is a sequence of discovered nodes. type lookupIterator struct { cur *enode.Node @@ -179,7 +186,7 @@ const lookupIteratorBuffer = 100 type filterFunc func(*enode.Node) bool -func (w *lookupWalker) newIterator(filter filterFunc) *lookupIterator { +func newLookupIterator(w *lookupWalker, filter filterFunc) *lookupIterator { if filter == nil { filter = func(*enode.Node) bool { return true } } @@ -189,21 +196,10 @@ func (w *lookupWalker) newIterator(filter filterFunc) *lookupIterator { buf: make([]*enode.Node, 0, lookupIteratorBuffer), } it.cond = sync.NewCond(&it.mu) - - // Register the iterator with walker. - select { - case w.newIterCh <- it: - case <-w.closeCh: - it.buf = nil - } return it } func (it *lookupIterator) Next() bool { - select { - case it.walker.triggerCh <- struct{}{}: - case <-it.walker.closeCh: - } it.cur = nil // Wait for the buffer to be filled. @@ -218,6 +214,8 @@ func (it *lookupIterator) Next() bool { it.cur = it.buf[0] copy(it.buf, it.buf[1:]) it.buf = it.buf[:len(it.buf)-1] + fmt.Println("read node", len(it.buf)) + it.walker.unblockLookup() return true } @@ -226,10 +224,7 @@ func (it *lookupIterator) Node() *enode.Node { } func (it *lookupIterator) Close() { - select { - case it.walker.delIterCh <- it: - case <-it.walker.closeCh: - } + it.walker.remove(it) it.close() } @@ -244,20 +239,19 @@ func (it *lookupIterator) close() { } // deliver places a node into the iterator buffer. -func (it *lookupIterator) deliver(n *enode.Node) { +func (it *lookupIterator) deliver(n *enode.Node) bool { it.mu.Lock() defer it.mu.Unlock() if it.buf == nil || !it.filter(n) { - return + return true } - // Place in buffer, overwriting a random entry when at capacity. - if len(it.buf) == lookupIteratorBuffer { - it.buf[rand.Intn(len(it.buf))] = n - } else { - it.buf = append(it.buf, n) + if len(it.buf) == cap(it.buf) { + return false } + it.buf = append(it.buf, n) it.cond.Signal() + return true } // needsNodes reports whether the iterator is low on nodes. diff --git a/p2p/discover/common_test.go b/p2p/discover/common_test.go index 3b21903c40..e257d9938c 100644 --- a/p2p/discover/common_test.go +++ b/p2p/discover/common_test.go @@ -69,7 +69,7 @@ func TestLookupIterator(t *testing.T) { test.serveOneLookup(testNodes[:10]) test.serveOneLookup(testNodes[10:20]) - test.serveOneLookup(testNodes[20:]) + test.serveOneLookup(testNodes[20:40]) wg.Wait() test.close() @@ -87,41 +87,6 @@ func TestLookupIteratorClose(t *testing.T) { it.Next() } -// This test checks that the lookup iterator drops nodes when they're not being -// read fast enough. -func TestLookupIteratorDropStale(t *testing.T) { - var ( - test = newLookupWalkerTest() - testNodes = makeTestNodes(2 * lookupIteratorBuffer) - lookupDone = make(chan struct{}) - ) - defer test.close() - it := test.newIterator(nil) - go func() { - test.serveOneLookup(testNodes) - close(lookupDone) - }() - - // The first call to NextNode triggers the lookup and receives the first result - // as soon as it becomes available. - it.Next() - if it.Node() != testNodes[0] { - t.Fatalf("wrong result %d: got %v, want %v", 0, it.Node().ID(), testNodes[0]) - } - - // Now wait for the lookup to finish and read the remaining nodes. - <-lookupDone - for i := 0; i < lookupIteratorBuffer; i++ { - it.Next() - for _, tn := range testNodes[lookupIteratorBuffer:] { - if it.Node() == tn { - return - } - } - } - t.Fatal("didn't find any node from second half of testNodes") -} - // This test checks that the iterator kicks off a lookup when Next is called. func TestLookupIteratorDrained(t *testing.T) { var ( @@ -177,7 +142,7 @@ func (t *lookupWalkerTest) serveOneLookup(nodes []*enode.Node) { <-t.nodes } -func (t *lookupWalkerTest) lookupFunc(callback func(*enode.Node)) { +func (t *lookupWalkerTest) lookupFunc(cancel <-chan struct{}, callback func(*enode.Node)) { if atomic.AddInt32(&t.running, 1) != 1 { panic("spawned more than one instance of lookupFunc") } @@ -189,7 +154,7 @@ func (t *lookupWalkerTest) lookupFunc(callback func(*enode.Node)) { callback(n) } t.nodes <- nil - case <-t.lookupWalker.closeCh: + case <-cancel: return } } diff --git a/p2p/discover/v4_udp.go b/p2p/discover/v4_udp.go index b5afd57187..ecfc983881 100644 --- a/p2p/discover/v4_udp.go +++ b/p2p/discover/v4_udp.go @@ -314,7 +314,7 @@ func (t *UDPv4) RandomNodes(filter func(*enode.Node) bool) discutil.Iterator { } // LookupRandom finds random nodes in the network. -func (t *UDPv4) randomLookupWithCallback(callback func(*enode.Node)) { +func (t *UDPv4) randomLookupWithCallback(cancel <-chan struct{}, callback func(*enode.Node)) { if t.tab.len() == 0 { // All nodes were dropped, refresh. The very first query will hit this // case and run the bootstrapping logic. @@ -322,7 +322,7 @@ func (t *UDPv4) randomLookupWithCallback(callback func(*enode.Node)) { } var target encPubkey crand.Read(target[:]) - t.lookup(target, callback) + t.lookup(target, cancel, callback) } // LookupPubkey finds the closest nodes to the given public key. @@ -332,25 +332,25 @@ func (t *UDPv4) LookupPubkey(key *ecdsa.PublicKey) []*enode.Node { // case and run the bootstrapping logic. <-t.tab.refresh() } - return unwrapNodes(t.lookup(encodePubkey(key), nil)) + return unwrapNodes(t.lookup(encodePubkey(key), t.tab.closeReq, nil)) } // for Table func (t *UDPv4) lookupRandom() []*enode.Node { var target encPubkey crand.Read(target[:]) - return unwrapNodes(t.lookup(target, nil)) + return unwrapNodes(t.lookup(target, t.tab.closeReq, nil)) } // for Table func (t *UDPv4) lookupSelf() []*enode.Node { - return unwrapNodes(t.lookup(encodePubkey(&t.priv.PublicKey), nil)) + return unwrapNodes(t.lookup(encodePubkey(&t.priv.PublicKey), t.tab.closeReq, nil)) } // lookup performs a network search for nodes close to the given target. It approaches the // target by querying nodes that are closer to it on each iteration. The given target does // not need to be an actual node identifier. -func (t *UDPv4) lookup(targetKey encPubkey, nodeCallback func(*enode.Node)) []*node { +func (t *UDPv4) lookup(targetKey encPubkey, cancel <-chan struct{}, nodeCallback func(*enode.Node)) []*node { var ( target = enode.ID(crypto.Keccak256Hash(targetKey[:])) asked = make(map[enode.ID]bool) @@ -390,7 +390,7 @@ func (t *UDPv4) lookup(targetKey encPubkey, nodeCallback func(*enode.Node)) []*n result.push(n, bucketSize) } } - case <-t.tab.closeReq: + case <-cancel: return nil // shutdown, no need to continue. } pendingQueries--