diff --git a/p2p/peer.go b/p2p/peer.go index facdaacb1c..a14a7f589e 100644 --- a/p2p/peer.go +++ b/p2p/peer.go @@ -116,22 +116,6 @@ func (self *Peer) init(conn net.Conn) { self.closed = make(chan struct{}) } -func (p *Peer) connect(server *Server, conn net.Conn) { - p.init(conn) - p.ourID = server.Identity - p.addPeer = server.AddPeer - p.getPeers = server.GetPeers - p.pubkeyHook = server.verifyPeer - p.runBaseProtocol = true - p.protocols = server.Protocols - - // laddr can be updated concurrently by NAT traversal. - // newServerPeer must be called with the server lock held. - if server.laddr != nil { - p.ourListenAddr = newPeerAddr(server.laddr, server.Identity.Pubkey()) - } -} - // ActiveAddresses returns addresses of all connected peers. // Actually if the peer selector keeps historical info , then once active peers // will be included too. diff --git a/p2p/peer_test.go b/p2p/peer_test.go index 4ee88f112b..8c04c137d7 100644 --- a/p2p/peer_test.go +++ b/p2p/peer_test.go @@ -30,7 +30,9 @@ var discard = Protocol{ func testPeer(protos []Protocol) (net.Conn, *Peer, <-chan error) { conn1, conn2 := net.Pipe() - peer := newPeer(conn1, protos, nil) + peer := &Peer{} + peer.init(conn1) + peer.protocols = protos peer.ourID = &peerId{} peer.pubkeyHook = func(*peerAddr) error { return nil } errc := make(chan error, 1) diff --git a/p2p/server.go b/p2p/server.go index ec09647011..a9d1d3989b 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -77,6 +77,8 @@ type Server struct { peerSlots chan int peerCount int + connectFunc func(*Peer, net.Conn) + quit chan struct{} wg sync.WaitGroup peerDisconnect chan *Peer @@ -92,8 +94,6 @@ type NAT interface { String() string } -type peerFunc func(srv *Server, c net.Conn, dialAddr *peerAddr) *Peer - // Peers returns all currently connected peers. func (srv *Server) Peers() (peers []*Peer) { srv.lock.RLock() @@ -145,6 +145,8 @@ func (srv *Server) SuggestPeer(addr string, pubkey []byte) error { // to decide if it is a worthwhile connection func (srv *Server) AddPeer(addr *peerAddr) (err error) { // need to look up nodeID first + srvlog.Infof("checking peer %v", addr) + peer := &Peer{ dialAddr: addr, lastActiveC: make(chan time.Time), @@ -173,15 +175,30 @@ func (srv *Server) dialPeer(peer *Peer) (err error) { srv.peerSlots <- slot return } + srvlog.Debugf("Connected to %v (slot %d)\n", peer.dialAddr, slot) peer.slot = slot - peer.connect(srv, conn) - srv.wg.Add(1) + srv.connectFunc(peer, conn) go srv.addPeer(peer) } return } +func (srv *Server) connect(p *Peer, conn net.Conn) { + p.init(conn) + p.ourID = srv.Identity + p.addPeer = srv.AddPeer + p.getPeers = srv.GetPeers + p.pubkeyHook = srv.verifyPeer + p.runBaseProtocol = true + p.protocols = srv.Protocols + + // laddr can be updated concurrently by NAT traversal. + if srv.laddr != nil { + p.ourListenAddr = newPeerAddr(srv.laddr, srv.Identity.Pubkey()) + } +} + // Broadcast sends an RLP-encoded message to all connected peers. // This method is deprecated and will be removed later. func (srv *Server) Broadcast(protocol string, code uint64, data ...interface{}) { @@ -251,6 +268,10 @@ func (srv *Server) Start() (err error) { srv.PeerSelector = &BaseSelector{} } + if srv.connectFunc == nil { + srv.connectFunc = srv.connect + } + // make all slots available for i := range srv.peers { srv.peerSlots <- i @@ -332,7 +353,7 @@ func (srv *Server) listenLoop() { } srvlog.Debugf("Accepted conn %v (slot %d) - peer selector check after handshake", conn.RemoteAddr(), slot) peer := &Peer{slot: slot} - peer.connect(srv, conn) + srv.connectFunc(peer, conn) srv.addPeer(peer) case <-srv.quit: return diff --git a/p2p/server_test.go b/p2p/server_test.go index ceb89e3f7f..2545aecfc8 100644 --- a/p2p/server_test.go +++ b/p2p/server_test.go @@ -9,12 +9,12 @@ import ( "time" ) -func startTestServer(t *testing.T, pf peerFunc) *Server { +func startTestServer(t *testing.T, cb func(*Peer, net.Conn)) *Server { server := &Server{ Identity: &peerId{}, MaxPeers: 10, ListenAddr: "127.0.0.1:0", - newPeerFunc: pf, + connectFunc: cb, } if err := server.Start(); err != nil { t.Fatalf("Could not start server: %v", err) @@ -27,16 +27,9 @@ func TestServerListen(t *testing.T) { // start the test server connected := make(chan *Peer) - srv := startTestServer(t, func(srv *Server, conn net.Conn, dialAddr *peerAddr) *Peer { - if conn == nil { - t.Error("peer func called with nil conn") - } - if dialAddr != nil { - t.Error("peer func called with non-nil dialAddr") - } - peer := newPeer(conn, nil, dialAddr) + srv := startTestServer(t, func(peer *Peer, conn net.Conn) { + peer.init(conn) connected <- peer - return peer }) defer close(connected) defer srv.Stop() @@ -50,10 +43,18 @@ func TestServerListen(t *testing.T) { select { case peer := <-connected: - if peer.conn.LocalAddr().String() != conn.RemoteAddr().String() { - t.Errorf("peer started with wrong conn: got %v, want %v", - peer.conn.LocalAddr(), conn.RemoteAddr()) + if peer.conn == nil { + t.Error("peer setup with nil conn") + } else { + if peer.conn.LocalAddr().String() != conn.RemoteAddr().String() { + t.Errorf("peer started with wrong conn: got %v, want %v", + peer.conn.LocalAddr(), conn.RemoteAddr()) + } } + if peer.dialAddr != nil { + t.Error("peer setup with non-nil dialAddr") + } + case <-time.After(1 * time.Second): t.Error("server did not accept within one second") } @@ -72,7 +73,7 @@ func TestServerDial(t *testing.T) { go func() { conn, err := listener.Accept() if err != nil { - t.Error("acccept error:", err) + t.Error("accept error:", err) } conn.Close() accepted <- conn @@ -80,28 +81,28 @@ func TestServerDial(t *testing.T) { // start the test server connected := make(chan *Peer) - srv := startTestServer(t, func(srv *Server, conn net.Conn, dialAddr *peerAddr) *Peer { - if conn == nil { - t.Error("peer func called with nil conn") - } - peer := newPeer(conn, nil, dialAddr) - connected <- peer - return peer + srv := startTestServer(t, func(peer *Peer, conn net.Conn) { + peer.init(conn) + go func() { connected <- peer }() }) defer close(connected) defer srv.Stop() // tell the server to connect. connAddr := newPeerAddr(listener.Addr(), nil) - srv.peerConnect <- connAddr + srv.AddPeer(connAddr) select { case conn := <-accepted: select { case peer := <-connected: - if peer.conn.RemoteAddr().String() != conn.LocalAddr().String() { - t.Errorf("peer started with wrong conn: got %v, want %v", - peer.conn.RemoteAddr(), conn.LocalAddr()) + if conn == nil { + t.Error("peer func called with nil conn") + } else { + if peer.conn.RemoteAddr().String() != conn.LocalAddr().String() { + t.Errorf("peer started with wrong conn: got %v, want %v", + peer.conn.RemoteAddr(), conn.LocalAddr()) + } } if peer.dialAddr != connAddr { t.Errorf("peer started with wrong dialAddr: got %v, want %v", @@ -119,11 +120,11 @@ func TestServerDial(t *testing.T) { func TestServerBroadcast(t *testing.T) { defer testlog(t).detach() var connected sync.WaitGroup - srv := startTestServer(t, func(srv *Server, c net.Conn, dialAddr *peerAddr) *Peer { - peer := newPeer(c, []Protocol{discard}, dialAddr) + srv := startTestServer(t, func(peer *Peer, conn net.Conn) { + peer.init(conn) + peer.protocols = []Protocol{discard} peer.startSubprotocols([]Cap{discard.cap()}) connected.Done() - return peer }) defer srv.Stop()