mirror of
https://github.com/ethereum/go-ethereum.git
synced 2026-07-22 12:46:44 +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{})
|
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.
|
// ActiveAddresses returns addresses of all connected peers.
|
||||||
// Actually if the peer selector keeps historical info , then once active peers
|
// Actually if the peer selector keeps historical info , then once active peers
|
||||||
// will be included too.
|
// will be included too.
|
||||||
|
|
|
||||||
|
|
@ -30,7 +30,9 @@ var discard = Protocol{
|
||||||
|
|
||||||
func testPeer(protos []Protocol) (net.Conn, *Peer, <-chan error) {
|
func testPeer(protos []Protocol) (net.Conn, *Peer, <-chan error) {
|
||||||
conn1, conn2 := net.Pipe()
|
conn1, conn2 := net.Pipe()
|
||||||
peer := newPeer(conn1, protos, nil)
|
peer := &Peer{}
|
||||||
|
peer.init(conn1)
|
||||||
|
peer.protocols = protos
|
||||||
peer.ourID = &peerId{}
|
peer.ourID = &peerId{}
|
||||||
peer.pubkeyHook = func(*peerAddr) error { return nil }
|
peer.pubkeyHook = func(*peerAddr) error { return nil }
|
||||||
errc := make(chan error, 1)
|
errc := make(chan error, 1)
|
||||||
|
|
|
||||||
|
|
@ -77,6 +77,8 @@ type Server struct {
|
||||||
peerSlots chan int
|
peerSlots chan int
|
||||||
peerCount int
|
peerCount int
|
||||||
|
|
||||||
|
connectFunc func(*Peer, net.Conn)
|
||||||
|
|
||||||
quit chan struct{}
|
quit chan struct{}
|
||||||
wg sync.WaitGroup
|
wg sync.WaitGroup
|
||||||
peerDisconnect chan *Peer
|
peerDisconnect chan *Peer
|
||||||
|
|
@ -92,8 +94,6 @@ type NAT interface {
|
||||||
String() string
|
String() string
|
||||||
}
|
}
|
||||||
|
|
||||||
type peerFunc func(srv *Server, c net.Conn, dialAddr *peerAddr) *Peer
|
|
||||||
|
|
||||||
// Peers returns all currently connected peers.
|
// Peers returns all currently connected peers.
|
||||||
func (srv *Server) Peers() (peers []*Peer) {
|
func (srv *Server) Peers() (peers []*Peer) {
|
||||||
srv.lock.RLock()
|
srv.lock.RLock()
|
||||||
|
|
@ -145,6 +145,8 @@ func (srv *Server) SuggestPeer(addr string, pubkey []byte) error {
|
||||||
// to decide if it is a worthwhile connection
|
// to decide if it is a worthwhile connection
|
||||||
func (srv *Server) AddPeer(addr *peerAddr) (err error) {
|
func (srv *Server) AddPeer(addr *peerAddr) (err error) {
|
||||||
// need to look up nodeID first
|
// need to look up nodeID first
|
||||||
|
srvlog.Infof("checking peer %v", addr)
|
||||||
|
|
||||||
peer := &Peer{
|
peer := &Peer{
|
||||||
dialAddr: addr,
|
dialAddr: addr,
|
||||||
lastActiveC: make(chan time.Time),
|
lastActiveC: make(chan time.Time),
|
||||||
|
|
@ -173,15 +175,30 @@ func (srv *Server) dialPeer(peer *Peer) (err error) {
|
||||||
srv.peerSlots <- slot
|
srv.peerSlots <- slot
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
srvlog.Debugf("Connected to %v (slot %d)\n", peer.dialAddr, slot)
|
||||||
peer.slot = slot
|
peer.slot = slot
|
||||||
peer.connect(srv, conn)
|
srv.connectFunc(peer, conn)
|
||||||
srv.wg.Add(1)
|
|
||||||
go srv.addPeer(peer)
|
go srv.addPeer(peer)
|
||||||
}
|
}
|
||||||
return
|
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.
|
// Broadcast sends an RLP-encoded message to all connected peers.
|
||||||
// This method is deprecated and will be removed later.
|
// This method is deprecated and will be removed later.
|
||||||
func (srv *Server) Broadcast(protocol string, code uint64, data ...interface{}) {
|
func (srv *Server) Broadcast(protocol string, code uint64, data ...interface{}) {
|
||||||
|
|
@ -251,6 +268,10 @@ func (srv *Server) Start() (err error) {
|
||||||
srv.PeerSelector = &BaseSelector{}
|
srv.PeerSelector = &BaseSelector{}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if srv.connectFunc == nil {
|
||||||
|
srv.connectFunc = srv.connect
|
||||||
|
}
|
||||||
|
|
||||||
// make all slots available
|
// make all slots available
|
||||||
for i := range srv.peers {
|
for i := range srv.peers {
|
||||||
srv.peerSlots <- i
|
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)
|
srvlog.Debugf("Accepted conn %v (slot %d) - peer selector check after handshake", conn.RemoteAddr(), slot)
|
||||||
peer := &Peer{slot: slot}
|
peer := &Peer{slot: slot}
|
||||||
peer.connect(srv, conn)
|
srv.connectFunc(peer, conn)
|
||||||
srv.addPeer(peer)
|
srv.addPeer(peer)
|
||||||
case <-srv.quit:
|
case <-srv.quit:
|
||||||
return
|
return
|
||||||
|
|
|
||||||
|
|
@ -9,12 +9,12 @@ import (
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
func startTestServer(t *testing.T, pf peerFunc) *Server {
|
func startTestServer(t *testing.T, cb func(*Peer, net.Conn)) *Server {
|
||||||
server := &Server{
|
server := &Server{
|
||||||
Identity: &peerId{},
|
Identity: &peerId{},
|
||||||
MaxPeers: 10,
|
MaxPeers: 10,
|
||||||
ListenAddr: "127.0.0.1:0",
|
ListenAddr: "127.0.0.1:0",
|
||||||
newPeerFunc: pf,
|
connectFunc: cb,
|
||||||
}
|
}
|
||||||
if err := server.Start(); err != nil {
|
if err := server.Start(); err != nil {
|
||||||
t.Fatalf("Could not start server: %v", err)
|
t.Fatalf("Could not start server: %v", err)
|
||||||
|
|
@ -27,16 +27,9 @@ func TestServerListen(t *testing.T) {
|
||||||
|
|
||||||
// start the test server
|
// start the test server
|
||||||
connected := make(chan *Peer)
|
connected := make(chan *Peer)
|
||||||
srv := startTestServer(t, func(srv *Server, conn net.Conn, dialAddr *peerAddr) *Peer {
|
srv := startTestServer(t, func(peer *Peer, conn net.Conn) {
|
||||||
if conn == nil {
|
peer.init(conn)
|
||||||
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)
|
|
||||||
connected <- peer
|
connected <- peer
|
||||||
return peer
|
|
||||||
})
|
})
|
||||||
defer close(connected)
|
defer close(connected)
|
||||||
defer srv.Stop()
|
defer srv.Stop()
|
||||||
|
|
@ -50,10 +43,18 @@ func TestServerListen(t *testing.T) {
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case peer := <-connected:
|
case peer := <-connected:
|
||||||
if peer.conn.LocalAddr().String() != conn.RemoteAddr().String() {
|
if peer.conn == nil {
|
||||||
t.Errorf("peer started with wrong conn: got %v, want %v",
|
t.Error("peer setup with nil conn")
|
||||||
peer.conn.LocalAddr(), conn.RemoteAddr())
|
} 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):
|
case <-time.After(1 * time.Second):
|
||||||
t.Error("server did not accept within one second")
|
t.Error("server did not accept within one second")
|
||||||
}
|
}
|
||||||
|
|
@ -72,7 +73,7 @@ func TestServerDial(t *testing.T) {
|
||||||
go func() {
|
go func() {
|
||||||
conn, err := listener.Accept()
|
conn, err := listener.Accept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Error("acccept error:", err)
|
t.Error("accept error:", err)
|
||||||
}
|
}
|
||||||
conn.Close()
|
conn.Close()
|
||||||
accepted <- conn
|
accepted <- conn
|
||||||
|
|
@ -80,28 +81,28 @@ func TestServerDial(t *testing.T) {
|
||||||
|
|
||||||
// start the test server
|
// start the test server
|
||||||
connected := make(chan *Peer)
|
connected := make(chan *Peer)
|
||||||
srv := startTestServer(t, func(srv *Server, conn net.Conn, dialAddr *peerAddr) *Peer {
|
srv := startTestServer(t, func(peer *Peer, conn net.Conn) {
|
||||||
if conn == nil {
|
peer.init(conn)
|
||||||
t.Error("peer func called with nil conn")
|
go func() { connected <- peer }()
|
||||||
}
|
|
||||||
peer := newPeer(conn, nil, dialAddr)
|
|
||||||
connected <- peer
|
|
||||||
return peer
|
|
||||||
})
|
})
|
||||||
defer close(connected)
|
defer close(connected)
|
||||||
defer srv.Stop()
|
defer srv.Stop()
|
||||||
|
|
||||||
// tell the server to connect.
|
// tell the server to connect.
|
||||||
connAddr := newPeerAddr(listener.Addr(), nil)
|
connAddr := newPeerAddr(listener.Addr(), nil)
|
||||||
srv.peerConnect <- connAddr
|
srv.AddPeer(connAddr)
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case conn := <-accepted:
|
case conn := <-accepted:
|
||||||
select {
|
select {
|
||||||
case peer := <-connected:
|
case peer := <-connected:
|
||||||
if peer.conn.RemoteAddr().String() != conn.LocalAddr().String() {
|
if conn == nil {
|
||||||
t.Errorf("peer started with wrong conn: got %v, want %v",
|
t.Error("peer func called with nil conn")
|
||||||
peer.conn.RemoteAddr(), conn.LocalAddr())
|
} 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 {
|
if peer.dialAddr != connAddr {
|
||||||
t.Errorf("peer started with wrong dialAddr: got %v, want %v",
|
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) {
|
func TestServerBroadcast(t *testing.T) {
|
||||||
defer testlog(t).detach()
|
defer testlog(t).detach()
|
||||||
var connected sync.WaitGroup
|
var connected sync.WaitGroup
|
||||||
srv := startTestServer(t, func(srv *Server, c net.Conn, dialAddr *peerAddr) *Peer {
|
srv := startTestServer(t, func(peer *Peer, conn net.Conn) {
|
||||||
peer := newPeer(c, []Protocol{discard}, dialAddr)
|
peer.init(conn)
|
||||||
|
peer.protocols = []Protocol{discard}
|
||||||
peer.startSubprotocols([]Cap{discard.cap()})
|
peer.startSubprotocols([]Cap{discard.cap()})
|
||||||
connected.Done()
|
connected.Done()
|
||||||
return peer
|
|
||||||
})
|
})
|
||||||
defer srv.Stop()
|
defer srv.Stop()
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue