From 8aec0164e1263d95865943a963f7673d19350df8 Mon Sep 17 00:00:00 2001 From: Daniel Liu <139250065@qq.com> Date: Tue, 21 Apr 2026 14:16:44 +0800 Subject: [PATCH] fix(p2p,eth): synchronize pair peer teardown Guard PairPeer access with synchronized helpers and snapshot the pair before launching async disconnect during peer shutdown. Add regression coverage for pair-peer teardown and keep the existing pair tracking checks using the synchronized accessors. --- eth/peer.go | 4 ++-- p2p/dial.go | 4 ++-- p2p/peer.go | 34 ++++++++++++++++++++++++++++++---- p2p/peer_test.go | 28 ++++++++++++++++++++++++++++ p2p/server.go | 7 +++---- p2p/server_test.go | 4 ++-- 6 files changed, 67 insertions(+), 14 deletions(-) diff --git a/eth/peer.go b/eth/peer.go index 458ae56fe406..b864f8cec244 100644 --- a/eth/peer.go +++ b/eth/peer.go @@ -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 diff --git a/p2p/dial.go b/p2p/dial.go index d98241c84524..62320af5e241 100644 --- a/p2p/dial.go +++ b/p2p/dial.go @@ -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: diff --git a/p2p/peer.go b/p2p/peer.go index 79eb400c387c..b8b257f56dbc 100644 --- a/p2p/peer.go +++ b/p2p/peer.go @@ -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. @@ -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) @@ -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 } diff --git a/p2p/peer_test.go b/p2p/peer_test.go index 20d020f186d8..09907bbb341c 100644 --- a/p2p/peer_test.go +++ b/p2p/peer_test.go @@ -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}} diff --git a/p2p/server.go b/p2p/server.go index 467147623f2a..72df20dade2d 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -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 @@ -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 } @@ -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 diff --git a/p2p/server_test.go b/p2p/server_test.go index 49c6d3130f9b..92933d2f43f1 100644 --- a/p2p/server_test.go +++ b/p2p/server_test.go @@ -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) @@ -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") } }