diff --git a/p2p/protocols/protocol.go b/p2p/protocols/protocol.go index bbe4198455..555cab6710 100644 --- a/p2p/protocols/protocol.go +++ b/p2p/protocols/protocol.go @@ -228,7 +228,7 @@ func (self *Peer) Send(msg interface{}) error { if !found { return errorf(ErrInvalidMsgType, "%v", code) } - log.Trace(fmt.Sprintf("=> msg #%d TO %v : %v", code, self.ID(), msg)) + log.Trace(fmt.Sprintf("=> msg %s#%d TO %v : %v", self.spec.Name, code, self.ID(), msg)) return p2p.Send(self.rw, code, msg) } @@ -257,7 +257,7 @@ func (self *Peer) handleIncoming(handle func(msg interface{}) error) error { if err := msg.Decode(val); err != nil { return errorf(ErrDecode, "<= %v: %v", msg, err) } - log.Trace(fmt.Sprintf("<= %v FROM %v %T %v", msg, self.ID(), val, val)) + log.Trace(fmt.Sprintf("<= %s/%v FROM %v %T %v", self.spec.Name, msg, self.ID(), val, val)) // call the registered handler callbacks // a registered callback take the decoded message as argument as an interface @@ -280,29 +280,24 @@ func (self *Peer) Handshake(ctx context.Context, hs interface{}) (interface{}, e } errc := make(chan error, 2) go func() { - if err := self.Send(hs); err != nil { - errc <- errorf(ErrHandshake, "cannot send: %v", err) - } + errc <- self.Send(hs) }() - hsc := make(chan interface{}) + var rhs interface{} go func() { - var rhs interface{} - err := self.handleIncoming(func(msg interface{}) error { + errc <- self.handleIncoming(func(msg interface{}) error { rhs = msg return nil }) - if err != nil { - errc <- err - return - } - hsc <- rhs }() - select { - case rhs := <-hsc: - return rhs, nil - case <-ctx.Done(): - return nil, ctx.Err() - case err := <-errc: - return nil, err + for i := 0; i < 2; i++ { + select { + case err := <-errc: + if err != nil { + return nil, errorf(ErrHandshake, err.Error()) + } + case <-ctx.Done(): + return nil, ctx.Err() + } } + return rhs, nil } diff --git a/p2p/server.go b/p2p/server.go index eb7c70e593..3127d5aa04 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -548,7 +548,11 @@ running: c.flags |= trustedConn } // TODO: track in-progress inbound node IDs (pre-Peer) to avoid dialing them. - c.cont <- srv.encHandshakeChecks(peers, c) + select { + case c.cont <- srv.encHandshakeChecks(peers, c): + case <-srv.quit: + break running + } case c := <-srv.addpeer: // At this point the connection is past the protocol handshake. // Its capabilities are known and the remote identity is verified. @@ -569,7 +573,11 @@ running: // The dialer logic relies on the assumption that // dial tasks complete after the peer has been added or // discarded. Unblock the task last. - c.cont <- err + select { + case c.cont <- err: + case <-srv.quit: + break running + } case pd := <-srv.delpeer: // A peer disconnected. d := common.PrettyDuration(mclock.Now() - pd.created) diff --git a/p2p/simulations/adapters/inproc.go b/p2p/simulations/adapters/inproc.go index f81611a470..3e69728d98 100644 --- a/p2p/simulations/adapters/inproc.go +++ b/p2p/simulations/adapters/inproc.go @@ -74,14 +74,29 @@ func (s *SimAdapter) NewNode(config *NodeConfig) (Node, error) { } } - node := &SimNode{ + n, err := node.New(&node.Config{ + P2P: p2p.Config{ + PrivateKey: config.PrivateKey, + MaxPeers: math.MaxInt32, + NoDiscovery: true, + Dialer: s, + EnableMsgEvents: false, + }, + NoUSB: true, + }) + if err != nil { + return nil, err + } + + simNode := &SimNode{ ID: id, config: config, + node: n, adapter: s, running: make(map[string]node.Service), } - s.nodes[id] = node - return node, nil + s.nodes[id] = simNode + return simNode, nil } func (s *SimAdapter) Dial(dest *discover.Node) (conn net.Conn, err error) { @@ -99,17 +114,11 @@ func (s *SimAdapter) Dial(dest *discover.Node) (conn net.Conn, err error) { } func (s *SimAdapter) DialRPC(id discover.NodeID) (*rpc.Client, error) { - simNode, ok := s.GetNode(id) + node, ok := s.GetNode(id) if !ok { return nil, fmt.Errorf("unknown node: %s", id) } - simNode.lock.RLock() - node := simNode.node - simNode.lock.RUnlock() - if node == nil { - return nil, errors.New("node not started") - } - handler, err := node.RPCHandler() + handler, err := node.node.RPCHandler() if err != nil { return nil, err } @@ -164,13 +173,7 @@ func (self *SimNode) Client() (*rpc.Client, error) { // ServeRPC serves RPC requests over the given connection by creating an // in-process client to the node's RPC server func (self *SimNode) ServeRPC(conn net.Conn) error { - self.lock.RLock() - node := self.node - self.lock.RUnlock() - if node == nil { - return errors.New("node not started") - } - handler, err := node.RPCHandler() + handler, err := self.node.RPCHandler() if err != nil { return err } @@ -181,9 +184,12 @@ func (self *SimNode) ServeRPC(conn net.Conn) error { // Snapshots creates snapshots of the services by calling the // simulation_snapshot RPC method func (self *SimNode) Snapshots() (map[string][]byte, error) { - self.lock.Lock() - services := self.running - self.lock.Unlock() + self.lock.RLock() + services := make(map[string]node.Service, len(self.running)) + for name, service := range self.running { + services[name] = service + } + self.lock.RUnlock() if len(services) == 0 { return nil, errors.New("no running services") } @@ -206,10 +212,6 @@ func (self *SimNode) Snapshots() (map[string][]byte, error) { func (self *SimNode) Start(snapshots map[string][]byte) error { self.lock.Lock() defer self.lock.Unlock() - if self.node != nil { - return errors.New("node already started") - } - newService := func(name string) func(ctx *node.ServiceContext) (node.Service, error) { return func(nodeCtx *node.ServiceContext) (node.Service, error) { ctx := &ServiceContext{ @@ -230,59 +232,34 @@ func (self *SimNode) Start(snapshots map[string][]byte) error { } } - node, err := node.New(&node.Config{ - P2P: p2p.Config{ - PrivateKey: self.config.PrivateKey, - MaxPeers: math.MaxInt32, - NoDiscovery: true, - Dialer: self.adapter, - EnableMsgEvents: false, - }, - NoUSB: true, - }) - if err != nil { - return err - } - for _, name := range self.config.Services { - if err := node.Register(newService(name)); err != nil { + if err := self.node.Register(newService(name)); err != nil { return err } } - if err := node.Start(); err != nil { + if err := self.node.Start(); err != nil { return err } // create an in-process RPC client - handler, err := node.RPCHandler() + handler, err := self.node.RPCHandler() if err != nil { return err } self.client = rpc.DialInProc(handler) - self.node = node - return nil } func (self *SimNode) Stop() error { - self.lock.Lock() - defer self.lock.Unlock() - if self.node == nil { - return nil - } - if err := self.node.Stop(); err != nil { - return err - } - self.node = nil - return nil + return self.node.Stop() } // Services returns the underlying services func (self *SimNode) Services() []node.Service { - self.lock.Lock() - defer self.lock.Unlock() + self.lock.RLock() + defer self.lock.RUnlock() services := make([]node.Service, 0, len(self.running)) for _, service := range self.running { services = append(services, service) @@ -291,11 +268,6 @@ func (self *SimNode) Services() []node.Service { } func (self *SimNode) Server() *p2p.Server { - self.lock.Lock() - defer self.lock.Unlock() - if self.node == nil { - return nil - } return self.node.Server() } diff --git a/swarm/network/kademlia.go b/swarm/network/kademlia.go index 0205c8404a..89cc6b0ffd 100644 --- a/swarm/network/kademlia.go +++ b/swarm/network/kademlia.go @@ -595,7 +595,7 @@ func (k *Kademlia) gotNearestNeighbours(peers [][]byte) (got bool) { for _, p := range peers { pm[string(p)] = true } - k.EachConn(nil, 255, func(p OverlayConn, po int, nn bool) bool { + k.eachConn(nil, 255, func(p OverlayConn, po int, nn bool) bool { if !nn { return false }