diff --git a/p2p/discover/common.go b/p2p/discover/common.go index cbd62f867d..193c6c12cc 100644 --- a/p2p/discover/common.go +++ b/p2p/discover/common.go @@ -17,8 +17,8 @@ package discover import ( - "context" "crypto/ecdsa" + "math/rand" "net" "sync" @@ -64,22 +64,27 @@ type lookupFunc func(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 { - lookup lookupFunc - newIterCh chan *lookupIterator delIterCh chan *lookupIterator triggerCh chan struct{} closeCh chan struct{} wg sync.WaitGroup + + lookup lookupFunc + lookupDone chan struct{} + foundNode chan *enode.Node + iters map[*lookupIterator]struct{} } 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.wg.Add(1) go w.loop() @@ -92,38 +97,33 @@ func (w *lookupWalker) close() { } func (w *lookupWalker) loop() { - var ( - iters = make(map[*lookupIterator]struct{}) - foundNode = make(chan *enode.Node) - lookupDone = make(chan struct{}, 1) - trigger = w.triggerCh - ) for { select { case it := <-w.newIterCh: - iters[it] = struct{}{} + w.iters[it] = struct{}{} + w.startNewLookup() case it := <-w.delIterCh: - delete(iters, it) + delete(w.iters, it) - case <-trigger: - trigger = nil // stop listening to trigger until lookupDone - go w.runLookup(foundNode, lookupDone) + case <-w.triggerCh: + w.startNewLookup() - case <-lookupDone: - trigger = w.triggerCh + case <-w.lookupDone: + w.lookupDone = nil + w.startNewLookup() - case n := <-foundNode: - for it := range iters { + case n := <-w.foundNode: + for it := range w.iters { it.deliver(n) } case <-w.closeCh: - for it := range iters { - it.drainAndClose() + for it := range w.iters { + it.close() } - if trigger == nil { - <-lookupDone + if w.lookupDone != nil { + <-w.lookupDone } w.wg.Done() return @@ -131,85 +131,112 @@ func (w *lookupWalker) loop() { } } -func (w *lookupWalker) runLookup(nodes chan<- *enode.Node, done chan<- struct{}) { +func (w *lookupWalker) startNewLookup() { + if w.lookupDone != nil { + return // already running + } + for it := range w.iters { + if it.needsNodes() { + w.lookupDone = make(chan struct{}) + go w.runLookup() + return + } + } + // all iterators have full buffer +} + +func (w *lookupWalker) runLookup() { w.lookup(func(n *enode.Node) { select { - case nodes <- n: + case w.foundNode <- n: case <-w.closeCh: } }) - done <- struct{}{} + w.lookupDone <- struct{}{} } // lookupIterator is a sequence of discovered nodes. type lookupIterator struct { - w *lookupWalker - buf chan *enode.Node - closed bool - closeOnce sync.Once + cur *enode.Node + w *lookupWalker + mu sync.Mutex + cond *sync.Cond + buf []*enode.Node } const lookupIteratorBuffer = 100 func (w *lookupWalker) newIterator() *lookupIterator { - it := &lookupIterator{w: w, buf: make(chan *enode.Node, lookupIteratorBuffer)} + it := &lookupIterator{w: w, buf: make([]*enode.Node, 0, lookupIteratorBuffer)} + it.cond = sync.NewCond(&it.mu) select { case w.newIterCh <- it: case <-w.closeCh: - it.closed = true - close(it.buf) + it.buf = nil } return it } -func (it *lookupIterator) NextNode(ctx context.Context) (n *enode.Node, isLive bool) { - for { - select { - case it.w.triggerCh <- struct{}{}: - // lookup triggered - case n, ok := <-it.buf: - if !ok { - it.closed = true - } - return n, !it.closed - case <-ctx.Done(): - return nil, !it.closed - } +func (it *lookupIterator) Next() bool { + select { + case it.w.triggerCh <- struct{}{}: + case <-it.w.closeCh: } + it.cur = nil + + // Wait for the buffer to be filled. + it.mu.Lock() + defer it.mu.Unlock() + for it.buf != nil && len(it.buf) == 0 { + it.cond.Wait() + } + if it.buf == nil { + return false // closed + } + it.cur = it.buf[0] + copy(it.buf, it.buf[1:]) + it.buf = it.buf[:len(it.buf)-1] + return true +} + +func (it *lookupIterator) Node() *enode.Node { + return it.cur } func (it *lookupIterator) Close() { - it.closeOnce.Do(func() { - select { - case it.w.delIterCh <- it: - case <-it.w.closeCh: - } - it.drainAndClose() - }) + select { + case it.w.delIterCh <- it: + case <-it.w.closeCh: + } + it.close() +} + +func (it *lookupIterator) close() { + it.mu.Lock() + defer it.mu.Unlock() + + if it.buf != nil { + it.buf = nil + it.cond.Signal() + } } // deliver sends a node to the iterator buffer. func (it *lookupIterator) deliver(n *enode.Node) { - // We don't want deliver to block and replace stale results when they're not being - // read. Check whether the buffer is full and allow one receive from the buffer if so. - // This is OK because there is only one writer. - var remove chan *enode.Node - if len(it.buf) == cap(it.buf) { - remove = it.buf - } - for { - select { - case it.buf <- n: - return - case <-remove: - remove = nil - } + it.mu.Lock() + defer it.mu.Unlock() + + if len(it.buf) == lookupIteratorBuffer { + it.buf[rand.Intn(len(it.buf))] = n + } else { + it.buf = append(it.buf, n) } + it.cond.Signal() } -func (it *lookupIterator) drainAndClose() { - for len(it.buf) > 0 { - <-it.buf - } - close(it.buf) +func (it *lookupIterator) needsNodes() bool { + it.mu.Lock() + defer it.mu.Unlock() + + return len(it.buf) < lookupIteratorBuffer/3 } diff --git a/p2p/discover/common_test.go b/p2p/discover/common_test.go index 3c55256d11..28df182755 100644 --- a/p2p/discover/common_test.go +++ b/p2p/discover/common_test.go @@ -17,7 +17,6 @@ package discover import ( - "context" "encoding/binary" "sync" "sync/atomic" @@ -33,32 +32,33 @@ import ( func TestLookupIterator(t *testing.T) { var ( test = newLookupWalkerTest() - testNodes = make([]*enode.Node, 100) + testNodes = makeTestNodes(lookupIteratorBuffer) wg sync.WaitGroup ) - for i := range testNodes { - testNodes[i] = testNode(i) - } + testIterator := func(it discutil.Iterator) { defer wg.Done() - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() - nodes := discutil.ReadNodes(ctx, it, 20) - sortByID(nodes) // ReadNodes may shuffle results + // Check reading nodes: + nodes := discutil.ReadNodes(it, 20) + sortByID(nodes) if err := checkNodesEqual(nodes, testNodes[:20]); err != nil { t.Error(err) } + nodes = discutil.ReadNodes(it, 20) + sortByID(nodes) + if err := checkNodesEqual(nodes, testNodes[20:40]); err != nil { + t.Error(err) + } + // Check close: it.Close() - n, isLive := it.NextNode(context.Background()) - if n != nil { - t.Error("iterator returned non-nil node after close") + if it.Next() { + t.Error("Next returned true after close") } - if isLive { - t.Error("iterator returned isLive == true after close") + if it.Node() != nil { + t.Error("iterator has non-nil node after close") } - it.Close() // shouldn't crash } @@ -68,30 +68,35 @@ func TestLookupIterator(t *testing.T) { } test.serveOneLookup(testNodes[:10]) - test.serveOneLookup(testNodes[10:]) + test.serveOneLookup(testNodes[10:20]) + test.serveOneLookup(testNodes[20:]) wg.Wait() test.close() } +func TestLookupIteratorClose(t *testing.T) { + test := newLookupWalkerTest() + defer test.close() + it := test.newIterator() + + go func() { + time.Sleep(200 * time.Millisecond) + it.Close() + }() + 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 = make([]*enode.Node, 2*lookupIteratorBuffer) + testNodes = makeTestNodes(2 * lookupIteratorBuffer) lookupDone = make(chan struct{}) ) defer test.close() - - for i := range testNodes { - testNodes[i] = testNode(i) - } - - // Create iterator first so all found nodes go through its buffer. it := test.newIterator() - - // Serve one lookup. go func() { test.serveOneLookup(testNodes) close(lookupDone) @@ -99,17 +104,17 @@ func TestLookupIteratorDropStale(t *testing.T) { // The first call to NextNode triggers the lookup and receives the first result // as soon as it becomes available. - n, _ := it.NextNode(context.Background()) - if n != testNodes[0] { - t.Fatalf("wrong result %d: got %v, want %v", 0, n.ID(), testNodes[0]) + 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++ { - n, _ := it.NextNode(context.Background()) + it.Next() for _, tn := range testNodes[lookupIteratorBuffer:] { - if n == tn { + if it.Node() == tn { return } } @@ -117,10 +122,40 @@ func TestLookupIteratorDropStale(t *testing.T) { t.Fatal("didn't find any node from second half of testNodes") } -func testNode(id int) *enode.Node { - var nodeID enode.ID - binary.BigEndian.PutUint64(nodeID[:], uint64(id)) - return enode.SignNull(new(enr.Record), nodeID) +// This test checks that the iterator kicks off a lookup when Next is called. +func TestLookupIteratorDrained(t *testing.T) { + var ( + test = newLookupWalkerTest() + it = test.newIterator() + testNodes = makeTestNodes(2 * lookupIteratorBuffer) + ) + + test.serveOneLookup(testNodes[:lookupIteratorBuffer]) + nodes := discutil.ReadNodes(it, lookupIteratorBuffer) + sortByID(nodes) + if err := checkNodesEqual(nodes, testNodes[:lookupIteratorBuffer]); err != nil { + t.Fatal(err) + } + + // Here the iterator buffer is drained and no lookup is running. + + // Request more nodes. This needs to start another lookup. + go test.serveOneLookup(testNodes[lookupIteratorBuffer:]) + nodes = discutil.ReadNodes(it, 10) + sortByID(nodes) + if err := checkNodesEqual(nodes, testNodes[lookupIteratorBuffer:lookupIteratorBuffer+10]); err != nil { + t.Fatal(err) + } +} + +func makeTestNodes(n int) []*enode.Node { + nodes := make([]*enode.Node, n) + for i := range nodes { + var nodeID enode.ID + binary.BigEndian.PutUint64(nodeID[:], uint64(i)) + nodes[i] = enode.SignNull(new(enr.Record), nodeID) + } + return nodes } type lookupWalkerTest struct {