diff --git a/p2p/cademlia.go b/p2p/cademlia.go index d1567b74b0..a89023aa02 100644 --- a/p2p/cademlia.go +++ b/p2p/cademlia.go @@ -1,6 +1,7 @@ package p2p import ( + "fmt" "sync" "time" @@ -61,17 +62,18 @@ func (self *Cademlia) Stop() { self.quitC = nil } -func (self *Cademlia) AddPeer(peer peerInfo) (needed bool) { +func (self *Cademlia) AddPeer(peer peerInfo) (err error) { index := self.commonPrefixLength(peer.Hash()) row := self.rows[index] - needed = row.insert(&entry{peer: peer}) + needed := row.insert(&entry{peer: peer}) if needed { if index >= self.depth { go self.updateDepth() } cadlogger.Infof("accept peer %x...", peer.Hash()[:8]) } else { - cadlogger.Infof("reject peer %x... no worse peer found", peer.Hash()[:8]) + err = fmt.Errorf("no worse peer found") + cadlogger.Infof("reject peer %x..: %v", peer.Hash()[:8], err) } return } diff --git a/p2p/peer_selector.go b/p2p/peer_selector.go index dc4f953abc..2c1b2260f8 100644 --- a/p2p/peer_selector.go +++ b/p2p/peer_selector.go @@ -16,7 +16,7 @@ type peerInfo interface { } type peerSelector interface { - AddPeer(peer peerInfo) bool + AddPeer(peer peerInfo) error GetPeers(target ...[]byte) []peerInfo Start() error Stop() error @@ -28,8 +28,8 @@ type BaseSelector struct { peers []peerInfo } -func (self *BaseSelector) AddPeer(peer peerInfo) bool { - return true +func (self *BaseSelector) AddPeer(peer peerInfo) error { + return nil } func (self *BaseSelector) GetPeers(target ...[]byte) []peerInfo { diff --git a/p2p/server.go b/p2p/server.go index 35dca03c1d..d8285ff243 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -152,7 +152,7 @@ func (srv *Server) AddPeer(addr *peerAddr) (err error) { lastActiveC: make(chan time.Time), lastActive: time.Now().Add(-24 * time.Hour), } - if srv.PeerSelector.AddPeer(peer) { + if err = srv.PeerSelector.AddPeer(peer); err == nil { srvlog.Infof("peer %v accepted by peer selection", addr) err = srv.dialPeer(peer) } else { @@ -436,7 +436,6 @@ func (srv *Server) verifyPeer(peer *Peer) error { return newPeerError(errSelfConnection, "") } srv.lock.RLock() - defer srv.lock.RUnlock() for _, peer := range srv.peers { if peer != nil { id := peer.Identity() @@ -445,6 +444,12 @@ func (srv *Server) verifyPeer(peer *Peer) error { } } } + srv.lock.RUnlock() + if peer.dialAddr == nil { + if err := srv.PeerSelector.AddPeer(peer); err != nil { + return newPeerError(errRejectedPeer, "%v", err) + } + } return nil }