Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions eth/peer.go
Original file line number Diff line number Diff line change
Expand Up @@ -534,9 +534,9 @@ func (ps *peerSet) Register(p *peer) error {
if existPeer.pairRw != nil {
return errAlreadyRegistered
}
existPeer.PairPeer = p.Peer
existPeer.SetPairPeer(p.Peer)
existPeer.pairRw = p.rw
p.PairPeer = existPeer.Peer
p.SetPairPeer(existPeer.Peer)
return p2p.ErrAddPairPeer
}
ps.peers[p.id] = p
Expand Down
4 changes: 2 additions & 2 deletions p2p/dial.go
Original file line number Diff line number Diff line change
Expand Up @@ -266,8 +266,8 @@ func (s *dialstate) checkDial(n *discover.Node, peers map[discover.NodeID]*Peer)
case dialing:
return errAlreadyDialing
case peers[n.ID] != nil:
exitsPeer := peers[n.ID]
if exitsPeer.PairPeer != nil {
existPeer := peers[n.ID]
if existPeer.PairPeer() != nil {
return errAlreadyConnected
}
case s.ntab != nil && n.ID == s.ntab.Self().ID:
Expand Down
34 changes: 30 additions & 4 deletions p2p/peer.go
Original file line number Diff line number Diff line change
Expand Up @@ -113,8 +113,10 @@ type Peer struct {
disc chan DiscReason

// events receives message send / receive events if set
events *event.Feed
PairPeer *Peer
events *event.Feed

pairPeerMu sync.RWMutex
pairPeer *Peer
}

// NewPeer returns a peer for testing purposes.
Expand Down Expand Up @@ -190,6 +192,30 @@ func (p *Peer) Log() log.Logger {
return p.log
}

func (p *Peer) PairPeer() *Peer {
p.pairPeerMu.RLock()
defer p.pairPeerMu.RUnlock()

return p.pairPeer
}

func (p *Peer) SetPairPeer(pair *Peer) {
p.pairPeerMu.Lock()
p.pairPeer = pair
p.pairPeerMu.Unlock()
}

func (p *Peer) ClearPairPeer(pair *Peer) bool {
p.pairPeerMu.Lock()
defer p.pairPeerMu.Unlock()

if p.pairPeer != pair {
return false
}
p.pairPeer = nil
return true
}

func (p *Peer) run() (remoteRequested bool, err error) {
var (
writeStart = make(chan struct{}, 1)
Expand Down Expand Up @@ -235,8 +261,8 @@ loop:
close(p.closed)
p.rw.close(reason)
p.wg.Wait()
if p.PairPeer != nil {
go func() { p.PairPeer.Disconnect(DiscPairPeerStop) }()
if pairPeer := p.PairPeer(); pairPeer != nil {
go pairPeer.Disconnect(DiscPairPeerStop)
}
return remoteRequested, err
}
Expand Down
28 changes: 28 additions & 0 deletions p2p/peer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -197,6 +197,34 @@ func TestPeerDisconnectRace(t *testing.T) {
}
}

func TestPeerRunDisconnectsPairPeer(t *testing.T) {
closer, _, peer, errc := testPeer(nil)
defer closer()

pairPeer := &Peer{
disc: make(chan DiscReason, 1),
closed: make(chan struct{}),
}
peer.SetPairPeer(pairPeer)

closer()

select {
case <-errc:
case <-time.After(2 * time.Second):
t.Fatal("peer did not stop")
}

select {
case reason := <-pairPeer.disc:
if reason != DiscPairPeerStop {
t.Fatalf("unexpected pair disconnect reason: got %v want %v", reason, DiscPairPeerStop)
}
case <-time.After(2 * time.Second):
t.Fatal("pair peer was not disconnected")
}
}

func TestNewPeer(t *testing.T) {
name := "nodename"
caps := []Cap{{"foo", 2}, {"bar", 3}}
Expand Down
7 changes: 3 additions & 4 deletions p2p/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -714,7 +714,7 @@ running:

go srv.runPeer(p)
if peers[c.id] != nil {
peers[c.id].PairPeer = p
peers[c.id].SetPairPeer(p)
srv.log.Debug("Adding p2p pair peer", "name", name, "addr", c.fd.RemoteAddr(), "connections", connCount)
} else {
peers[c.id] = p
Expand Down Expand Up @@ -777,8 +777,7 @@ func removePeerTracking(peers map[discover.NodeID]*Peer, pd peerDrop, connCount
}
if current := peers[pd.ID()]; current == pd.Peer {
delete(peers, pd.ID())
} else if current != nil && current.PairPeer == pd.Peer {
current.PairPeer = nil
} else if current != nil && current.ClearPairPeer(pd.Peer) {
}
return connCount
}
Expand All @@ -801,7 +800,7 @@ func (srv *Server) encHandshakeChecks(peers map[discover.NodeID]*Peer, inboundCo
return DiscTooManyPeers
case peers[c.id] != nil:
exitPeer := peers[c.id]
if exitPeer.PairPeer != nil {
if exitPeer.PairPeer() != nil {
return DiscAlreadyConnected
}
return nil
Expand Down
4 changes: 2 additions & 2 deletions p2p/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -490,7 +490,7 @@ func TestRemovePeerTrackingKeepsPrimaryOnPairDrop(t *testing.T) {
id := randomID()
primary := newPeer(&conn{id: id}, nil)
pair := newPeer(&conn{id: id}, nil)
primary.PairPeer = pair
primary.SetPairPeer(pair)

peers := map[discover.NodeID]*Peer{id: primary}
connCount := removePeerTracking(peers, peerDrop{Peer: pair}, 2)
Expand All @@ -501,7 +501,7 @@ func TestRemovePeerTrackingKeepsPrimaryOnPairDrop(t *testing.T) {
if peers[id] != primary {
t.Fatal("primary peer was removed while dropping pair peer")
}
if primary.PairPeer != nil {
if primary.PairPeer() != nil {
t.Fatal("primary peer still references dropped pair peer")
}
}
Expand Down
Loading