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:
zelig 2015-01-07 18:08:10 +00:00
parent 34c7c6b767
commit 211ee7e4f3
4 changed files with 59 additions and 51 deletions

View file

@ -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.

View file

@ -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)

View file

@ -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

View file

@ -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()