diff --git a/p2p/discover/common.go b/p2p/discover/common.go index 193c6c12cc..06bcc9f9a7 100644 --- a/p2p/discover/common.go +++ b/p2p/discover/common.go @@ -21,6 +21,7 @@ import ( "math/rand" "net" "sync" + "sync/atomic" "github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/p2p/enode" @@ -70,22 +71,20 @@ type lookupWalker struct { closeCh chan struct{} wg sync.WaitGroup - lookup lookupFunc - lookupDone chan struct{} - foundNode chan *enode.Node - iters map[*lookupIterator]struct{} + lookup lookupFunc + lookupDone chan struct{} + liveItersVal atomic.Value // []*lookupIterator } func newLookupWalker(fn lookupFunc) *lookupWalker { w := &lookupWalker{ lookup: fn, - iters: make(map[*lookupIterator]struct{}), newIterCh: make(chan *lookupIterator), delIterCh: make(chan *lookupIterator), triggerCh: make(chan struct{}), closeCh: make(chan struct{}), - foundNode: make(chan *enode.Node), } + w.setLiveIters(nil) w.wg.Add(1) go w.loop() return w @@ -96,34 +95,39 @@ func (w *lookupWalker) close() { w.wg.Wait() } +// 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{}) + for { + if lookupDone == nil && anyIterNeedsNodes(iters) { + lookupDone = make(chan struct{}) + go w.runLookup(lookupDone) + } + select { case it := <-w.newIterCh: - w.iters[it] = struct{}{} - w.startNewLookup() + iters[it] = struct{}{} + w.setLiveIters(iters) case it := <-w.delIterCh: - delete(w.iters, it) + delete(iters, it) + w.setLiveIters(iters) case <-w.triggerCh: - w.startNewLookup() - case <-w.lookupDone: - w.lookupDone = nil - w.startNewLookup() - - case n := <-w.foundNode: - for it := range w.iters { - it.deliver(n) - } + case <-lookupDone: + lookupDone = nil case <-w.closeCh: - for it := range w.iters { + w.setLiveIters(nil) + for it := range iters { it.close() } - if w.lookupDone != nil { - <-w.lookupDone + if lookupDone != nil { + <-lookupDone } w.wg.Done() return @@ -131,44 +135,62 @@ func (w *lookupWalker) loop() { } } -func (w *lookupWalker) startNewLookup() { - if w.lookupDone != nil { - return // already running - } - for it := range w.iters { +func anyIterNeedsNodes(iters map[*lookupIterator]struct{}) bool { + for it := range iters { if it.needsNodes() { - w.lookupDone = make(chan struct{}) - go w.runLookup() - return + return true } } - // all iterators have full buffer + return false } -func (w *lookupWalker) runLookup() { +func (w *lookupWalker) runLookup(done chan struct{}) { w.lookup(func(n *enode.Node) { - select { - case w.foundNode <- n: - case <-w.closeCh: + for _, it := range w.liveIters() { + it.deliver(n) } }) - w.lookupDone <- struct{}{} + 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 - w *lookupWalker - mu sync.Mutex - cond *sync.Cond - buf []*enode.Node + cur *enode.Node + walker *lookupWalker + filter filterFunc + mu sync.Mutex + cond *sync.Cond + buf []*enode.Node } const lookupIteratorBuffer = 100 -func (w *lookupWalker) newIterator() *lookupIterator { - it := &lookupIterator{w: w, buf: make([]*enode.Node, 0, lookupIteratorBuffer)} +type filterFunc func(*enode.Node) bool + +func (w *lookupWalker) newIterator(filter filterFunc) *lookupIterator { + if filter == nil { + filter = func(*enode.Node) bool { return true } + } + it := &lookupIterator{ + walker: w, + filter: filter, + 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: @@ -179,8 +201,8 @@ func (w *lookupWalker) newIterator() *lookupIterator { func (it *lookupIterator) Next() bool { select { - case it.w.triggerCh <- struct{}{}: - case <-it.w.closeCh: + case it.walker.triggerCh <- struct{}{}: + case <-it.walker.closeCh: } it.cur = nil @@ -205,8 +227,8 @@ func (it *lookupIterator) Node() *enode.Node { func (it *lookupIterator) Close() { select { - case it.w.delIterCh <- it: - case <-it.w.closeCh: + case it.walker.delIterCh <- it: + case <-it.walker.closeCh: } it.close() } @@ -221,11 +243,15 @@ func (it *lookupIterator) close() { } } -// deliver sends a node to the iterator buffer. +// deliver places a node into the iterator buffer. func (it *lookupIterator) deliver(n *enode.Node) { it.mu.Lock() defer it.mu.Unlock() + if it.buf == nil || !it.filter(n) { + return + } + // Place in buffer, overwriting a random entry when at capacity. if len(it.buf) == lookupIteratorBuffer { it.buf[rand.Intn(len(it.buf))] = n } else { @@ -234,6 +260,7 @@ func (it *lookupIterator) deliver(n *enode.Node) { it.cond.Signal() } +// needsNodes reports whether the iterator is low on nodes. func (it *lookupIterator) needsNodes() bool { it.mu.Lock() defer it.mu.Unlock() diff --git a/p2p/discover/common_test.go b/p2p/discover/common_test.go index 28df182755..3b21903c40 100644 --- a/p2p/discover/common_test.go +++ b/p2p/discover/common_test.go @@ -64,7 +64,7 @@ func TestLookupIterator(t *testing.T) { for i := 0; i < 10; i++ { wg.Add(1) - go testIterator(test.newIterator()) + go testIterator(test.newIterator(nil)) } test.serveOneLookup(testNodes[:10]) @@ -78,7 +78,7 @@ func TestLookupIterator(t *testing.T) { func TestLookupIteratorClose(t *testing.T) { test := newLookupWalkerTest() defer test.close() - it := test.newIterator() + it := test.newIterator(nil) go func() { time.Sleep(200 * time.Millisecond) @@ -96,7 +96,7 @@ func TestLookupIteratorDropStale(t *testing.T) { lookupDone = make(chan struct{}) ) defer test.close() - it := test.newIterator() + it := test.newIterator(nil) go func() { test.serveOneLookup(testNodes) close(lookupDone) @@ -126,7 +126,7 @@ func TestLookupIteratorDropStale(t *testing.T) { func TestLookupIteratorDrained(t *testing.T) { var ( test = newLookupWalkerTest() - it = test.newIterator() + it = test.newIterator(nil) testNodes = makeTestNodes(2 * lookupIteratorBuffer) ) diff --git a/p2p/discover/v4_udp.go b/p2p/discover/v4_udp.go index e1561365a6..b5afd57187 100644 --- a/p2p/discover/v4_udp.go +++ b/p2p/discover/v4_udp.go @@ -305,8 +305,12 @@ func (t *UDPv4) Close() { } // RandomNodes is an iterator yielding nodes from a random walk of the DHT. -func (t *UDPv4) RandomNodes() discutil.Iterator { - return t.randomWalk.newIterator() +// +// All iterators share the same random walk to minimize network traffic. Discovered nodes +// are checked against the filter function and returned by the iterator only when the +// filter returns true. +func (t *UDPv4) RandomNodes(filter func(*enode.Node) bool) discutil.Iterator { + return t.randomWalk.newIterator(filter) } // LookupRandom finds random nodes in the network.