diff --git a/p2p/peer.go b/p2p/peer.go index a14a7f589e..87aaa10afe 100644 --- a/p2p/peer.go +++ b/p2p/peer.go @@ -88,12 +88,12 @@ type Peer struct { // These fields are kept so base protocol can access them. // TODO: this should be one or more interfaces - hash []byte // hash of pubkey used as address - ourID ClientIdentity // client id of the Server - ourListenAddr *peerAddr // listen addr of Server, nil if not listening - addPeer func(*peerAddr) error // tell server about received peers - getPeers func(...[]byte) []*peerAddr // should return the list of all peers - pubkeyHook func(*peerAddr) error // called at end of handshake to validate pubkey + hash []byte // hash of pubkey used as address + ourID ClientIdentity // client id of the Server + ourListenAddr *peerAddr // listen addr of Server, nil if not listening + addPeer func(*peerAddr) error // tell server about received peers + getPeers func(...[]byte) []*peerAddr // should return the list of all peers + verifyPeerHook func(*Peer) error // called at end of handshake to validate peer } // NewPeer returns a peer for testing purposes. @@ -152,12 +152,13 @@ func (self *Peer) Hash() []byte { func (self *Peer) Pubkey() (pubkey []byte) { self.infolock.Lock() defer self.infolock.Unlock() - if self.dialAddr != nil { + switch { + case self.identity != nil: + pubkey = self.identity.Pubkey() + case self.dialAddr != nil: pubkey = self.dialAddr.Pubkey - } else { - if self.listenAddr != nil { - pubkey = self.listenAddr.Pubkey - } + case self.listenAddr != nil: + pubkey = self.listenAddr.Pubkey } return } diff --git a/p2p/peer_error.go b/p2p/peer_error.go index 0eb7ec838d..73587ddca4 100644 --- a/p2p/peer_error.go +++ b/p2p/peer_error.go @@ -14,7 +14,11 @@ const ( errP2PVersionMismatch errPubkeyMissing errPubkeyInvalid - errPubkeyForbidden + errPubkeyMismatch + errBlacklistedPeer + errSelfConnection + errConnectedPeer + errRejectedPeer errProtocolBreach errPingTimeout errInvalidNetworkId @@ -31,7 +35,11 @@ var errorToString = map[int]string{ errP2PVersionMismatch: "P2P Version Mismatch", errPubkeyMissing: "Public key missing", errPubkeyInvalid: "Public key invalid", - errPubkeyForbidden: "Public key forbidden", + errPubkeyMismatch: "Public key mismatch", + errBlacklistedPeer: "Blacklisted peer", + errSelfConnection: "Self connection", + errConnectedPeer: "Connected peer", + errRejectedPeer: "Rejected peer", errProtocolBreach: "Protocol Breach", errPingTimeout: "Ping timeout", errInvalidNetworkId: "Invalid network id", @@ -117,9 +125,9 @@ func discReasonForError(err error) DiscReason { switch peerError.Code { case errP2PVersionMismatch: return DiscIncompatibleVersion - case errPubkeyMissing, errPubkeyInvalid: + case errPubkeyMissing, errPubkeyMismatch, errPubkeyInvalid: return DiscInvalidIdentity - case errPubkeyForbidden: + case errBlacklistedPeer, errSelfConnection, errConnectedPeer, errRejectedPeer: return DiscUselessPeer case errInvalidMsgCode, errMagicTokenMismatch, errProtocolBreach: return DiscProtocolError diff --git a/p2p/peer_test.go b/p2p/peer_test.go index 8c04c137d7..c8ccf25020 100644 --- a/p2p/peer_test.go +++ b/p2p/peer_test.go @@ -34,7 +34,7 @@ func testPeer(protos []Protocol) (net.Conn, *Peer, <-chan error) { peer.init(conn1) peer.protocols = protos peer.ourID = &peerId{} - peer.pubkeyHook = func(*peerAddr) error { return nil } + peer.verifyPeerHook = func(*Peer) error { return nil } errc := make(chan error, 1) go func() { _, err := peer.loop() diff --git a/p2p/protocol.go b/p2p/protocol.go index 7aabe2860d..908cdc8161 100644 --- a/p2p/protocol.go +++ b/p2p/protocol.go @@ -249,13 +249,9 @@ func (bp *baseProtocol) readHandshake() error { // verify that the peer we wanted to connect to // actually holds the target public key. if da.Pubkey != nil && !bytes.Equal(da.Pubkey, hs.NodeID) { - return newPeerError(errPubkeyForbidden, "dial address pubkey mismatch") + return newPeerError(errPubkeyMismatch, "dial address pubkey mismatch: %x vs %x", da.Pubkey, hs.NodeID) } } - pa := newPeerAddr(bp.peer.conn.RemoteAddr(), hs.NodeID) - if err := bp.peer.pubkeyHook(pa); err != nil { - return newPeerError(errPubkeyForbidden, "%v", err) - } // TODO: remove Caps with empty name var addr *peerAddr if hs.ListenPort != 0 { @@ -263,6 +259,9 @@ func (bp *baseProtocol) readHandshake() error { addr.Port = hs.ListenPort } bp.peer.setHandshakeInfo(&hs, addr, hs.Caps) + if err := bp.peer.verifyPeerHook(bp.peer); err != nil { + return err + } bp.peer.startSubprotocols(hs.Caps) return nil } diff --git a/p2p/protocol_test.go b/p2p/protocol_test.go index 052f1d10a7..6a71e30add 100644 --- a/p2p/protocol_test.go +++ b/p2p/protocol_test.go @@ -29,7 +29,7 @@ func (self *peerId) Pubkey() (pubkey []byte) { func newTestPeer() (peer *Peer) { peer = NewPeer(&peerId{}, []Cap{}) - peer.pubkeyHook = func(*peerAddr) error { return nil } + peer.verifyPeerHook = func(*Peer) error { return nil } peer.ourID = &peerId{} peer.listenAddr = &peerAddr{} peer.getPeers = func(...[]byte) []*peerAddr { return nil } @@ -110,7 +110,7 @@ func TestBaseProtocolPeers(t *testing.T) { func TestBaseProtocolDisconnect(t *testing.T) { peer := NewPeer(&peerId{}, nil) peer.ourID = &peerId{} - peer.pubkeyHook = func(*peerAddr) error { return nil } + peer.verifyPeerHook = func(*Peer) error { return nil } rw1, rw2 := MsgPipe() done := make(chan struct{}) diff --git a/p2p/server.go b/p2p/server.go index a9d1d3989b..35dca03c1d 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -189,7 +189,7 @@ func (srv *Server) connect(p *Peer, conn net.Conn) { p.ourID = srv.Identity p.addPeer = srv.AddPeer p.getPeers = srv.GetPeers - p.pubkeyHook = srv.verifyPeer + p.verifyPeerHook = srv.verifyPeer p.runBaseProtocol = true p.protocols = srv.Protocols @@ -427,20 +427,21 @@ func (srv *Server) removePeer(peer *Peer) { srv.peerSlots <- peer.slot } -func (srv *Server) verifyPeer(addr *peerAddr) error { - if srv.Blacklist.Exists(addr.Pubkey) { - return errors.New("blacklisted") +func (srv *Server) verifyPeer(peer *Peer) error { + pubkey := peer.Pubkey() + if srv.Blacklist.Exists(pubkey) { + return newPeerError(errBlacklistedPeer, "") } - if bytes.Equal(srv.Identity.Pubkey()[1:], addr.Pubkey) { - return newPeerError(errPubkeyForbidden, "not allowed to connect to srv") + if bytes.Equal(srv.Identity.Pubkey()[1:], pubkey) { + return newPeerError(errSelfConnection, "") } srv.lock.RLock() defer srv.lock.RUnlock() for _, peer := range srv.peers { if peer != nil { id := peer.Identity() - if id != nil && bytes.Equal(id.Pubkey(), addr.Pubkey) { - return errors.New("already connected") + if id != nil && bytes.Equal(id.Pubkey(), pubkey) { + return newPeerError(errConnectedPeer, "") } } }