diff --git a/p2p/devp2p.go b/p2p/devp2p.go new file mode 100644 index 0000000000..9868ef0f4f --- /dev/null +++ b/p2p/devp2p.go @@ -0,0 +1,245 @@ +// Copyright 2015 The go-ethereum Authors +// This file is part of the go-ethereum library. +// +// The go-ethereum library is free software: you can redistribute it and/or modify +// it under the terms of the GNU Lesser General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// The go-ethereum library is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Lesser General Public License for more details. +// +// You should have received a copy of the GNU Lesser General Public License +// along with the go-ethereum library. If not, see . + +package p2p + +import ( + "crypto/ecdsa" + "errors" + "fmt" + "io" + "math" + "net" + "time" + + "github.com/ethereum/go-ethereum/p2p/discover" + "github.com/ethereum/go-ethereum/p2p/rlpx" + "github.com/ethereum/go-ethereum/rlp" +) + +const ( + // devp2p message codes + handshakeMsg = 0x00 + discMsg = 0x01 + pingMsg = 0x02 + pongMsg = 0x03 + getPeersMsg = 0x04 + peersMsg = 0x05 +) + +const ( + baseProtocolVersion = 4 + baseProtocolLength = uint64(16) + baseProtocolMaxMsgSize = 2 * 1024 +) + +var errMsgTooBig = errors.New("encoded message size exceeds uint32") + +// DiscReason indicates why a connection is being disconnected. +type DiscReason uint + +const ( + DiscRequested DiscReason = iota + DiscNetworkError + DiscProtocolError + DiscUselessPeer + DiscTooManyPeers + DiscAlreadyConnected + DiscIncompatibleVersion + DiscInvalidIdentity + DiscQuitting + DiscUnexpectedIdentity + DiscSelf + DiscReadTimeout + DiscSubprotocolError = 0x10 +) + +var discReasonToString = [...]string{ + DiscRequested: "Disconnect requested", + DiscNetworkError: "Network error", + DiscProtocolError: "Breach of protocol", + DiscUselessPeer: "Useless peer", + DiscTooManyPeers: "Too many peers", + DiscAlreadyConnected: "Already connected", + DiscIncompatibleVersion: "Incompatible P2P protocol version", + DiscInvalidIdentity: "Invalid node identity", + DiscQuitting: "Client quitting", + DiscUnexpectedIdentity: "Unexpected identity", + DiscSelf: "Connected to self", + DiscReadTimeout: "Read timeout", + DiscSubprotocolError: "Subprotocol error", +} + +func (d DiscReason) String() string { + if len(discReasonToString) < int(d) { + return fmt.Sprintf("Unknown Reason(%d)", d) + } + return discReasonToString[d] +} + +func (d DiscReason) Error() string { + return d.String() +} + +func discReasonForError(err error) DiscReason { + if reason, ok := err.(DiscReason); ok { + return reason + } + peerError, ok := err.(*peerError) + if ok { + switch peerError.code { + case errInvalidMsgCode, errInvalidMsg: + return DiscProtocolError + default: + return DiscSubprotocolError + } + } + return DiscSubprotocolError +} + +// protoHandshake is the RLP structure of the protocol handshake. +type protoHandshake struct { + Version uint64 + Name string + Caps []Cap + ListenPort uint64 + ID discover.NodeID +} + +// devConn implements the devp2p the messaging layer atop RLPx. +type devConn struct { + *rlpx.Conn + // contains negotiated protocol sessions. + // protocol zero is pre-negotiated and carries the + // built-in devp2p packets. + protocols []*devProtocol +} + +// devProtocol represents a running subprotocol. +type devProtocol struct { + p *rlpx.Protocol +} + +func newDevConn(fd net.Conn, key *ecdsa.PrivateKey, remote *ecdsa.PublicKey) *devConn { + c := new(devConn) + if remote == nil { + c.Conn = rlpx.Server(fd, &rlpx.Config{Key: key}) + } else { + c.Conn = rlpx.Client(fd, remote, &rlpx.Config{Key: key}) + } + c.protocols = []*devProtocol{{c.Conn.Protocol(0)}} + return c +} + +func (t *devConn) addProtocols(n int) { + for i := 0; i < n; i++ { + p := t.Conn.Protocol(uint16(len(t.protocols))) + t.protocols = append(t.protocols, &devProtocol{p}) + } +} + +// protoHandshake negotiates RLPx subprotocols. +// the protocol handshake is the first authenticated message +// and also verifies whether the RLPx encryption handshake 'worked' and the +// remote side actually provided the right public key. +func (t *devConn) doProtoHandshake(our *protoHandshake) (their *protoHandshake, err error) { + // Writing our handshake happens concurrently, we prefer + // returning the handshake read error. If the remote side + // disconnects us early with a valid reason, we should return it + // as the error so it can be tracked elsewhere. + werr := make(chan error, 1) + go func() { werr <- Send(t.protocols[0], handshakeMsg, our) }() + if their, err = readProtocolHandshake(t.protocols[0], our); err != nil { + <-werr // make sure the write terminates too + return nil, err + } + if err := <-werr; err != nil { + return nil, fmt.Errorf("write error: %v", err) + } + return their, nil +} + +func readProtocolHandshake(rw MsgReader, our *protoHandshake) (*protoHandshake, error) { + msg, err := rw.ReadMsg() + if err != nil { + return nil, err + } + if msg.Size > baseProtocolMaxMsgSize { + return nil, fmt.Errorf("message too big") + } + if msg.Code == discMsg { + // Disconnect before protocol handshake is valid according to the + // spec and we send it ourself if the posthanshake checks fail. + // We can't return the reason directly, though, because it is echoed + // back otherwise. Wrap it in a string instead. + var reason [1]DiscReason + rlp.Decode(msg.Payload, &reason) + return nil, reason[0] + } + if msg.Code != handshakeMsg { + return nil, fmt.Errorf("expected handshake, got %x", msg.Code) + } + var hs protoHandshake + if err := msg.Decode(&hs); err != nil { + return nil, err + } + // validate handshake info + if hs.Version != our.Version { + return nil, DiscIncompatibleVersion + } + if (hs.ID == discover.NodeID{}) { + return nil, DiscInvalidIdentity + } + return &hs, nil +} + +func (t *devConn) close(err error) { + // Tell the remote end why we're disconnecting if possible. + // TODO: if t.DidHandshake() + if r, ok := err.(DiscReason); ok && r != DiscNetworkError { + SendItems(t.protocols[0], discMsg, r) + } + t.Close() +} + +func (p *devProtocol) WriteMsg(msg Msg) error { + codelen, code, _ := rlp.EncodeToReader(msg.Code) + if msg.Size > math.MaxUint32-uint32(codelen) { + return errMsgTooBig + } + plen := msg.Size + uint32(codelen) + return p.p.SendPacket(plen, io.MultiReader(code, msg.Payload)) +} + +func (p *devProtocol) ReadMsg() (msg Msg, err error) { + len, r, err := p.p.ReadPacket() + if err != nil { + return msg, err + } + // Parse the message code, which is prepended to the protocol payload. + // r must be recognized as buffered by package rlp to prevent it from + // reading into the payload. The interface assertion ensures that it is. + // The input limit is 9, which is as large as an encoded uint64 can get. + s := rlp.NewStream(r.(rlp.ByteReader), 9) + if err := s.Decode(&msg.Code); err != nil { + return msg, err + } + // Remaining data in r belongs to the protocol. + msg.Payload = r + msg.Size = len - uint32(rlp.IntSize(msg.Code)) + msg.ReceivedAt = time.Now() + return msg, nil +} diff --git a/p2p/devp2p_test.go b/p2p/devp2p_test.go new file mode 100644 index 0000000000..8bf906c8f9 --- /dev/null +++ b/p2p/devp2p_test.go @@ -0,0 +1,125 @@ +// Copyright 2015 The go-ethereum Authors +// This file is part of the go-ethereum library. +// +// The go-ethereum library is free software: you can redistribute it and/or modify +// it under the terms of the GNU Lesser General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// The go-ethereum library is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Lesser General Public License for more details. +// +// You should have received a copy of the GNU Lesser General Public License +// along with the go-ethereum library. If not, see . + +package p2p + +import ( + "errors" + "net" + "reflect" + "sync" + "testing" + + "github.com/davecgh/go-spew/spew" + "github.com/ethereum/go-ethereum/p2p/discover" +) + +func TestProtocolHandshake(t *testing.T) { + var ( + prv0, prv1 = newkey(), newkey() + fd0, fd1 = net.Pipe() + hs0 = &protoHandshake{Version: 3, ID: discover.PubkeyID(&prv0.PublicKey), Caps: []Cap{{"a", 0}, {"b", 2}}} + hs1 = &protoHandshake{Version: 3, ID: discover.PubkeyID(&prv1.PublicKey), Caps: []Cap{{"c", 1}, {"d", 3}}} + wg sync.WaitGroup + ) + + wg.Add(2) + go func() { + defer wg.Done() + conn := newDevConn(fd0, prv0, nil) + + phs, err := conn.doProtoHandshake(hs0) + if err != nil { + t.Errorf("dial side proto handshake error: %v", err) + return + } + if !reflect.DeepEqual(phs, hs1) { + t.Errorf("dial side proto handshake mismatch:\ngot: %s\nwant: %s\n", spew.Sdump(phs), spew.Sdump(hs1)) + return + } + conn.close(DiscQuitting) + }() + go func() { + defer wg.Done() + conn := newDevConn(fd1, prv1, &prv0.PublicKey) + + phs, err := conn.doProtoHandshake(hs1) + if err != nil { + t.Errorf("listen side proto handshake error: %v", err) + return + } + if !reflect.DeepEqual(phs, hs0) { + t.Errorf("listen side proto handshake mismatch:\ngot: %s\nwant: %s\n", spew.Sdump(phs), spew.Sdump(hs0)) + return + } + + if err := ExpectMsg(conn.protocols[0], discMsg, []DiscReason{DiscQuitting}); err != nil { + t.Errorf("error receiving disconnect: %v", err) + } + }() + wg.Wait() +} + +func TestProtocolHandshakeErrors(t *testing.T) { + our := &protoHandshake{Version: 3, Caps: []Cap{{"foo", 2}, {"bar", 3}}, Name: "quux"} + id := randomID() + tests := []struct { + code uint64 + msg interface{} + err error + }{ + { + code: discMsg, + msg: []DiscReason{DiscQuitting}, + err: DiscQuitting, + }, + { + code: 0x989898, + msg: []byte{1}, + err: errors.New("expected handshake, got 989898"), + }, + { + code: handshakeMsg, + msg: make([]byte, baseProtocolMaxMsgSize+2), + err: errors.New("message too big"), + }, + { + code: handshakeMsg, + msg: []byte{1, 2, 3}, + err: newPeerError(errInvalidMsg, "(code 0) (size 4) rlp: expected input list for p2p.protoHandshake"), + }, + { + code: handshakeMsg, + msg: &protoHandshake{Version: 9944, ID: id}, + err: DiscIncompatibleVersion, + }, + { + code: handshakeMsg, + msg: &protoHandshake{Version: 3}, + err: DiscInvalidIdentity, + }, + } + + for i, test := range tests { + p1, p2 := MsgPipe() + go Send(p1, test.code, test.msg) + _, err := readProtocolHandshake(p2, our) + if !reflect.DeepEqual(err, test.err) { + t.Errorf("test %d: error mismatch: got %q, want %q", i, err, test.err) + } + p1.Close() + } +} diff --git a/p2p/dial.go b/p2p/dial.go index 0fd3a4cf52..32290a3d83 100644 --- a/p2p/dial.go +++ b/p2p/dial.go @@ -133,7 +133,7 @@ func (s *dialstate) newTasks(nRunning int, peers map[discover.NodeID]*Peer, now // Compute number of dynamic dials necessary at this point. needDynDials := s.maxDynDials for _, p := range peers { - if p.rw.is(dynDialedConn) { + if p.conn.is(dynDialedConn) { needDynDials-- } } @@ -205,6 +205,11 @@ func (s *dialstate) taskDone(t task, now time.Time) { } func (t *dialTask) Do(srv *Server) { + remotePubkey, err := t.dest.ID.Pubkey() + if err != nil { + glog.V(logger.Warn).Infof("aborted dialing (invalid pubkey) %v\n", t.dest) + return + } addr := &net.TCPAddr{IP: t.dest.IP, Port: int(t.dest.TCP)} glog.V(logger.Debug).Infof("dialing %v\n", t.dest) fd, err := srv.Dialer.Dial("tcp", addr.String()) @@ -213,9 +218,10 @@ func (t *dialTask) Do(srv *Server) { return } mfd := newMeteredConn(fd, false) - - srv.setupConn(mfd, t.flags, t.dest) + dc := newDevConn(mfd, srv.PrivateKey, remotePubkey) + srv.setupConn(dc, t.flags, t.dest) } + func (t *dialTask) String() string { return fmt.Sprintf("%v %x %v:%d", t.flags, t.dest.ID[:8], t.dest.IP, t.dest.TCP) } diff --git a/p2p/dial_test.go b/p2p/dial_test.go index d24e03e292..99bc016de3 100644 --- a/p2p/dial_test.go +++ b/p2p/dial_test.go @@ -28,6 +28,9 @@ import ( func init() { spew.Config.Indent = "\t" + spew.Config.DisableMethods = true + // glog.SetV(8) + // glog.SetToStderr(true) } type dialtest struct { @@ -49,7 +52,7 @@ func runDialTest(t *testing.T, test dialtest) { pm := func(ps []*Peer) map[discover.NodeID]*Peer { m := make(map[discover.NodeID]*Peer) for _, p := range ps { - m[p.rw.id] = p + m[p.conn.id] = p } return m } @@ -94,18 +97,18 @@ func TestDialStateDynDial(t *testing.T) { // A discovery query is launched. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, new: []task{&discoverTask{bootstrap: true}}, }, // Dynamic dials are launched when it completes. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, done: []task{ &discoverTask{bootstrap: true, results: []*discover.Node{ @@ -127,11 +130,11 @@ func TestDialStateDynDial(t *testing.T) { // the sum of active dial count and dynamic peer count is == maxDynDials. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(3)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(4)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(3)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(4)}}, }, done: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(3)}}, @@ -142,12 +145,12 @@ func TestDialStateDynDial(t *testing.T) { // maxDynDials has been reached. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(3)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(4)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(5)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(3)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(4)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(5)}}, }, done: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(5)}}, @@ -160,11 +163,11 @@ func TestDialStateDynDial(t *testing.T) { // results from last discovery lookup are reused. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(3)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(4)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(5)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(3)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(4)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(5)}}, }, new: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(6)}}, @@ -175,9 +178,9 @@ func TestDialStateDynDial(t *testing.T) { // and a new one is spawned because more candidates are needed. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(5)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(5)}}, }, done: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(6)}}, @@ -192,10 +195,10 @@ func TestDialStateDynDial(t *testing.T) { // no new is started. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(5)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(7)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(5)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(7)}}, }, done: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(7)}}, @@ -205,10 +208,10 @@ func TestDialStateDynDial(t *testing.T) { // should be immediately requested. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(0)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(5)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(7)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(0)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(5)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(7)}}, }, done: []task{ &discoverTask{}, @@ -259,8 +262,8 @@ func TestDialStateDynDialFromTable(t *testing.T) { // Dialing nodes 1,2 succeeds. Dials from the lookup are launched. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, done: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(1)}}, @@ -281,11 +284,11 @@ func TestDialStateDynDialFromTable(t *testing.T) { // Dialing nodes 3,4,5 fails. The dials from the lookup succeed. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(10)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(11)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(12)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(10)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(11)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(12)}}, }, done: []task{ &dialTask{dynDialedConn, &discover.Node{ID: uintID(3)}}, @@ -300,11 +303,11 @@ func TestDialStateDynDialFromTable(t *testing.T) { // discovery query is still running. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(10)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(11)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(12)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(10)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(11)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(12)}}, }, }, // Nodes 3,4 are not tried again because only the first two @@ -312,11 +315,11 @@ func TestDialStateDynDialFromTable(t *testing.T) { // already connected. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(10)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(11)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(12)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(10)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(11)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(12)}}, }, }, }, @@ -340,8 +343,8 @@ func TestDialStateStaticDial(t *testing.T) { // aren't yet connected. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, new: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(3)}}, @@ -353,9 +356,9 @@ func TestDialStateStaticDial(t *testing.T) { // nodes are either connected or still being dialed. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(3)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(3)}}, }, done: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(3)}}, @@ -365,11 +368,11 @@ func TestDialStateStaticDial(t *testing.T) { // nodes are now connected. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(3)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(4)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(5)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(3)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(4)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(5)}}, }, done: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(4)}}, @@ -382,20 +385,20 @@ func TestDialStateStaticDial(t *testing.T) { // Wait a round for dial history to expire, no new tasks should spawn. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(3)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(4)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(5)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(3)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(4)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(5)}}, }, }, // If a static node is dropped, it should be immediately redialed, // irrespective whether it was originally static or dynamic. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(3)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(5)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(3)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(5)}}, }, new: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(2)}}, @@ -431,8 +434,8 @@ func TestDialStateCache(t *testing.T) { // nodes are either connected or still being dialed. { peers: []*Peer{ - {rw: &conn{flags: staticDialedConn, id: uintID(1)}}, - {rw: &conn{flags: staticDialedConn, id: uintID(2)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(1)}}, + {conn: &conn{flags: staticDialedConn, id: uintID(2)}}, }, done: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(1)}}, @@ -443,8 +446,8 @@ func TestDialStateCache(t *testing.T) { // entry to expire. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, done: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(3)}}, @@ -456,15 +459,15 @@ func TestDialStateCache(t *testing.T) { // Still waiting for node 3's entry to expire in the cache. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, }, // The cache entry for node 3 has expired and is retried. { peers: []*Peer{ - {rw: &conn{flags: dynDialedConn, id: uintID(1)}}, - {rw: &conn{flags: dynDialedConn, id: uintID(2)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(1)}}, + {conn: &conn{flags: dynDialedConn, id: uintID(2)}}, }, new: []task{ &dialTask{staticDialedConn, &discover.Node{ID: uintID(3)}}, diff --git a/p2p/discover/database.go b/p2p/discover/database.go index e8e3371ff9..daa7af99c7 100644 --- a/p2p/discover/database.go +++ b/p2p/discover/database.go @@ -180,7 +180,7 @@ func (db *nodeDB) storeInt64(key []byte, n int64) error { func (db *nodeDB) node(id NodeID) *Node { blob, err := db.lvl.Get(makeKey(id, nodeDBDiscoverRoot), nil) if err != nil { - glog.V(logger.Detail).Infof("failed to retrieve node %v: %v", id, err) + glog.V(logger.Detail).Infof("node %x: %v", id[:8], err) return nil } node := new(Node) diff --git a/p2p/message.go b/p2p/message.go index 1292d21213..7758796a4e 100644 --- a/p2p/message.go +++ b/p2p/message.go @@ -22,8 +22,6 @@ import ( "fmt" "io" "io/ioutil" - "net" - "sync" "sync/atomic" "time" @@ -38,9 +36,10 @@ import ( // structure, encode the payload into a byte array and create a // separate Msg with a bytes.Reader as Payload for each send. type Msg struct { - Code uint64 - Size uint32 // size of the paylod - Payload io.Reader + Code uint64 + Size uint32 // size of the paylod + Payload io.Reader + ReceivedAt time.Time } @@ -66,10 +65,12 @@ func (msg Msg) Discard() error { return err } +// MsgReader wraps the ReadMsg operation. type MsgReader interface { ReadMsg() (Msg, error) } +// MsgWriter wraps the WriteMsg operation. type MsgWriter interface { // WriteMsg sends a message. It will block until the message's // Payload has been consumed by the other end. @@ -110,30 +111,6 @@ func SendItems(w MsgWriter, msgcode uint64, elems ...interface{}) error { return Send(w, msgcode, elems) } -// netWrapper wraps a MsgReadWriter with locks around -// ReadMsg/WriteMsg and applies read/write deadlines. -type netWrapper struct { - rmu, wmu sync.Mutex - - rtimeout, wtimeout time.Duration - conn net.Conn - wrapped MsgReadWriter -} - -func (rw *netWrapper) ReadMsg() (Msg, error) { - rw.rmu.Lock() - defer rw.rmu.Unlock() - rw.conn.SetReadDeadline(time.Now().Add(rw.rtimeout)) - return rw.wrapped.ReadMsg() -} - -func (rw *netWrapper) WriteMsg(msg Msg) error { - rw.wmu.Lock() - defer rw.wmu.Unlock() - rw.conn.SetWriteDeadline(time.Now().Add(rw.wtimeout)) - return rw.wrapped.WriteMsg(msg) -} - // eofSignal wraps a reader with eof signaling. the eof channel is // closed when the wrapped reader returns an error or when count bytes // have been read. diff --git a/p2p/message_test.go b/p2p/message_test.go index 8599b7e87c..5b58654aee 100644 --- a/p2p/message_test.go +++ b/p2p/message_test.go @@ -18,11 +18,9 @@ package p2p import ( "bytes" - "encoding/hex" "fmt" "io" "runtime" - "strings" "testing" "time" ) @@ -141,11 +139,3 @@ func TestEOFSignal(t *testing.T) { default: } } - -func unhex(str string) []byte { - b, err := hex.DecodeString(strings.Replace(str, "\n", "", -1)) - if err != nil { - panic(fmt.Sprintf("invalid hex string: %q", str)) - } - return b -} diff --git a/p2p/peer.go b/p2p/peer.go index 72ed4069c9..aeecc11b21 100644 --- a/p2p/peer.go +++ b/p2p/peer.go @@ -25,44 +25,21 @@ import ( "sync" "time" + "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" "github.com/ethereum/go-ethereum/p2p/discover" "github.com/ethereum/go-ethereum/rlp" ) -const ( - baseProtocolVersion = 4 - baseProtocolLength = uint64(16) - baseProtocolMaxMsgSize = 2 * 1024 - - pingInterval = 15 * time.Second -) - -const ( - // devp2p message codes - handshakeMsg = 0x00 - discMsg = 0x01 - pingMsg = 0x02 - pongMsg = 0x03 - getPeersMsg = 0x04 - peersMsg = 0x05 -) - -// protoHandshake is the RLP structure of the protocol handshake. -type protoHandshake struct { - Version uint64 - Name string - Caps []Cap - ListenPort uint64 - ID discover.NodeID -} +const pingInterval = 15 * time.Second // Peer represents a connected remote node. type Peer struct { - rw *conn - running map[string]*protoRW + // contains an element for each running subprotocol (excluding devp2p). + running []*protoRW + conn *conn wg sync.WaitGroup protoErr chan error closed chan struct{} @@ -72,7 +49,9 @@ type Peer struct { // NewPeer returns a peer for testing purposes. func NewPeer(id discover.NodeID, name string, caps []Cap) *Peer { pipe, _ := net.Pipe() - conn := &conn{fd: pipe, transport: nil, id: id, caps: caps, name: name} + randomPriv, _ := crypto.GenerateKey() + dc := newDevConn(pipe, randomPriv, nil) + conn := &conn{transport: dc, id: id, caps: caps, name: name} peer := newPeer(conn, nil) close(peer.closed) // ensures Disconnect doesn't block return peer @@ -80,28 +59,28 @@ func NewPeer(id discover.NodeID, name string, caps []Cap) *Peer { // ID returns the node's public key. func (p *Peer) ID() discover.NodeID { - return p.rw.id + return p.conn.id } // Name returns the node name that the remote node advertised. func (p *Peer) Name() string { - return p.rw.name + return p.conn.name } // Caps returns the capabilities (supported subprotocols) of the remote peer. func (p *Peer) Caps() []Cap { // TODO: maybe return copy - return p.rw.caps + return p.conn.caps } // RemoteAddr returns the remote address of the network connection. func (p *Peer) RemoteAddr() net.Addr { - return p.rw.fd.RemoteAddr() + return p.conn.RemoteAddr() } // LocalAddr returns the local address of the network connection. func (p *Peer) LocalAddr() net.Addr { - return p.rw.fd.LocalAddr() + return p.conn.LocalAddr() } // Disconnect terminates the peer connection with the given reason. @@ -115,13 +94,13 @@ func (p *Peer) Disconnect(reason DiscReason) { // String implements fmt.Stringer. func (p *Peer) String() string { - return fmt.Sprintf("Peer %x %v", p.rw.id[:8], p.RemoteAddr()) + return fmt.Sprintf("Peer %x %v", p.conn.id[:8], p.RemoteAddr()) } func newPeer(conn *conn, protocols []Protocol) *Peer { - protomap := matchProtocols(protocols, conn.caps, conn) + protomap := matchProtocols(protocols, conn.caps) p := &Peer{ - rw: conn, + conn: conn, running: protomap, disc: make(chan DiscReason), protoErr: make(chan error, len(protomap)+1), // protocols + pingLoop @@ -132,33 +111,40 @@ func newPeer(conn *conn, protocols []Protocol) *Peer { func (p *Peer) run() DiscReason { var ( - writeStart = make(chan struct{}, 1) - writeErr = make(chan error, 1) - readErr = make(chan error, 1) - reason DiscReason - requested bool + writeErr = make(chan error, len(p.running)) + readErr = make(chan error, 1) + reason DiscReason + requested bool + // While most of the code works with the transport interface so it + // can be tested, using the connection requires an actual + // *devConn. + devconn = p.conn.transport.(*devConn) ) + + // Ensure that the RLPx handshake is done. The only time this will + // actually do anything is while testing because the tests don't + // trigger the handshake explicitly. + if err := devconn.Handshake(); err != nil { + return DiscProtocolError + } + p.wg.Add(2) - go p.readLoop(readErr) - go p.pingLoop() + go p.readLoop(devconn.protocols[0], readErr) + go p.pingLoop(devconn.protocols[0]) // Start all protocol handlers. - writeStart <- struct{}{} - p.startProtocols(writeStart, writeErr) + p.startProtocols(devconn, writeErr) // Wait for an error or disconnect. loop: for { select { case err := <-writeErr: - // A write finished. Allow the next write to start if - // there was no error. if err != nil { glog.V(logger.Detail).Infof("%v: write error: %v\n", p, err) reason = DiscNetworkError break loop } - writeStart <- struct{}{} case err := <-readErr: if r, ok := err.(DiscReason); ok { glog.V(logger.Debug).Infof("%v: remote requested disconnect: %v\n", p, r) @@ -180,7 +166,7 @@ loop: } close(p.closed) - p.rw.close(reason) + p.conn.close(reason) p.wg.Wait() if requested { reason = DiscRequested @@ -188,14 +174,14 @@ loop: return reason } -func (p *Peer) pingLoop() { +func (p *Peer) pingLoop(devp2p *devProtocol) { ping := time.NewTicker(pingInterval) defer p.wg.Done() defer ping.Stop() for { select { case <-ping.C: - if err := SendItems(p.rw, pingMsg); err != nil { + if err := SendItems(devp2p, pingMsg); err != nil { p.protoErr <- err return } @@ -205,27 +191,27 @@ func (p *Peer) pingLoop() { } } -func (p *Peer) readLoop(errc chan<- error) { +func (p *Peer) readLoop(devp2p *devProtocol, errc chan<- error) { defer p.wg.Done() for { - msg, err := p.rw.ReadMsg() + msg, err := devp2p.ReadMsg() if err != nil { errc <- err return } - msg.ReceivedAt = time.Now() - if err = p.handle(msg); err != nil { + if err = p.handle(devp2p, msg); err != nil { errc <- err return } } } -func (p *Peer) handle(msg Msg) error { +func (p *Peer) handle(devp2p *devProtocol, msg Msg) (err error) { switch { case msg.Code == pingMsg: msg.Discard() - go SendItems(p.rw, pongMsg) + go SendItems(devp2p, pongMsg) + return case msg.Code == discMsg: var reason [1]DiscReason // This is the last message. We don't need to discard or @@ -236,8 +222,10 @@ func (p *Peer) handle(msg Msg) error { // ignore other base protocol messages return msg.Discard() default: - // it's a subprotocol message + // Dispatch as subprotocol message by message code offset. + // This is how dispatch worked before chunking was implemented. proto, err := p.getProto(msg.Code) + msg.Code -= proto.offset if err != nil { return fmt.Errorf("msg code out of range: %v", msg.Code) } @@ -248,7 +236,6 @@ func (p *Peer) handle(msg Msg) error { return io.EOF } } - return nil } func countMatchingProtocols(protocols []Protocol, caps []Cap) int { @@ -263,24 +250,27 @@ func countMatchingProtocols(protocols []Protocol, caps []Cap) int { return n } -// matchProtocols creates structures for matching named subprotocols. -func matchProtocols(protocols []Protocol, caps []Cap, rw MsgReadWriter) map[string]*protoRW { +// matchProtocols creates protoRWs for matching named subprotocols. +func matchProtocols(protocols []Protocol, caps []Cap) []*protoRW { sort.Sort(capsByNameAndVersion(caps)) + i := 0 offset := baseProtocolLength - result := make(map[string]*protoRW) - + var result []*protoRW outer: for _, cap := range caps { for _, proto := range protocols { if proto.Name == cap.Name && proto.Version == cap.Version { - // If an old protocol version matched, revert it - if old := result[cap.Name]; old != nil { - offset -= old.Length + if i > 0 && result[i-1].Name == cap.Name { + // If the previous match was for the same protocol + // (with a lower version), reset the offset and replace it. + offset -= result[i-1].Protocol.Length + } else { + // Otherwise, append a new protocol. + result = append(result, nil) + i++ } - // Assign the new match - result[cap.Name] = &protoRW{Protocol: proto, offset: offset, in: make(chan Msg), w: rw} + result[i-1] = &protoRW{Protocol: proto, offset: offset} offset += proto.Length - continue outer } } @@ -288,13 +278,33 @@ outer: return result } -func (p *Peer) startProtocols(writeStart <-chan struct{}, writeErr chan<- error) { +func (p *Peer) startProtocols(dc *devConn, writeErr chan<- error) { + switch dc.Version() { + case 5: + // Acknowledge the protocols on the RLPx layer. This creates + // *devProtocol wrappers, dc.protocols[i] contains entries in + // range 1..len(p.running). + dc.addProtocols(len(p.running)) + for i, proto := range p.running { + proto.offset = 0 + proto.werr = writeErr + proto.rw = dc.protocols[i+1] + } + case 4: + // This is a legacy connection with offset-based dispatch. + for _, proto := range p.running { + proto.closed = p.closed + proto.in = make(chan Msg) + proto.werr = writeErr + proto.rw = dc.protocols[0] + } + default: + panic("conn has no version") + } + // Spawn Run for all protocols. p.wg.Add(len(p.running)) for _, proto := range p.running { proto := proto - proto.closed = p.closed - proto.wstart = writeStart - proto.werr = writeErr glog.V(logger.Detail).Infof("%v: Starting protocol %s/%d\n", p, proto.Name, proto.Version) go func() { err := proto.Run(p, proto) @@ -314,7 +324,7 @@ func (p *Peer) startProtocols(writeStart <-chan struct{}, writeErr chan<- error) // the given message code. func (p *Peer) getProto(code uint64) (*protoRW, error) { for _, proto := range p.running { - if code >= proto.offset && code < proto.offset+proto.Length { + if proto.offset > 0 && code >= proto.offset && code < proto.offset+proto.Length { return proto, nil } } @@ -323,37 +333,40 @@ func (p *Peer) getProto(code uint64) (*protoRW, error) { type protoRW struct { Protocol + offset uint64 + rw MsgReadWriter + werr chan<- error // for write results + + // for RLPx V4 offset-based dispatch in chan Msg // receices read messages closed <-chan struct{} // receives when peer is shutting down - wstart <-chan struct{} // receives when write may start - werr chan<- error // for write results - offset uint64 - w MsgWriter + index uint16 } -func (rw *protoRW) WriteMsg(msg Msg) (err error) { +func (rw *protoRW) WriteMsg(msg Msg) error { if msg.Code >= rw.Length { return newPeerError(errInvalidMsgCode, "not handled") } msg.Code += rw.offset - select { - case <-rw.wstart: - err = rw.w.WriteMsg(msg) - // Report write status back to Peer.run. It will initiate - // shutdown if the error is non-nil and unblock the next write - // otherwise. The calling protocol code should exit for errors - // as well but we don't want to rely on that. + err := rw.rw.WriteMsg(msg) + // Report write status back to Peer.run. It will initiate shutdown + // if the error is non-nil otherwise. The calling protocol should + // exit soon after, but might not return the error correctly. + if err != nil { rw.werr <- err - case <-rw.closed: - err = fmt.Errorf("shutting down") } + // TODO: maybe make the error sticky to prevent further writes return err } func (rw *protoRW) ReadMsg() (Msg, error) { + if rw.offset == 0 { + // RLPx version 5 + return rw.rw.ReadMsg() + } + // RLPx version 4 select { case msg := <-rw.in: - msg.Code -= rw.offset return msg, nil case <-rw.closed: return Msg{}, io.EOF diff --git a/p2p/peer_error.go b/p2p/peer_error.go index 62c7b665dd..e0fd5778d2 100644 --- a/p2p/peer_error.go +++ b/p2p/peer_error.go @@ -50,64 +50,3 @@ func newPeerError(code int, format string, v ...interface{}) *peerError { func (self *peerError) Error() string { return self.message } - -type DiscReason uint - -const ( - DiscRequested DiscReason = iota - DiscNetworkError - DiscProtocolError - DiscUselessPeer - DiscTooManyPeers - DiscAlreadyConnected - DiscIncompatibleVersion - DiscInvalidIdentity - DiscQuitting - DiscUnexpectedIdentity - DiscSelf - DiscReadTimeout - DiscSubprotocolError = 0x10 -) - -var discReasonToString = [...]string{ - DiscRequested: "Disconnect requested", - DiscNetworkError: "Network error", - DiscProtocolError: "Breach of protocol", - DiscUselessPeer: "Useless peer", - DiscTooManyPeers: "Too many peers", - DiscAlreadyConnected: "Already connected", - DiscIncompatibleVersion: "Incompatible P2P protocol version", - DiscInvalidIdentity: "Invalid node identity", - DiscQuitting: "Client quitting", - DiscUnexpectedIdentity: "Unexpected identity", - DiscSelf: "Connected to self", - DiscReadTimeout: "Read timeout", - DiscSubprotocolError: "Subprotocol error", -} - -func (d DiscReason) String() string { - if len(discReasonToString) < int(d) { - return fmt.Sprintf("Unknown Reason(%d)", d) - } - return discReasonToString[d] -} - -func (d DiscReason) Error() string { - return d.String() -} - -func discReasonForError(err error) DiscReason { - if reason, ok := err.(DiscReason); ok { - return reason - } - peerError, ok := err.(*peerError) - if ok { - switch peerError.code { - case errInvalidMsgCode, errInvalidMsg: - return DiscProtocolError - default: - return DiscSubprotocolError - } - } - return DiscSubprotocolError -} diff --git a/p2p/peer_test.go b/p2p/peer_test.go index 6f96a823b4..5617a92b33 100644 --- a/p2p/peer_test.go +++ b/p2p/peer_test.go @@ -24,6 +24,8 @@ import ( "reflect" "testing" "time" + + "github.com/davecgh/go-spew/spew" ) var discard = Protocol{ @@ -43,21 +45,20 @@ var discard = Protocol{ }, } -func testPeer(protos []Protocol) (func(), *conn, *Peer, <-chan DiscReason) { +func testPeer(protos []Protocol) (*devConn, *Peer, <-chan DiscReason) { fd1, fd2 := net.Pipe() - c1 := &conn{fd: fd1, transport: newTestTransport(randomID(), fd1)} - c2 := &conn{fd: fd2, transport: newTestTransport(randomID(), fd2)} + k1, k2 := newkey(), newkey() + c1 := &conn{transport: newDevConn(fd1, k1, &k2.PublicKey)} for _, p := range protos { c1.caps = append(c1.caps, p.cap()) - c2.caps = append(c2.caps, p.cap()) } - peer := newPeer(c1, protos) errc := make(chan DiscReason, 1) go func() { errc <- peer.run() }() - closer := func() { c2.close(errors.New("close func called")) } - return closer, c2, peer, errc + c2 := newDevConn(fd2, k2, nil) + c2.addProtocols(len(protos)) + return c2, peer, errc } func TestPeerProtoReadMsg(t *testing.T) { @@ -80,12 +81,12 @@ func TestPeerProtoReadMsg(t *testing.T) { }, } - closer, rw, _, errc := testPeer([]Protocol{proto}) - defer closer() + conn, _, errc := testPeer([]Protocol{proto}) + defer conn.Close() - Send(rw, baseProtocolLength+2, []uint{1}) - Send(rw, baseProtocolLength+3, []uint{2}) - Send(rw, baseProtocolLength+4, []uint{3}) + Send(conn.protocols[1], 2, []uint{1}) + Send(conn.protocols[1], 3, []uint{2}) + Send(conn.protocols[1], 4, []uint{3}) select { case <-done: @@ -110,29 +111,29 @@ func TestPeerProtoEncodeMsg(t *testing.T) { return nil }, } - closer, rw, _, _ := testPeer([]Protocol{proto}) - defer closer() + conn, _, _ := testPeer([]Protocol{proto}) + defer conn.Close() - if err := ExpectMsg(rw, 17, []string{"foo", "bar"}); err != nil { + if err := ExpectMsg(conn.protocols[1], 1, []string{"foo", "bar"}); err != nil { t.Error(err) } } func TestPeerPing(t *testing.T) { - closer, rw, _, _ := testPeer(nil) - defer closer() - if err := SendItems(rw, pingMsg); err != nil { + conn, _, _ := testPeer(nil) + defer conn.Close() + if err := SendItems(conn.protocols[0], pingMsg); err != nil { t.Fatal(err) } - if err := ExpectMsg(rw, pongMsg, nil); err != nil { + if err := ExpectMsg(conn.protocols[0], pongMsg, nil); err != nil { t.Error(err) } } func TestPeerDisconnect(t *testing.T) { - closer, rw, _, disc := testPeer(nil) - defer closer() - if err := SendItems(rw, discMsg, DiscQuitting); err != nil { + conn, _, disc := testPeer(nil) + defer conn.Close() + if err := SendItems(conn.protocols[0], discMsg, DiscQuitting); err != nil { t.Fatal(err) } select { @@ -150,10 +151,10 @@ func TestPeerDisconnect(t *testing.T) { func TestPeerDisconnectRace(t *testing.T) { maybe := func() bool { return rand.Intn(1) == 1 } - for i := 0; i < 1000; i++ { - protoclose := make(chan error) - protodisc := make(chan DiscReason) - closer, rw, p, disc := testPeer([]Protocol{ + for i := 0; i < 100; i++ { + protoclose := make(chan error, 1) + protodisc := make(chan DiscReason, 1) + conn, p, disc := testPeer([]Protocol{ { Name: "closereq", Run: func(p *Peer, rw MsgReadWriter) error { return <-protoclose }, @@ -165,12 +166,13 @@ func TestPeerDisconnectRace(t *testing.T) { Length: 1, }, }) + conn.Handshake() // Simulate incoming messages. - go SendItems(rw, baseProtocolLength+1) - go SendItems(rw, baseProtocolLength+2) + go SendItems(conn.protocols[1], 1) + go SendItems(conn.protocols[2], 2) // Close the network connection. - go closer() + go conn.Close() // Make protocol "closereq" return. protoclose <- errors.New("protocol closed") // Make protocol "disconnect" call peer.Disconnect @@ -181,7 +183,7 @@ func TestPeerDisconnectRace(t *testing.T) { } // In some cases, simulate remote requesting a disconnect. if maybe() { - go SendItems(rw, discMsg, DiscQuitting) + go SendItems(conn.protocols[0], discMsg, DiscQuitting) } select { @@ -214,96 +216,78 @@ func TestNewPeer(t *testing.T) { } func TestMatchProtocols(t *testing.T) { - tests := []struct { + tests := map[string]struct { Remote []Cap Local []Protocol - Match map[string]protoRW + Match []*protoRW }{ - { - // No remote capabilities + "no remote caps": { Local: []Protocol{{Name: "a"}}, }, - { - // No local protocols + "no local protocols": { Remote: []Cap{{Name: "a"}}, }, - { - // No mutual protocols + "no mutual protocols": { Remote: []Cap{{Name: "a"}}, Local: []Protocol{{Name: "b"}}, }, - { - // Some matches, some differences + "some matches": { Remote: []Cap{{Name: "local"}, {Name: "match1"}, {Name: "match2"}}, Local: []Protocol{{Name: "match1"}, {Name: "match2"}, {Name: "remote"}}, - Match: map[string]protoRW{"match1": {Protocol: Protocol{Name: "match1"}}, "match2": {Protocol: Protocol{Name: "match2"}}}, + Match: []*protoRW{ + {Protocol: Protocol{Name: "match1"}, offset: 16}, + {Protocol: Protocol{Name: "match2"}, offset: 16}, + }, }, - { - // Various alphabetical ordering + "alphabetical ordering": { Remote: []Cap{{Name: "aa"}, {Name: "ab"}, {Name: "bb"}, {Name: "ba"}}, Local: []Protocol{{Name: "ba"}, {Name: "bb"}, {Name: "ab"}, {Name: "aa"}}, - Match: map[string]protoRW{"aa": {Protocol: Protocol{Name: "aa"}}, "ab": {Protocol: Protocol{Name: "ab"}}, "ba": {Protocol: Protocol{Name: "ba"}}, "bb": {Protocol: Protocol{Name: "bb"}}}, + Match: []*protoRW{ + {Protocol: Protocol{Name: "aa"}, offset: 16}, + {Protocol: Protocol{Name: "ab"}, offset: 16}, + {Protocol: Protocol{Name: "ba"}, offset: 16}, + {Protocol: Protocol{Name: "bb"}, offset: 16}, + }, }, - { - // No mutual versions + "no mutual versions": { Remote: []Cap{{Version: 1}}, Local: []Protocol{{Version: 2}}, }, - { - // Multiple versions, single common + "multiple versions, single common": { Remote: []Cap{{Version: 1}, {Version: 2}}, Local: []Protocol{{Version: 2}, {Version: 3}}, - Match: map[string]protoRW{"": {Protocol: Protocol{Version: 2}}}, + Match: []*protoRW{ + {Protocol: Protocol{Version: 2}, offset: 16}, + }, }, - { - // Multiple versions, multiple common + "multiple versions, multiple common": { Remote: []Cap{{Version: 1}, {Version: 2}, {Version: 3}, {Version: 4}}, Local: []Protocol{{Version: 2}, {Version: 3}}, - Match: map[string]protoRW{"": {Protocol: Protocol{Version: 3}}}, + Match: []*protoRW{ + {Protocol: Protocol{Version: 3}, offset: 16}, + }, }, - { - // Various version orderings + "version ordering": { Remote: []Cap{{Version: 4}, {Version: 1}, {Version: 3}, {Version: 2}}, Local: []Protocol{{Version: 2}, {Version: 3}, {Version: 1}}, - Match: map[string]protoRW{"": {Protocol: Protocol{Version: 3}}}, + Match: []*protoRW{ + {Protocol: Protocol{Version: 3}, offset: 16}, + }, }, - { - // Versions overriding sub-protocol lengths + "versions overriding subprotocol lengths": { Remote: []Cap{{Version: 1}, {Version: 2}, {Version: 3}, {Name: "a"}}, Local: []Protocol{{Version: 1, Length: 1}, {Version: 2, Length: 2}, {Version: 3, Length: 3}, {Name: "a"}}, - Match: map[string]protoRW{"": {Protocol: Protocol{Version: 3}}, "a": {Protocol: Protocol{Name: "a"}, offset: 3}}, + Match: []*protoRW{ + {Protocol: Protocol{Version: 3, Length: 3}, offset: 16}, + {Protocol: Protocol{Name: "a"}, offset: 19}, + }, }, } - for i, tt := range tests { - result := matchProtocols(tt.Local, tt.Remote, nil) - if len(result) != len(tt.Match) { - t.Errorf("test %d: negotiation mismatch: have %v, want %v", i, len(result), len(tt.Match)) - continue - } - // Make sure all negotiated protocols are needed and correct - for name, proto := range result { - match, ok := tt.Match[name] - if !ok { - t.Errorf("test %d, proto '%s': negotiated but shouldn't have", i, name) - continue - } - if proto.Name != match.Name { - t.Errorf("test %d, proto '%s': name mismatch: have %v, want %v", i, name, proto.Name, match.Name) - } - if proto.Version != match.Version { - t.Errorf("test %d, proto '%s': version mismatch: have %v, want %v", i, name, proto.Version, match.Version) - } - if proto.offset-baseProtocolLength != match.offset { - t.Errorf("test %d, proto '%s': offset mismatch: have %v, want %v", i, name, proto.offset-baseProtocolLength, match.offset) - } - } - // Make sure no protocols missed negotiation - for name, _ := range tt.Match { - if _, ok := result[name]; !ok { - t.Errorf("test %d, proto '%s': not negotiated, should have", i, name) - continue - } + for tname, tt := range tests { + result := matchProtocols(tt.Local, tt.Remote) + if !reflect.DeepEqual(result, tt.Match) { + t.Errorf("%s: wrong result\ngot %s\nwant: %s", tname, spew.Sdump(result), spew.Sdump(tt.Match)) } } } diff --git a/p2p/rlpx.go b/p2p/rlpx.go deleted file mode 100644 index 8f429d6ec1..0000000000 --- a/p2p/rlpx.go +++ /dev/null @@ -1,635 +0,0 @@ -// Copyright 2015 The go-ethereum Authors -// This file is part of the go-ethereum library. -// -// The go-ethereum library is free software: you can redistribute it and/or modify -// it under the terms of the GNU Lesser General Public License as published by -// the Free Software Foundation, either version 3 of the License, or -// (at your option) any later version. -// -// The go-ethereum library is distributed in the hope that it will be useful, -// but WITHOUT ANY WARRANTY; without even the implied warranty of -// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the -// GNU Lesser General Public License for more details. -// -// You should have received a copy of the GNU Lesser General Public License -// along with the go-ethereum library. If not, see . - -package p2p - -import ( - "bytes" - "crypto/aes" - "crypto/cipher" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/hmac" - "crypto/rand" - "errors" - "fmt" - "hash" - "io" - "net" - "sync" - "time" - - "github.com/ethereum/go-ethereum/crypto" - "github.com/ethereum/go-ethereum/crypto/ecies" - "github.com/ethereum/go-ethereum/crypto/secp256k1" - "github.com/ethereum/go-ethereum/crypto/sha3" - "github.com/ethereum/go-ethereum/p2p/discover" - "github.com/ethereum/go-ethereum/rlp" -) - -const ( - maxUint24 = ^uint32(0) >> 8 - - sskLen = 16 // ecies.MaxSharedKeyLength(pubKey) / 2 - sigLen = 65 // elliptic S256 - pubLen = 64 // 512 bit pubkey in uncompressed representation without format byte - shaLen = 32 // hash length (for nonce etc) - - authMsgLen = sigLen + shaLen + pubLen + shaLen + 1 - authRespLen = pubLen + shaLen + 1 - - eciesBytes = 65 + 16 + 32 - encAuthMsgLen = authMsgLen + eciesBytes // size of the final ECIES payload sent as initiator's handshake - encAuthRespLen = authRespLen + eciesBytes // size of the final ECIES payload sent as receiver's handshake - - // total timeout for encryption handshake and protocol - // handshake in both directions. - handshakeTimeout = 5 * time.Second - - // This is the timeout for sending the disconnect reason. - // This is shorter than the usual timeout because we don't want - // to wait if the connection is known to be bad anyway. - discWriteTimeout = 1 * time.Second -) - -// rlpx is the transport protocol used by actual (non-test) connections. -// It wraps the frame encoder with locks and read/write deadlines. -type rlpx struct { - fd net.Conn - - rmu, wmu sync.Mutex - rw *rlpxFrameRW -} - -func newRLPX(fd net.Conn) transport { - fd.SetDeadline(time.Now().Add(handshakeTimeout)) - return &rlpx{fd: fd} -} - -func (t *rlpx) ReadMsg() (Msg, error) { - t.rmu.Lock() - defer t.rmu.Unlock() - t.fd.SetReadDeadline(time.Now().Add(frameReadTimeout)) - return t.rw.ReadMsg() -} - -func (t *rlpx) WriteMsg(msg Msg) error { - t.wmu.Lock() - defer t.wmu.Unlock() - t.fd.SetWriteDeadline(time.Now().Add(frameWriteTimeout)) - return t.rw.WriteMsg(msg) -} - -func (t *rlpx) close(err error) { - t.wmu.Lock() - defer t.wmu.Unlock() - // Tell the remote end why we're disconnecting if possible. - if t.rw != nil { - if r, ok := err.(DiscReason); ok && r != DiscNetworkError { - t.fd.SetWriteDeadline(time.Now().Add(discWriteTimeout)) - SendItems(t.rw, discMsg, r) - } - } - t.fd.Close() -} - -// doEncHandshake runs the protocol handshake using authenticated -// messages. the protocol handshake is the first authenticated message -// and also verifies whether the encryption handshake 'worked' and the -// remote side actually provided the right public key. -func (t *rlpx) doProtoHandshake(our *protoHandshake) (their *protoHandshake, err error) { - // Writing our handshake happens concurrently, we prefer - // returning the handshake read error. If the remote side - // disconnects us early with a valid reason, we should return it - // as the error so it can be tracked elsewhere. - werr := make(chan error, 1) - go func() { werr <- Send(t.rw, handshakeMsg, our) }() - if their, err = readProtocolHandshake(t.rw, our); err != nil { - <-werr // make sure the write terminates too - return nil, err - } - if err := <-werr; err != nil { - return nil, fmt.Errorf("write error: %v", err) - } - return their, nil -} - -func readProtocolHandshake(rw MsgReader, our *protoHandshake) (*protoHandshake, error) { - msg, err := rw.ReadMsg() - if err != nil { - return nil, err - } - if msg.Size > baseProtocolMaxMsgSize { - return nil, fmt.Errorf("message too big") - } - if msg.Code == discMsg { - // Disconnect before protocol handshake is valid according to the - // spec and we send it ourself if the posthanshake checks fail. - // We can't return the reason directly, though, because it is echoed - // back otherwise. Wrap it in a string instead. - var reason [1]DiscReason - rlp.Decode(msg.Payload, &reason) - return nil, reason[0] - } - if msg.Code != handshakeMsg { - return nil, fmt.Errorf("expected handshake, got %x", msg.Code) - } - var hs protoHandshake - if err := msg.Decode(&hs); err != nil { - return nil, err - } - // validate handshake info - if hs.Version != our.Version { - return nil, DiscIncompatibleVersion - } - if (hs.ID == discover.NodeID{}) { - return nil, DiscInvalidIdentity - } - return &hs, nil -} - -func (t *rlpx) doEncHandshake(prv *ecdsa.PrivateKey, dial *discover.Node) (discover.NodeID, error) { - var ( - sec secrets - err error - ) - if dial == nil { - sec, err = receiverEncHandshake(t.fd, prv, nil) - } else { - sec, err = initiatorEncHandshake(t.fd, prv, dial.ID, nil) - } - if err != nil { - return discover.NodeID{}, err - } - t.wmu.Lock() - t.rw = newRLPXFrameRW(t.fd, sec) - t.wmu.Unlock() - return sec.RemoteID, nil -} - -// encHandshake contains the state of the encryption handshake. -type encHandshake struct { - initiator bool - remoteID discover.NodeID - - remotePub *ecies.PublicKey // remote-pubk - initNonce, respNonce []byte // nonce - randomPrivKey *ecies.PrivateKey // ecdhe-random - remoteRandomPub *ecies.PublicKey // ecdhe-random-pubk -} - -// secrets represents the connection secrets -// which are negotiated during the encryption handshake. -type secrets struct { - RemoteID discover.NodeID - AES, MAC []byte - EgressMAC, IngressMAC hash.Hash - Token []byte -} - -// secrets is called after the handshake is completed. -// It extracts the connection secrets from the handshake values. -func (h *encHandshake) secrets(auth, authResp []byte) (secrets, error) { - ecdheSecret, err := h.randomPrivKey.GenerateShared(h.remoteRandomPub, sskLen, sskLen) - if err != nil { - return secrets{}, err - } - - // derive base secrets from ephemeral key agreement - sharedSecret := crypto.Sha3(ecdheSecret, crypto.Sha3(h.respNonce, h.initNonce)) - aesSecret := crypto.Sha3(ecdheSecret, sharedSecret) - s := secrets{ - RemoteID: h.remoteID, - AES: aesSecret, - MAC: crypto.Sha3(ecdheSecret, aesSecret), - Token: crypto.Sha3(sharedSecret), - } - - // setup sha3 instances for the MACs - mac1 := sha3.NewKeccak256() - mac1.Write(xor(s.MAC, h.respNonce)) - mac1.Write(auth) - mac2 := sha3.NewKeccak256() - mac2.Write(xor(s.MAC, h.initNonce)) - mac2.Write(authResp) - if h.initiator { - s.EgressMAC, s.IngressMAC = mac1, mac2 - } else { - s.EgressMAC, s.IngressMAC = mac2, mac1 - } - - return s, nil -} - -func (h *encHandshake) ecdhShared(prv *ecdsa.PrivateKey) ([]byte, error) { - return ecies.ImportECDSA(prv).GenerateShared(h.remotePub, sskLen, sskLen) -} - -// initiatorEncHandshake negotiates a session token on conn. -// it should be called on the dialing side of the connection. -// -// prv is the local client's private key. -// token is the token from a previous session with this node. -func initiatorEncHandshake(conn io.ReadWriter, prv *ecdsa.PrivateKey, remoteID discover.NodeID, token []byte) (s secrets, err error) { - h, err := newInitiatorHandshake(remoteID) - if err != nil { - return s, err - } - auth, err := h.authMsg(prv, token) - if err != nil { - return s, err - } - if _, err = conn.Write(auth); err != nil { - return s, err - } - - response := make([]byte, encAuthRespLen) - if _, err = io.ReadFull(conn, response); err != nil { - return s, err - } - if err := h.decodeAuthResp(response, prv); err != nil { - return s, err - } - return h.secrets(auth, response) -} - -func newInitiatorHandshake(remoteID discover.NodeID) (*encHandshake, error) { - rpub, err := remoteID.Pubkey() - if err != nil { - return nil, fmt.Errorf("bad remoteID: %v", err) - } - // generate random initiator nonce - n := make([]byte, shaLen) - if _, err := rand.Read(n); err != nil { - return nil, err - } - // generate random keypair to use for signing - randpriv, err := ecies.GenerateKey(rand.Reader, secp256k1.S256(), nil) - if err != nil { - return nil, err - } - h := &encHandshake{ - initiator: true, - remoteID: remoteID, - remotePub: ecies.ImportECDSAPublic(rpub), - initNonce: n, - randomPrivKey: randpriv, - } - return h, nil -} - -// authMsg creates an encrypted initiator handshake message. -func (h *encHandshake) authMsg(prv *ecdsa.PrivateKey, token []byte) ([]byte, error) { - var tokenFlag byte - if token == nil { - // no session token found means we need to generate shared secret. - // ecies shared secret is used as initial session token for new peers - // generate shared key from prv and remote pubkey - var err error - if token, err = h.ecdhShared(prv); err != nil { - return nil, err - } - } else { - // for known peers, we use stored token from the previous session - tokenFlag = 0x01 - } - - // sign known message: - // ecdh-shared-secret^nonce for new peers - // token^nonce for old peers - signed := xor(token, h.initNonce) - signature, err := crypto.Sign(signed, h.randomPrivKey.ExportECDSA()) - if err != nil { - return nil, err - } - - // encode auth message - // signature || sha3(ecdhe-random-pubk) || pubk || nonce || token-flag - msg := make([]byte, authMsgLen) - n := copy(msg, signature) - n += copy(msg[n:], crypto.Sha3(exportPubkey(&h.randomPrivKey.PublicKey))) - n += copy(msg[n:], crypto.FromECDSAPub(&prv.PublicKey)[1:]) - n += copy(msg[n:], h.initNonce) - msg[n] = tokenFlag - - // encrypt auth message using remote-pubk - return ecies.Encrypt(rand.Reader, h.remotePub, msg, nil, nil) -} - -// decodeAuthResp decode an encrypted authentication response message. -func (h *encHandshake) decodeAuthResp(auth []byte, prv *ecdsa.PrivateKey) error { - msg, err := crypto.Decrypt(prv, auth) - if err != nil { - return fmt.Errorf("could not decrypt auth response (%v)", err) - } - h.respNonce = msg[pubLen : pubLen+shaLen] - h.remoteRandomPub, err = importPublicKey(msg[:pubLen]) - if err != nil { - return err - } - // ignore token flag for now - return nil -} - -// receiverEncHandshake negotiates a session token on conn. -// it should be called on the listening side of the connection. -// -// prv is the local client's private key. -// token is the token from a previous session with this node. -func receiverEncHandshake(conn io.ReadWriter, prv *ecdsa.PrivateKey, token []byte) (s secrets, err error) { - // read remote auth sent by initiator. - auth := make([]byte, encAuthMsgLen) - if _, err := io.ReadFull(conn, auth); err != nil { - return s, err - } - h, err := decodeAuthMsg(prv, token, auth) - if err != nil { - return s, err - } - - // send auth response - resp, err := h.authResp(prv, token) - if err != nil { - return s, err - } - if _, err = conn.Write(resp); err != nil { - return s, err - } - - return h.secrets(auth, resp) -} - -func decodeAuthMsg(prv *ecdsa.PrivateKey, token []byte, auth []byte) (*encHandshake, error) { - var err error - h := new(encHandshake) - // generate random keypair for session - h.randomPrivKey, err = ecies.GenerateKey(rand.Reader, secp256k1.S256(), nil) - if err != nil { - return nil, err - } - // generate random nonce - h.respNonce = make([]byte, shaLen) - if _, err = rand.Read(h.respNonce); err != nil { - return nil, err - } - - msg, err := crypto.Decrypt(prv, auth) - if err != nil { - return nil, fmt.Errorf("could not decrypt auth message (%v)", err) - } - - // decode message parameters - // signature || sha3(ecdhe-random-pubk) || pubk || nonce || token-flag - h.initNonce = msg[authMsgLen-shaLen-1 : authMsgLen-1] - copy(h.remoteID[:], msg[sigLen+shaLen:sigLen+shaLen+pubLen]) - rpub, err := h.remoteID.Pubkey() - if err != nil { - return nil, fmt.Errorf("bad remoteID: %#v", err) - } - h.remotePub = ecies.ImportECDSAPublic(rpub) - - // recover remote random pubkey from signed message. - if token == nil { - // TODO: it is an error if the initiator has a token and we don't. check that. - - // no session token means we need to generate shared secret. - // ecies shared secret is used as initial session token for new peers. - // generate shared key from prv and remote pubkey. - if token, err = h.ecdhShared(prv); err != nil { - return nil, err - } - } - signedMsg := xor(token, h.initNonce) - remoteRandomPub, err := secp256k1.RecoverPubkey(signedMsg, msg[:sigLen]) - if err != nil { - return nil, err - } - - // validate the sha3 of recovered pubkey - remoteRandomPubMAC := msg[sigLen : sigLen+shaLen] - shaRemoteRandomPub := crypto.Sha3(remoteRandomPub[1:]) - if !bytes.Equal(remoteRandomPubMAC, shaRemoteRandomPub) { - return nil, fmt.Errorf("sha3 of recovered ephemeral pubkey does not match checksum in auth message") - } - - h.remoteRandomPub, _ = importPublicKey(remoteRandomPub) - return h, nil -} - -// authResp generates the encrypted authentication response message. -func (h *encHandshake) authResp(prv *ecdsa.PrivateKey, token []byte) ([]byte, error) { - // responder auth message - // E(remote-pubk, ecdhe-random-pubk || nonce || 0x0) - resp := make([]byte, authRespLen) - n := copy(resp, exportPubkey(&h.randomPrivKey.PublicKey)) - n += copy(resp[n:], h.respNonce) - if token == nil { - resp[n] = 0 - } else { - resp[n] = 1 - } - // encrypt using remote-pubk - return ecies.Encrypt(rand.Reader, h.remotePub, resp, nil, nil) -} - -// importPublicKey unmarshals 512 bit public keys. -func importPublicKey(pubKey []byte) (*ecies.PublicKey, error) { - var pubKey65 []byte - switch len(pubKey) { - case 64: - // add 'uncompressed key' flag - pubKey65 = append([]byte{0x04}, pubKey...) - case 65: - pubKey65 = pubKey - default: - return nil, fmt.Errorf("invalid public key length %v (expect 64/65)", len(pubKey)) - } - // TODO: fewer pointless conversions - return ecies.ImportECDSAPublic(crypto.ToECDSAPub(pubKey65)), nil -} - -func exportPubkey(pub *ecies.PublicKey) []byte { - if pub == nil { - panic("nil pubkey") - } - return elliptic.Marshal(pub.Curve, pub.X, pub.Y)[1:] -} - -func xor(one, other []byte) (xor []byte) { - xor = make([]byte, len(one)) - for i := 0; i < len(one); i++ { - xor[i] = one[i] ^ other[i] - } - return xor -} - -var ( - // this is used in place of actual frame header data. - // TODO: replace this when Msg contains the protocol type code. - zeroHeader = []byte{0xC2, 0x80, 0x80} - // sixteen zero bytes - zero16 = make([]byte, 16) -) - -// rlpxFrameRW implements a simplified version of RLPx framing. -// chunked messages are not supported and all headers are equal to -// zeroHeader. -// -// rlpxFrameRW is not safe for concurrent use from multiple goroutines. -type rlpxFrameRW struct { - conn io.ReadWriter - enc cipher.Stream - dec cipher.Stream - - macCipher cipher.Block - egressMAC hash.Hash - ingressMAC hash.Hash -} - -func newRLPXFrameRW(conn io.ReadWriter, s secrets) *rlpxFrameRW { - macc, err := aes.NewCipher(s.MAC) - if err != nil { - panic("invalid MAC secret: " + err.Error()) - } - encc, err := aes.NewCipher(s.AES) - if err != nil { - panic("invalid AES secret: " + err.Error()) - } - // we use an all-zeroes IV for AES because the key used - // for encryption is ephemeral. - iv := make([]byte, encc.BlockSize()) - return &rlpxFrameRW{ - conn: conn, - enc: cipher.NewCTR(encc, iv), - dec: cipher.NewCTR(encc, iv), - macCipher: macc, - egressMAC: s.EgressMAC, - ingressMAC: s.IngressMAC, - } -} - -func (rw *rlpxFrameRW) WriteMsg(msg Msg) error { - ptype, _ := rlp.EncodeToBytes(msg.Code) - - // write header - headbuf := make([]byte, 32) - fsize := uint32(len(ptype)) + msg.Size - if fsize > maxUint24 { - return errors.New("message size overflows uint24") - } - putInt24(fsize, headbuf) // TODO: check overflow - copy(headbuf[3:], zeroHeader) - rw.enc.XORKeyStream(headbuf[:16], headbuf[:16]) // first half is now encrypted - - // write header MAC - copy(headbuf[16:], updateMAC(rw.egressMAC, rw.macCipher, headbuf[:16])) - if _, err := rw.conn.Write(headbuf); err != nil { - return err - } - - // write encrypted frame, updating the egress MAC hash with - // the data written to conn. - tee := cipher.StreamWriter{S: rw.enc, W: io.MultiWriter(rw.conn, rw.egressMAC)} - if _, err := tee.Write(ptype); err != nil { - return err - } - if _, err := io.Copy(tee, msg.Payload); err != nil { - return err - } - if padding := fsize % 16; padding > 0 { - if _, err := tee.Write(zero16[:16-padding]); err != nil { - return err - } - } - - // write frame MAC. egress MAC hash is up to date because - // frame content was written to it as well. - fmacseed := rw.egressMAC.Sum(nil) - mac := updateMAC(rw.egressMAC, rw.macCipher, fmacseed) - _, err := rw.conn.Write(mac) - return err -} - -func (rw *rlpxFrameRW) ReadMsg() (msg Msg, err error) { - // read the header - headbuf := make([]byte, 32) - if _, err := io.ReadFull(rw.conn, headbuf); err != nil { - return msg, err - } - // verify header mac - shouldMAC := updateMAC(rw.ingressMAC, rw.macCipher, headbuf[:16]) - if !hmac.Equal(shouldMAC, headbuf[16:]) { - return msg, errors.New("bad header MAC") - } - rw.dec.XORKeyStream(headbuf[:16], headbuf[:16]) // first half is now decrypted - fsize := readInt24(headbuf) - // ignore protocol type for now - - // read the frame content - var rsize = fsize // frame size rounded up to 16 byte boundary - if padding := fsize % 16; padding > 0 { - rsize += 16 - padding - } - framebuf := make([]byte, rsize) - if _, err := io.ReadFull(rw.conn, framebuf); err != nil { - return msg, err - } - - // read and validate frame MAC. we can re-use headbuf for that. - rw.ingressMAC.Write(framebuf) - fmacseed := rw.ingressMAC.Sum(nil) - if _, err := io.ReadFull(rw.conn, headbuf[:16]); err != nil { - return msg, err - } - shouldMAC = updateMAC(rw.ingressMAC, rw.macCipher, fmacseed) - if !hmac.Equal(shouldMAC, headbuf[:16]) { - return msg, errors.New("bad frame MAC") - } - - // decrypt frame content - rw.dec.XORKeyStream(framebuf, framebuf) - - // decode message code - content := bytes.NewReader(framebuf[:fsize]) - if err := rlp.Decode(content, &msg.Code); err != nil { - return msg, err - } - msg.Size = uint32(content.Len()) - msg.Payload = content - return msg, nil -} - -// updateMAC reseeds the given hash with encrypted seed. -// it returns the first 16 bytes of the hash sum after seeding. -func updateMAC(mac hash.Hash, block cipher.Block, seed []byte) []byte { - aesbuf := make([]byte, aes.BlockSize) - block.Encrypt(aesbuf, mac.Sum(nil)) - for i := range aesbuf { - aesbuf[i] ^= seed[i] - } - mac.Write(aesbuf) - return mac.Sum(nil)[:16] -} - -func readInt24(b []byte) uint32 { - return uint32(b[2]) | uint32(b[1])<<8 | uint32(b[0])<<16 -} - -func putInt24(v uint32, b []byte) { - b[0] = byte(v >> 16) - b[1] = byte(v >> 8) - b[2] = byte(v) -} diff --git a/p2p/rlpx_test.go b/p2p/rlpx_test.go deleted file mode 100644 index 7cc7548e2b..0000000000 --- a/p2p/rlpx_test.go +++ /dev/null @@ -1,376 +0,0 @@ -// Copyright 2015 The go-ethereum Authors -// This file is part of the go-ethereum library. -// -// The go-ethereum library is free software: you can redistribute it and/or modify -// it under the terms of the GNU Lesser General Public License as published by -// the Free Software Foundation, either version 3 of the License, or -// (at your option) any later version. -// -// The go-ethereum library is distributed in the hope that it will be useful, -// but WITHOUT ANY WARRANTY; without even the implied warranty of -// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the -// GNU Lesser General Public License for more details. -// -// You should have received a copy of the GNU Lesser General Public License -// along with the go-ethereum library. If not, see . - -package p2p - -import ( - "bytes" - "crypto/rand" - "errors" - "fmt" - "io/ioutil" - "net" - "reflect" - "strings" - "sync" - "testing" - "time" - - "github.com/davecgh/go-spew/spew" - "github.com/ethereum/go-ethereum/crypto" - "github.com/ethereum/go-ethereum/crypto/ecies" - "github.com/ethereum/go-ethereum/crypto/sha3" - "github.com/ethereum/go-ethereum/p2p/discover" - "github.com/ethereum/go-ethereum/rlp" -) - -func TestSharedSecret(t *testing.T) { - prv0, _ := crypto.GenerateKey() // = ecdsa.GenerateKey(crypto.S256(), rand.Reader) - pub0 := &prv0.PublicKey - prv1, _ := crypto.GenerateKey() - pub1 := &prv1.PublicKey - - ss0, err := ecies.ImportECDSA(prv0).GenerateShared(ecies.ImportECDSAPublic(pub1), sskLen, sskLen) - if err != nil { - return - } - ss1, err := ecies.ImportECDSA(prv1).GenerateShared(ecies.ImportECDSAPublic(pub0), sskLen, sskLen) - if err != nil { - return - } - t.Logf("Secret:\n%v %x\n%v %x", len(ss0), ss0, len(ss0), ss1) - if !bytes.Equal(ss0, ss1) { - t.Errorf("dont match :(") - } -} - -func TestEncHandshake(t *testing.T) { - for i := 0; i < 10; i++ { - start := time.Now() - if err := testEncHandshake(nil); err != nil { - t.Fatalf("i=%d %v", i, err) - } - t.Logf("(without token) %d %v\n", i+1, time.Since(start)) - } - for i := 0; i < 10; i++ { - tok := make([]byte, shaLen) - rand.Reader.Read(tok) - start := time.Now() - if err := testEncHandshake(tok); err != nil { - t.Fatalf("i=%d %v", i, err) - } - t.Logf("(with token) %d %v\n", i+1, time.Since(start)) - } -} - -func testEncHandshake(token []byte) error { - type result struct { - side string - id discover.NodeID - err error - } - var ( - prv0, _ = crypto.GenerateKey() - prv1, _ = crypto.GenerateKey() - fd0, fd1 = net.Pipe() - c0, c1 = newRLPX(fd0).(*rlpx), newRLPX(fd1).(*rlpx) - output = make(chan result) - ) - - go func() { - r := result{side: "initiator"} - defer func() { output <- r }() - defer fd0.Close() - - dest := &discover.Node{ID: discover.PubkeyID(&prv1.PublicKey)} - r.id, r.err = c0.doEncHandshake(prv0, dest) - if r.err != nil { - return - } - id1 := discover.PubkeyID(&prv1.PublicKey) - if r.id != id1 { - r.err = fmt.Errorf("remote ID mismatch: got %v, want: %v", r.id, id1) - } - }() - go func() { - r := result{side: "receiver"} - defer func() { output <- r }() - defer fd1.Close() - - r.id, r.err = c1.doEncHandshake(prv1, nil) - if r.err != nil { - return - } - id0 := discover.PubkeyID(&prv0.PublicKey) - if r.id != id0 { - r.err = fmt.Errorf("remote ID mismatch: got %v, want: %v", r.id, id0) - } - }() - - // wait for results from both sides - r1, r2 := <-output, <-output - if r1.err != nil { - return fmt.Errorf("%s side error: %v", r1.side, r1.err) - } - if r2.err != nil { - return fmt.Errorf("%s side error: %v", r2.side, r2.err) - } - - // compare derived secrets - if !reflect.DeepEqual(c0.rw.egressMAC, c1.rw.ingressMAC) { - return fmt.Errorf("egress mac mismatch:\n c0.rw: %#v\n c1.rw: %#v", c0.rw.egressMAC, c1.rw.ingressMAC) - } - if !reflect.DeepEqual(c0.rw.ingressMAC, c1.rw.egressMAC) { - return fmt.Errorf("ingress mac mismatch:\n c0.rw: %#v\n c1.rw: %#v", c0.rw.ingressMAC, c1.rw.egressMAC) - } - if !reflect.DeepEqual(c0.rw.enc, c1.rw.enc) { - return fmt.Errorf("enc cipher mismatch:\n c0.rw: %#v\n c1.rw: %#v", c0.rw.enc, c1.rw.enc) - } - if !reflect.DeepEqual(c0.rw.dec, c1.rw.dec) { - return fmt.Errorf("dec cipher mismatch:\n c0.rw: %#v\n c1.rw: %#v", c0.rw.dec, c1.rw.dec) - } - return nil -} - -func TestProtocolHandshake(t *testing.T) { - var ( - prv0, _ = crypto.GenerateKey() - node0 = &discover.Node{ID: discover.PubkeyID(&prv0.PublicKey), IP: net.IP{1, 2, 3, 4}, TCP: 33} - hs0 = &protoHandshake{Version: 3, ID: node0.ID, Caps: []Cap{{"a", 0}, {"b", 2}}} - - prv1, _ = crypto.GenerateKey() - node1 = &discover.Node{ID: discover.PubkeyID(&prv1.PublicKey), IP: net.IP{5, 6, 7, 8}, TCP: 44} - hs1 = &protoHandshake{Version: 3, ID: node1.ID, Caps: []Cap{{"c", 1}, {"d", 3}}} - - fd0, fd1 = net.Pipe() - wg sync.WaitGroup - ) - - wg.Add(2) - go func() { - defer wg.Done() - rlpx := newRLPX(fd0) - remid, err := rlpx.doEncHandshake(prv0, node1) - if err != nil { - t.Errorf("dial side enc handshake failed: %v", err) - return - } - if remid != node1.ID { - t.Errorf("dial side remote id mismatch: got %v, want %v", remid, node1.ID) - return - } - - phs, err := rlpx.doProtoHandshake(hs0) - if err != nil { - t.Errorf("dial side proto handshake error: %v", err) - return - } - if !reflect.DeepEqual(phs, hs1) { - t.Errorf("dial side proto handshake mismatch:\ngot: %s\nwant: %s\n", spew.Sdump(phs), spew.Sdump(hs1)) - return - } - rlpx.close(DiscQuitting) - }() - go func() { - defer wg.Done() - rlpx := newRLPX(fd1) - remid, err := rlpx.doEncHandshake(prv1, nil) - if err != nil { - t.Errorf("listen side enc handshake failed: %v", err) - return - } - if remid != node0.ID { - t.Errorf("listen side remote id mismatch: got %v, want %v", remid, node0.ID) - return - } - - phs, err := rlpx.doProtoHandshake(hs1) - if err != nil { - t.Errorf("listen side proto handshake error: %v", err) - return - } - if !reflect.DeepEqual(phs, hs0) { - t.Errorf("listen side proto handshake mismatch:\ngot: %s\nwant: %s\n", spew.Sdump(phs), spew.Sdump(hs0)) - return - } - - if err := ExpectMsg(rlpx, discMsg, []DiscReason{DiscQuitting}); err != nil { - t.Errorf("error receiving disconnect: %v", err) - } - }() - wg.Wait() -} - -func TestProtocolHandshakeErrors(t *testing.T) { - our := &protoHandshake{Version: 3, Caps: []Cap{{"foo", 2}, {"bar", 3}}, Name: "quux"} - id := randomID() - tests := []struct { - code uint64 - msg interface{} - err error - }{ - { - code: discMsg, - msg: []DiscReason{DiscQuitting}, - err: DiscQuitting, - }, - { - code: 0x989898, - msg: []byte{1}, - err: errors.New("expected handshake, got 989898"), - }, - { - code: handshakeMsg, - msg: make([]byte, baseProtocolMaxMsgSize+2), - err: errors.New("message too big"), - }, - { - code: handshakeMsg, - msg: []byte{1, 2, 3}, - err: newPeerError(errInvalidMsg, "(code 0) (size 4) rlp: expected input list for p2p.protoHandshake"), - }, - { - code: handshakeMsg, - msg: &protoHandshake{Version: 9944, ID: id}, - err: DiscIncompatibleVersion, - }, - { - code: handshakeMsg, - msg: &protoHandshake{Version: 3}, - err: DiscInvalidIdentity, - }, - } - - for i, test := range tests { - p1, p2 := MsgPipe() - go Send(p1, test.code, test.msg) - _, err := readProtocolHandshake(p2, our) - if !reflect.DeepEqual(err, test.err) { - t.Errorf("test %d: error mismatch: got %q, want %q", i, err, test.err) - } - } -} - -func TestRLPXFrameFake(t *testing.T) { - buf := new(bytes.Buffer) - hash := fakeHash([]byte{1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1}) - rw := newRLPXFrameRW(buf, secrets{ - AES: crypto.Sha3(), - MAC: crypto.Sha3(), - IngressMAC: hash, - EgressMAC: hash, - }) - - golden := unhex(` -00828ddae471818bb0bfa6b551d1cb42 -01010101010101010101010101010101 -ba628a4ba590cb43f7848f41c4382885 -01010101010101010101010101010101 -`) - - // Check WriteMsg. This puts a message into the buffer. - if err := Send(rw, 8, []uint{1, 2, 3, 4}); err != nil { - t.Fatalf("WriteMsg error: %v", err) - } - written := buf.Bytes() - if !bytes.Equal(written, golden) { - t.Fatalf("output mismatch:\n got: %x\n want: %x", written, golden) - } - - // Check ReadMsg. It reads the message encoded by WriteMsg, which - // is equivalent to the golden message above. - msg, err := rw.ReadMsg() - if err != nil { - t.Fatalf("ReadMsg error: %v", err) - } - if msg.Size != 5 { - t.Errorf("msg size mismatch: got %d, want %d", msg.Size, 5) - } - if msg.Code != 8 { - t.Errorf("msg code mismatch: got %d, want %d", msg.Code, 8) - } - payload, _ := ioutil.ReadAll(msg.Payload) - wantPayload := unhex("C401020304") - if !bytes.Equal(payload, wantPayload) { - t.Errorf("msg payload mismatch:\ngot %x\nwant %x", payload, wantPayload) - } -} - -type fakeHash []byte - -func (fakeHash) Write(p []byte) (int, error) { return len(p), nil } -func (fakeHash) Reset() {} -func (fakeHash) BlockSize() int { return 0 } - -func (h fakeHash) Size() int { return len(h) } -func (h fakeHash) Sum(b []byte) []byte { return append(b, h...) } - -func TestRLPXFrameRW(t *testing.T) { - var ( - aesSecret = make([]byte, 16) - macSecret = make([]byte, 16) - egressMACinit = make([]byte, 32) - ingressMACinit = make([]byte, 32) - ) - for _, s := range [][]byte{aesSecret, macSecret, egressMACinit, ingressMACinit} { - rand.Read(s) - } - conn := new(bytes.Buffer) - - s1 := secrets{ - AES: aesSecret, - MAC: macSecret, - EgressMAC: sha3.NewKeccak256(), - IngressMAC: sha3.NewKeccak256(), - } - s1.EgressMAC.Write(egressMACinit) - s1.IngressMAC.Write(ingressMACinit) - rw1 := newRLPXFrameRW(conn, s1) - - s2 := secrets{ - AES: aesSecret, - MAC: macSecret, - EgressMAC: sha3.NewKeccak256(), - IngressMAC: sha3.NewKeccak256(), - } - s2.EgressMAC.Write(ingressMACinit) - s2.IngressMAC.Write(egressMACinit) - rw2 := newRLPXFrameRW(conn, s2) - - // send some messages - for i := 0; i < 10; i++ { - // write message into conn buffer - wmsg := []interface{}{"foo", "bar", strings.Repeat("test", i)} - err := Send(rw1, uint64(i), wmsg) - if err != nil { - t.Fatalf("WriteMsg error (i=%d): %v", i, err) - } - - // read message that rw1 just wrote - msg, err := rw2.ReadMsg() - if err != nil { - t.Fatalf("ReadMsg error (i=%d): %v", i, err) - } - if msg.Code != uint64(i) { - t.Fatalf("msg code mismatch: got %d, want %d", msg.Code, i) - } - payload, _ := ioutil.ReadAll(msg.Payload) - wantPayload, _ := rlp.EncodeToBytes(wmsg) - if !bytes.Equal(payload, wantPayload) { - t.Fatalf("msg payload mismatch:\ngot %x\nwant %x", payload, wantPayload) - } - } -} diff --git a/p2p/server.go b/p2p/server.go index 7991585f17..071e0e4647 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -123,8 +123,7 @@ type Server struct { // Hooks for testing. These are useful because we can inhibit // the whole protocol stack. - newTransport func(net.Conn) transport - newPeerHook func(*Peer) + peerRunFunction func(*Peer) DiscReason lock sync.Mutex // protects running running bool @@ -157,30 +156,27 @@ const ( trustedConn ) +// used in place of devConn so server tests can substitute a mock. +type transport interface { + // devConn + doProtoHandshake(our *protoHandshake) (their *protoHandshake, err error) + close(err error) + // rlpx.Conn + Handshake() error + RemoteAddr() net.Addr + LocalAddr() net.Addr + RemoteID() *ecdsa.PublicKey +} + // conn wraps a network connection with information gathered // during the two handshakes. type conn struct { - fd net.Conn transport + id discover.NodeID flags connFlag - cont chan error // The run loop uses cont to signal errors to setupConn. - id discover.NodeID // valid after the encryption handshake - caps []Cap // valid after the protocol handshake - name string // valid after the protocol handshake -} - -type transport interface { - // The two handshakes. - doEncHandshake(prv *ecdsa.PrivateKey, dialDest *discover.Node) (discover.NodeID, error) - doProtoHandshake(our *protoHandshake) (*protoHandshake, error) - // The MsgReadWriter can only be used after the encryption - // handshake has completed. The code uses conn.id to track this - // by setting it to a non-nil value after the encryption handshake. - MsgReadWriter - // transports must provide Close because we use MsgPipe in some of - // the tests. Closing the actual network connection doesn't do - // anything in those tests because NsgPipe doesn't use it. - close(err error) + cont chan error // The run loop uses cont to signal errors to setupConn. + caps []Cap // valid after the protocol handshake + name string // valid after the protocol handshake } func (c *conn) String() string { @@ -188,7 +184,7 @@ func (c *conn) String() string { if (c.id != discover.NodeID{}) { s += fmt.Sprintf(" %x", c.id[:8]) } - s += " " + c.fd.RemoteAddr().String() + s += " " + c.RemoteAddr().String() return s } @@ -314,9 +310,6 @@ func (srv *Server) Start() (err error) { if srv.PrivateKey == nil { return fmt.Errorf("Server.PrivateKey must be set to a non-nil key") } - if srv.newTransport == nil { - srv.newTransport = newRLPX - } if srv.Dialer == nil { srv.Dialer = &net.Dialer{Timeout: defaultDialTimeout} } @@ -505,7 +498,7 @@ running: } // Disconnect all peers. for _, p := range peers { - p.Disconnect(DiscQuitting) + p.conn.close(DiscQuitting) } // Wait for peers to shut down. Pending connections and tasks are // not handled here and will terminate soon-ish because srv.quit @@ -588,7 +581,7 @@ func (srv *Server) listenLoop() { // Spawn the handler. It will give the slot back when the connection // has been established. go func() { - srv.setupConn(fd, inboundConn, nil) + srv.setupConn(newDevConn(fd, srv.PrivateKey, nil), inboundConn, nil) slots <- struct{}{} }() } @@ -597,24 +590,24 @@ func (srv *Server) listenLoop() { // setupConn runs the handshakes and attempts to add the connection // as a peer. It returns when the connection has been added as a peer // or the handshakes have failed. -func (srv *Server) setupConn(fd net.Conn, flags connFlag, dialDest *discover.Node) { +func (srv *Server) setupConn(t transport, flags connFlag, dialDest *discover.Node) { // Prevent leftover pending conns from entering the handshake. srv.lock.Lock() running := srv.running srv.lock.Unlock() - c := &conn{fd: fd, transport: srv.newTransport(fd), flags: flags, cont: make(chan error)} + c := &conn{transport: t, flags: flags, cont: make(chan error)} if !running { c.close(errServerStopped) return } // Run the encryption handshake. - var err error - if c.id, err = c.doEncHandshake(srv.PrivateKey, dialDest); err != nil { - glog.V(logger.Debug).Infof("%v faild enc handshake: %v", c, err) + if err := c.Handshake(); err != nil { + glog.V(logger.Debug).Infof("%v failed enc handshake: %v", c, err) c.close(err) return } // For dialed connections, check that the remote public key matches. + c.id = discover.PubkeyID(c.RemoteID()) if dialDest != nil && c.id != dialDest.ID { c.close(DiscUnexpectedIdentity) glog.V(logger.Debug).Infof("%v dialed identity mismatch, want %x", c, dialDest.ID[:8]) @@ -675,15 +668,17 @@ func (srv *Server) runPeer(p *Peer) { NumConnections: srv.PeerCount(), }) - if srv.newPeerHook != nil { - srv.newPeerHook(p) + var reason DiscReason + if srv.peerRunFunction != nil { + reason = srv.peerRunFunction(p) + } else { + reason = p.run() } - discreason := p.run() // Note: run waits for existing peers to be sent on srv.delpeer // before returning, so this send should not select on srv.quit. srv.delpeer <- p - glog.V(logger.Debug).Infof("Removed %v (%v)\n", p, discreason) + glog.V(logger.Debug).Infof("Removed %v (%v)\n", p, reason) srvjslog.LogJson(&logger.P2PDisconnected{ RemoteId: p.ID().String(), NumConnections: srv.PeerCount(), diff --git a/p2p/server_test.go b/p2p/server_test.go index 976d5baf59..be74e89d7e 100644 --- a/p2p/server_test.go +++ b/p2p/server_test.go @@ -26,54 +26,17 @@ import ( "time" "github.com/ethereum/go-ethereum/crypto" - "github.com/ethereum/go-ethereum/crypto/sha3" "github.com/ethereum/go-ethereum/p2p/discover" ) -func init() { - // glog.SetV(6) - // glog.SetToStderr(true) -} - -type testTransport struct { - id discover.NodeID - *rlpx - - closeErr error -} - -func newTestTransport(id discover.NodeID, fd net.Conn) transport { - wrapped := newRLPX(fd).(*rlpx) - wrapped.rw = newRLPXFrameRW(fd, secrets{ - MAC: zero16, - AES: zero16, - IngressMAC: sha3.NewKeccak256(), - EgressMAC: sha3.NewKeccak256(), - }) - return &testTransport{id: id, rlpx: wrapped} -} - -func (c *testTransport) doEncHandshake(prv *ecdsa.PrivateKey, dialDest *discover.Node) (discover.NodeID, error) { - return c.id, nil -} - -func (c *testTransport) doProtoHandshake(our *protoHandshake) (*protoHandshake, error) { - return &protoHandshake{ID: c.id, Name: "test"}, nil -} - -func (c *testTransport) close(err error) { - c.rlpx.fd.Close() - c.closeErr = err -} - -func startTestServer(t *testing.T, id discover.NodeID, pf func(*Peer)) *Server { +func startTestServer(t *testing.T, trustedNodes []*discover.Node, pf func(*Peer) DiscReason) *Server { server := &Server{ - Name: "test", - MaxPeers: 10, - ListenAddr: "127.0.0.1:0", - PrivateKey: newkey(), - newPeerHook: pf, - newTransport: func(fd net.Conn) transport { return newTestTransport(id, fd) }, + Name: "test", + MaxPeers: 10, + ListenAddr: "127.0.0.1:0", + PrivateKey: newkey(), + TrustedNodes: trustedNodes, + peerRunFunction: pf, } if err := server.Start(); err != nil { t.Fatalf("Could not start server: %v", err) @@ -84,8 +47,10 @@ func startTestServer(t *testing.T, id discover.NodeID, pf func(*Peer)) *Server { func TestServerListen(t *testing.T) { // start the test server connected := make(chan *Peer) - remid := randomID() - srv := startTestServer(t, remid, func(p *Peer) { + remkey := newkey() + remid := discover.PubkeyID(&remkey.PublicKey) + quitPeers := make(chan struct{}) + srv := startTestServer(t, nil, func(p *Peer) DiscReason { if p.ID() != remid { t.Error("peer func called with wrong node id") } @@ -93,16 +58,21 @@ func TestServerListen(t *testing.T) { t.Error("peer func called with nil conn") } connected <- p + <-quitPeers + return DiscQuitting }) defer close(connected) defer srv.Stop() + defer close(quitPeers) // dial the test server - conn, err := net.DialTimeout("tcp", srv.ListenAddr, 5*time.Second) + fd, err := net.DialTimeout("tcp", srv.ListenAddr, 5*time.Second) if err != nil { t.Fatalf("could not dial: %v", err) } - defer conn.Close() + defer fd.Close() + conn := newDevConn(fd, remkey, &srv.PrivateKey.PublicKey) + conn.doProtoHandshake(&protoHandshake{Version: baseProtocolVersion, ID: remid}) select { case peer := <-connected: @@ -120,6 +90,18 @@ func TestServerListen(t *testing.T) { } func TestServerDial(t *testing.T) { + // start the server + connected := make(chan *Peer) + quitPeers := make(chan struct{}) + srv := startTestServer(t, nil, func(p *Peer) DiscReason { + connected <- p + <-quitPeers + return DiscQuitting + }) + defer close(connected) + defer srv.Stop() + defer close(quitPeers) + // run a one-shot TCP server to handle the connection. listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { @@ -127,22 +109,19 @@ func TestServerDial(t *testing.T) { } defer listener.Close() accepted := make(chan net.Conn) + remkey := newkey() + remid := discover.PubkeyID(&remkey.PublicKey) go func() { - conn, err := listener.Accept() + fd, err := listener.Accept() if err != nil { t.Error("accept error:", err) return } - accepted <- conn + accepted <- fd + conn := newDevConn(fd, remkey, nil) + conn.doProtoHandshake(&protoHandshake{Version: baseProtocolVersion, ID: remid, Name: "test"}) }() - // start the server - connected := make(chan *Peer) - remid := randomID() - srv := startTestServer(t, remid, func(p *Peer) { connected <- p }) - defer close(connected) - defer srv.Stop() - // tell the server to connect tcpAddr := listener.Addr().(*net.TCPAddr) srv.AddPeer(&discover.Node{ID: remid, IP: tcpAddr.IP, TCP: uint16(tcpAddr.Port)}) @@ -263,21 +242,21 @@ func (t *testTask) Do(srv *Server) { // at capacity. Trusted connections should still be accepted. func TestServerAtCap(t *testing.T) { trustedID := randomID() - srv := &Server{ - PrivateKey: newkey(), - MaxPeers: 10, - NoDial: true, - TrustedNodes: []*discover.Node{{ID: trustedID}}, - } - if err := srv.Start(); err != nil { - t.Fatalf("could not start: %v", err) - } + quitPeers := make(chan struct{}) + srv := startTestServer(t, []*discover.Node{{ID: trustedID}}, func(*Peer) DiscReason { + <-quitPeers + return DiscQuitting + }) defer srv.Stop() + defer close(quitPeers) newconn := func(id discover.NodeID) *conn { - fd, _ := net.Pipe() - tx := newTestTransport(id, fd) - return &conn{fd: fd, transport: tx, flags: inboundConn, id: id, cont: make(chan error)} + return &conn{ + transport: &fakeTransport{id: id}, + flags: inboundConn, + id: id, + cont: make(chan error), + } } // Inject a few connections to fill up the peer set. @@ -300,16 +279,15 @@ func TestServerAtCap(t *testing.T) { if !c.is(trustedConn) { t.Error("Server did not set trusted flag") } - } func TestServerSetupConn(t *testing.T) { - id := randomID() - srvkey := newkey() + remkey, srvkey := newkey(), newkey() + id := discover.PubkeyID(&remkey.PublicKey) srvid := discover.PubkeyID(&srvkey.PublicKey) tests := []struct { dontstart bool - tt *setupTransport + tt *fakeTransport flags connFlag dialDest *discover.Node @@ -318,45 +296,45 @@ func TestServerSetupConn(t *testing.T) { }{ { dontstart: true, - tt: &setupTransport{id: id}, + tt: &fakeTransport{id: id}, wantCalls: "close,", wantCloseErr: errServerStopped, }, { - tt: &setupTransport{id: id, encHandshakeErr: errors.New("read error")}, + tt: &fakeTransport{id: id, encHandshakeErr: errors.New("read error")}, flags: inboundConn, wantCalls: "doEncHandshake,close,", wantCloseErr: errors.New("read error"), }, { - tt: &setupTransport{id: id}, + tt: &fakeTransport{id: id}, dialDest: &discover.Node{ID: randomID()}, flags: dynDialedConn, wantCalls: "doEncHandshake,close,", wantCloseErr: DiscUnexpectedIdentity, }, { - tt: &setupTransport{id: id, phs: &protoHandshake{ID: randomID()}}, + tt: &fakeTransport{id: id, phs: &protoHandshake{ID: randomID()}}, dialDest: &discover.Node{ID: id}, flags: dynDialedConn, wantCalls: "doEncHandshake,doProtoHandshake,close,", wantCloseErr: DiscUnexpectedIdentity, }, { - tt: &setupTransport{id: id, protoHandshakeErr: errors.New("foo")}, + tt: &fakeTransport{id: id, protoHandshakeErr: errors.New("foo")}, dialDest: &discover.Node{ID: id}, flags: dynDialedConn, wantCalls: "doEncHandshake,doProtoHandshake,close,", wantCloseErr: errors.New("foo"), }, { - tt: &setupTransport{id: srvid, phs: &protoHandshake{ID: srvid}}, + tt: &fakeTransport{id: srvid, phs: &protoHandshake{ID: srvid}}, flags: inboundConn, wantCalls: "doEncHandshake,close,", wantCloseErr: DiscSelf, }, { - tt: &setupTransport{id: id, phs: &protoHandshake{ID: id}}, + tt: &fakeTransport{id: id, phs: &protoHandshake{ID: id}}, flags: inboundConn, wantCalls: "doEncHandshake,doProtoHandshake,close,", wantCloseErr: DiscUselessPeer, @@ -365,19 +343,17 @@ func TestServerSetupConn(t *testing.T) { for i, test := range tests { srv := &Server{ - PrivateKey: srvkey, - MaxPeers: 10, - NoDial: true, - Protocols: []Protocol{discard}, - newTransport: func(fd net.Conn) transport { return test.tt }, + PrivateKey: srvkey, + MaxPeers: 10, + NoDial: true, + Protocols: []Protocol{discard}, } if !test.dontstart { if err := srv.Start(); err != nil { t.Fatalf("couldn't start server: %v", err) } } - p1, _ := net.Pipe() - srv.setupConn(p1, test.flags, test.dialDest) + srv.setupConn(test.tt, test.flags, test.dialDest) if !reflect.DeepEqual(test.tt.closeErr, test.wantCloseErr) { t.Errorf("test %d: close error mismatch: got %q, want %q", i, test.tt.closeErr, test.wantCloseErr) } @@ -387,7 +363,7 @@ func TestServerSetupConn(t *testing.T) { } } -type setupTransport struct { +type fakeTransport struct { id discover.NodeID encHandshakeErr error @@ -398,28 +374,30 @@ type setupTransport struct { closeErr error } -func (c *setupTransport) doEncHandshake(prv *ecdsa.PrivateKey, dialDest *discover.Node) (discover.NodeID, error) { +func (c *fakeTransport) Handshake() error { c.calls += "doEncHandshake," - return c.id, c.encHandshakeErr + return c.encHandshakeErr } -func (c *setupTransport) doProtoHandshake(our *protoHandshake) (*protoHandshake, error) { +func (c *fakeTransport) doProtoHandshake(our *protoHandshake) (*protoHandshake, error) { c.calls += "doProtoHandshake," if c.protoHandshakeErr != nil { return nil, c.protoHandshakeErr } return c.phs, nil } -func (c *setupTransport) close(err error) { +func (c *fakeTransport) close(err error) { c.calls += "close," c.closeErr = err } - -// setupConn shouldn't write to/read from the connection. -func (c *setupTransport) WriteMsg(Msg) error { - panic("WriteMsg called on setupTransport") +func (c *fakeTransport) RemoteID() *ecdsa.PublicKey { + key, _ := c.id.Pubkey() + return key } -func (c *setupTransport) ReadMsg() (Msg, error) { - panic("ReadMsg called on setupTransport") +func (c *fakeTransport) RemoteAddr() net.Addr { + return &net.TCPAddr{Port: 33, IP: net.IP{0, 0, 0, 1}} +} +func (c *fakeTransport) LocalAddr() net.Addr { + return &net.TCPAddr{Port: 44, IP: net.IP{0, 0, 0, 2}} } func newkey() *ecdsa.PrivateKey {