diff --git a/p2p/discover/table.go b/p2p/discover/table.go index e6dafb0dca..564fe86172 100644 --- a/p2p/discover/table.go +++ b/p2p/discover/table.go @@ -459,6 +459,25 @@ func (tab *Table) findnodeByID(target enode.ID, nresults int, preferLive bool) * return nodes } +func (tab *Table) appendBucketNodes(dist uint, result []*enode.Node) []*enode.Node { + if dist > 256 { + return result + } + if dist == 0 { + return append(result, tab.self()) + } + + tab.mutex.Lock() + defer tab.mutex.Unlock() + for _, n := range tab.bucketAtDistance(int(dist)).entries { + if n.livenessChecks > 1 { + node := n.Node // avoid handing out pointer to struct field + result = append(result, &node) + } + } + return result +} + // len returns the number of nodes in the table. func (tab *Table) len() (n int) { tab.mutex.Lock() diff --git a/p2p/discover/v5_udp.go b/p2p/discover/v5_udp.go index 6ba7a90618..4b96a1dec7 100644 --- a/p2p/discover/v5_udp.go +++ b/p2p/discover/v5_udp.go @@ -852,6 +852,7 @@ func (t *UDPv5) handleFindnode(p *v5wire.Findnode, fromID enode.ID, fromAddr *ne // collectTableNodes creates a FINDNODE result set for the given distances. func (t *UDPv5) collectTableNodes(rip net.IP, distances []uint, limit int) []*enode.Node { var nodes []*enode.Node + var bn []*enode.Node var processed = make(map[uint]struct{}) for _, dist := range distances { // Reject duplicate / invalid distances. @@ -860,20 +861,10 @@ func (t *UDPv5) collectTableNodes(rip net.IP, distances []uint, limit int) []*en continue } - // Get the nodes. - var bn []*enode.Node - if dist == 0 { - bn = []*enode.Node{t.Self()} - } else if dist <= 256 { - t.tab.mutex.Lock() - bn = unwrapNodes(t.tab.bucketAtDistance(int(dist)).entries) - t.tab.mutex.Unlock() - } - processed[dist] = struct{}{} - - // Apply some pre-checks to avoid sending invalid nodes. + bn = t.tab.appendBucketNodes(dist, limit, bn[:0]) for _, n := range bn { - // TODO livenessChecks > 1 + // Apply some pre-checks to avoid sending invalid nodes. + // Note liveness is checked by appendBucketNodes. if netutil.CheckRelayIP(rip, n.IP()) != nil { continue }