mirror of
https://github.com/ethereum/go-ethereum.git
synced 2026-07-22 04:36:42 +00:00
tests pass
- connectFunc on server, by default it calles srv.connect(peer, conn) - srv.connect(peer, conn) is the init code from peer - this still provides a nice hook for testing - adapt server and peer tests
This commit is contained in:
parent
34c7c6b767
commit
211ee7e4f3
4 changed files with 59 additions and 51 deletions
16
p2p/peer.go
16
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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == 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,29 +81,29 @@ 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 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",
|
||||
peer.dialAddr, connAddr)
|
||||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue