diff --git a/p2p/protocol.go b/p2p/protocol.go index 3f52205f59..50ea3a968d 100644 --- a/p2p/protocol.go +++ b/p2p/protocol.go @@ -3,8 +3,6 @@ package p2p import ( "bytes" "time" - - "github.com/ethereum/go-ethereum/ethutil" ) // Protocol represents a P2P subprotocol implementation. @@ -89,20 +87,26 @@ type baseProtocol struct { func runBaseProtocol(peer *Peer, rw MsgReadWriter) error { bp := &baseProtocol{rw, peer} - if err := bp.doHandshake(rw); err != nil { + errc := make(chan error, 1) + go func() { errc <- rw.WriteMsg(bp.handshakeMsg()) }() + if err := bp.readHandshake(); err != nil { return err } + // handle write error + if err := <-errc; err != nil { + return err + } + // run main loop - quit := make(chan error, 1) go func() { for { if err := bp.handle(rw); err != nil { - quit <- err + errc <- err break } } }() - return bp.loop(quit) + return bp.loop(errc) } var pingTimeout = 2 * time.Second @@ -174,7 +178,7 @@ func (bp *baseProtocol) handle(rw MsgReadWriter) error { // // TODO: add event mechanism to notify baseProtocol for new peers if len(peers) > 0 { - return bp.rw.EncodeMsg(peersMsg, peers) + return bp.rw.EncodeMsg(peersMsg, peers...) } case peersMsg: @@ -193,14 +197,9 @@ func (bp *baseProtocol) handle(rw MsgReadWriter) error { return nil } -func (bp *baseProtocol) doHandshake(rw MsgReadWriter) error { - // send our handshake - if err := rw.WriteMsg(bp.handshakeMsg()); err != nil { - return err - } - +func (bp *baseProtocol) readHandshake() error { // read and handle remote handshake - msg, err := rw.ReadMsg() + msg, err := bp.rw.ReadMsg() if err != nil { return err } @@ -271,9 +270,9 @@ func (bp *baseProtocol) handshakeMsg() Msg { ) } -func (bp *baseProtocol) peerList() []ethutil.RlpEncodable { +func (bp *baseProtocol) peerList() []interface{} { peers := bp.peer.otherPeers() - ds := make([]ethutil.RlpEncodable, 0, len(peers)) + ds := make([]interface{}, 0, len(peers)) for _, p := range peers { p.infolock.Lock() addr := p.listenAddr diff --git a/p2p/protocol_test.go b/p2p/protocol_test.go index 65f26fb12d..d8a6c666b6 100644 --- a/p2p/protocol_test.go +++ b/p2p/protocol_test.go @@ -2,6 +2,8 @@ package p2p import ( "fmt" + "net" + "reflect" "testing" ) @@ -43,6 +45,52 @@ func TestBaseProtocolDisconnect(t *testing.T) { <-done } +func TestBaseProtocolPeers(t *testing.T) { + id1 := NewSimpleClientIdentity("p1", "", "", "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") + id2 := NewSimpleClientIdentity("p2", "", "", "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb") + cannedPeerList := []*peerAddr{ + {IP: net.ParseIP("1.2.3.4"), Port: 2222, Pubkey: []byte{}}, + {IP: net.ParseIP("5.6.7.8"), Port: 3333, Pubkey: []byte{}}, + } + rw1, rw2 := MsgPipe() + + // run matcher, close pipe when addresses have arrived + addrChan := make(chan *peerAddr, len(cannedPeerList)) + go func() { + for _, want := range cannedPeerList { + got := <-addrChan + t.Logf("got peer: %+v", got) + if !reflect.DeepEqual(want, got) { + t.Errorf("mismatch: got %#v, want %#v", got, want) + } + } + rw1.Close() + }() + + // run first peer + peer1 := NewPeer(id2, nil) + peer1.ourID = id1 + peer1.pubkeyHook = func(*peerAddr) error { return nil } + peer1.otherPeers = func() []*Peer { + pl := make([]*Peer, len(cannedPeerList)) + for i, addr := range cannedPeerList { + pl[i] = &Peer{listenAddr: addr} + } + return pl + } + go runBaseProtocol(peer1, rw2) + + // run second peer + peer2 := NewPeer(id1, nil) + peer2.ourID = id2 + peer2.pubkeyHook = func(*peerAddr) error { return nil } + peer2.otherPeers = func() []*Peer { return nil } + peer2.newPeerAddr = addrChan // feed peer suggestions into matcher + if err := runBaseProtocol(peer2, rw1); err != ErrPipeClosed { + t.Errorf("peer2 terminated with unexpected error: %v", err) + } +} + func expectMsg(r MsgReader, code uint64) error { msg, err := r.ReadMsg() if err != nil {