p2p: fix RemovePeer to disconnect the peer again

Also make RemovePeer synchronous and add a test.
This commit is contained in:
Felix Lange 2020-02-09 16:04:33 +01:00
parent af30b4d982
commit 92e1a62dd9
2 changed files with 112 additions and 42 deletions

View file

@ -300,42 +300,55 @@ func (srv *Server) LocalNode() *enode.LocalNode {
// Peers returns all connected peers. // Peers returns all connected peers.
func (srv *Server) Peers() []*Peer { func (srv *Server) Peers() []*Peer {
var ps []*Peer var ps []*Peer
select { srv.doPeerOp(func(peers map[enode.ID]*Peer) {
// Note: We'd love to put this function into a variable but
// that seems to cause a weird compiler error in some
// environments.
case srv.peerOp <- func(peers map[enode.ID]*Peer) {
for _, p := range peers { for _, p := range peers {
ps = append(ps, p) ps = append(ps, p)
} }
}: })
<-srv.peerOpDone
case <-srv.quit:
}
return ps return ps
} }
// PeerCount returns the number of connected peers. // PeerCount returns the number of connected peers.
func (srv *Server) PeerCount() int { func (srv *Server) PeerCount() int {
var count int var count int
select { srv.doPeerOp(func(ps map[enode.ID]*Peer) {
case srv.peerOp <- func(ps map[enode.ID]*Peer) { count = len(ps) }: count = len(ps)
<-srv.peerOpDone })
case <-srv.quit:
}
return count return count
} }
// AddPeer connects to the given node and maintains the connection until the // AddPeer adds the given node to the static node set. When there is room in the peer set,
// server is shut down. If the connection fails for any reason, the server will // the server will connect to the node. If the connection fails for any reason, the server
// attempt to reconnect the peer. // will attempt to reconnect the peer.
func (srv *Server) AddPeer(node *enode.Node) { func (srv *Server) AddPeer(node *enode.Node) {
srv.dialsched.addStatic(node) srv.dialsched.addStatic(node)
} }
// RemovePeer disconnects from the given node // RemovePeer removes a node from the static node set. It also disconnects from the given
// node if it is currently connected as a peer.
func (srv *Server) RemovePeer(node *enode.Node) { func (srv *Server) RemovePeer(node *enode.Node) {
srv.dialsched.removeStatic(node) var (
ch chan *PeerEvent
sub event.Subscription
)
// Disconnect the peer on the main loop.
srv.doPeerOp(func(peers map[enode.ID]*Peer) {
srv.dialsched.removeStatic(node)
if peer := peers[node.ID()]; peer != nil {
ch = make(chan *PeerEvent, 1)
sub = srv.peerFeed.Subscribe(ch)
peer.Disconnect(DiscRequested)
}
})
// Wait for the peer connection to end.
if ch != nil {
defer sub.Unsubscribe()
for ev := range ch {
if ev.Peer == node.ID() && ev.Type == PeerEventTypeDrop {
return
}
}
}
} }
// AddTrustedPeer adds the given node to a reserved whitelist which allows the // AddTrustedPeer adds the given node to a reserved whitelist which allows the
@ -661,6 +674,16 @@ func (srv *Server) setupListening() error {
return nil return nil
} }
// doPeerOp runs fn on the main loop.
func (srv *Server) doPeerOp(fn peerOpFunc) {
select {
case srv.peerOp <- fn:
<-srv.peerOpDone
case <-srv.quit:
}
}
// run is the main loop of the server.
func (srv *Server) run() { func (srv *Server) run() {
srv.log.Info("Started P2P networking", "self", srv.localnode.Node().URLv4()) srv.log.Info("Started P2P networking", "self", srv.localnode.Node().URLv4())
defer srv.loopWG.Done() defer srv.loopWG.Done()
@ -999,14 +1022,10 @@ func (srv *Server) launchPeer(c *conn) *Peer {
} }
// runPeer runs in its own goroutine for each peer. // runPeer runs in its own goroutine for each peer.
// it waits until the Peer logic returns and removes
// the peer.
func (srv *Server) runPeer(p *Peer) { func (srv *Server) runPeer(p *Peer) {
if srv.newPeerHook != nil { if srv.newPeerHook != nil {
srv.newPeerHook(p) srv.newPeerHook(p)
} }
// broadcast peer add
srv.peerFeed.Send(&PeerEvent{ srv.peerFeed.Send(&PeerEvent{
Type: PeerEventTypeAdd, Type: PeerEventTypeAdd,
Peer: p.ID(), Peer: p.ID(),
@ -1014,10 +1033,18 @@ func (srv *Server) runPeer(p *Peer) {
LocalAddress: p.LocalAddr().String(), LocalAddress: p.LocalAddr().String(),
}) })
// run the protocol // Run the per-peer main loop.
remoteRequested, err := p.run() remoteRequested, err := p.run()
// broadcast peer drop // Announce disconnect on the main loop to update the peer set.
// The main loop waits for existing peers to be sent on srv.delpeer
// before returning, so this send should not select on srv.quit.
srv.delpeer <- peerDrop{p, err, remoteRequested}
// Broadcast peer drop to external subscribers. This needs to be
// after the send to delpeer so subscribers have a consistent view of
// the peer set (i.e. Server.Peers() doesn't include the peer when the
// event is received.
srv.peerFeed.Send(&PeerEvent{ srv.peerFeed.Send(&PeerEvent{
Type: PeerEventTypeDrop, Type: PeerEventTypeDrop,
Peer: p.ID(), Peer: p.ID(),
@ -1025,10 +1052,6 @@ func (srv *Server) runPeer(p *Peer) {
RemoteAddress: p.RemoteAddr().String(), RemoteAddress: p.RemoteAddr().String(),
LocalAddress: p.LocalAddr().String(), LocalAddress: p.LocalAddr().String(),
}) })
// Note: run waits for existing peers to be sent on srv.delpeer
// before returning, so this send should not select on srv.quit.
srv.delpeer <- peerDrop{p, err, remoteRequested}
} }
// NodeInfo represents a short summary of the information known about the host. // NodeInfo represents a short summary of the information known about the host.

View file

@ -34,10 +34,6 @@ import (
"golang.org/x/crypto/sha3" "golang.org/x/crypto/sha3"
) )
// func init() {
// log.Root().SetHandler(log.LvlFilterHandler(log.LvlTrace, log.StreamHandler(os.Stderr, log.TerminalFormat(false))))
// }
type testTransport struct { type testTransport struct {
rpub *ecdsa.PublicKey rpub *ecdsa.PublicKey
*rlpx *rlpx
@ -72,11 +68,12 @@ func (c *testTransport) close(err error) {
func startTestServer(t *testing.T, remoteKey *ecdsa.PublicKey, pf func(*Peer)) *Server { func startTestServer(t *testing.T, remoteKey *ecdsa.PublicKey, pf func(*Peer)) *Server {
config := Config{ config := Config{
Name: "test", Name: "test",
MaxPeers: 10, MaxPeers: 10,
ListenAddr: "127.0.0.1:0", ListenAddr: "127.0.0.1:0",
PrivateKey: newkey(), NoDiscovery: true,
Logger: testlog.Logger(t, log.LvlTrace), PrivateKey: newkey(),
Logger: testlog.Logger(t, log.LvlTrace),
} }
server := &Server{ server := &Server{
Config: config, Config: config,
@ -204,9 +201,38 @@ func TestServerDial(t *testing.T) {
} }
} }
// This test checks that connections are disconnected // This test checks that RemovePeer disconnects the peer if it is connected.
// just after the encryption handshake when the server is func TestServerRemovePeerDisconnect(t *testing.T) {
// at capacity. Trusted connections should still be accepted. srv1 := &Server{Config: Config{
PrivateKey: newkey(),
MaxPeers: 1,
NoDiscovery: true,
Logger: testlog.Logger(t, log.LvlTrace).New("server", "1"),
}}
srv2 := &Server{Config: Config{
PrivateKey: newkey(),
MaxPeers: 1,
NoDiscovery: true,
NoDial: true,
ListenAddr: "127.0.0.1:0",
Logger: testlog.Logger(t, log.LvlTrace).New("server", "2"),
}}
srv1.Start()
defer srv1.Stop()
srv2.Start()
defer srv2.Stop()
if !syncAddPeer(srv1, srv2.Self()) {
t.Fatal("peer not connected")
}
srv1.RemovePeer(srv2.Self())
if srv1.PeerCount() > 0 {
t.Fatal("removed peer still connected")
}
}
// This test checks that connections are disconnected just after the encryption handshake
// when the server is at capacity. Trusted connections should still be accepted.
func TestServerAtCap(t *testing.T) { func TestServerAtCap(t *testing.T) {
trustedNode := newkey() trustedNode := newkey()
trustedID := enode.PubkeyToIDV4(&trustedNode.PublicKey) trustedID := enode.PubkeyToIDV4(&trustedNode.PublicKey)
@ -217,6 +243,7 @@ func TestServerAtCap(t *testing.T) {
NoDial: true, NoDial: true,
NoDiscovery: true, NoDiscovery: true,
TrustedNodes: []*enode.Node{newNode(trustedID, "")}, TrustedNodes: []*enode.Node{newNode(trustedID, "")},
Logger: testlog.Logger(t, log.LvlTrace),
}, },
} }
if err := srv.Start(); err != nil { if err := srv.Start(); err != nil {
@ -292,9 +319,9 @@ func TestServerPeerLimits(t *testing.T) {
NoDial: true, NoDial: true,
NoDiscovery: true, NoDiscovery: true,
Protocols: []Protocol{discard}, Protocols: []Protocol{discard},
Logger: testlog.Logger(t, log.LvlTrace),
}, },
newTransport: func(fd net.Conn) transport { return tp }, newTransport: func(fd net.Conn) transport { return tp },
log: log.New(),
} }
if err := srv.Start(); err != nil { if err := srv.Start(); err != nil {
t.Fatalf("couldn't start server: %v", err) t.Fatalf("couldn't start server: %v", err)
@ -577,3 +604,23 @@ func (l *fakeAddrListener) Accept() (net.Conn, error) {
func (c *fakeAddrConn) RemoteAddr() net.Addr { func (c *fakeAddrConn) RemoteAddr() net.Addr {
return c.remoteAddr return c.remoteAddr
} }
func syncAddPeer(srv *Server, node *enode.Node) bool {
var (
ch = make(chan *PeerEvent)
sub = srv.SubscribeEvents(ch)
timeout = time.After(2 * time.Second)
)
defer sub.Unsubscribe()
srv.AddPeer(node)
for {
select {
case ev := <-ch:
if ev.Type == PeerEventTypeAdd && ev.Peer == node.ID() {
return true
}
case <-timeout:
return false
}
}
}