diff --git a/p2p/dial.go b/p2p/dial.go index b832fc233e..10f16f5166 100644 --- a/p2p/dial.go +++ b/p2p/dial.go @@ -47,6 +47,19 @@ const ( maxResolveDelay = time.Hour ) +type NodeDialer interface { + Dial(*discover.Node) (net.Conn, error) +} + +type TCPDialer struct { + *net.Dialer +} + +func (t TCPDialer) Dial(dest *discover.Node) (net.Conn, error) { + addr := &net.TCPAddr{IP: dest.IP, Port: int(dest.TCP)} + return t.Dialer.Dial("tcp", addr.String()) +} + // dialstate schedules dials and discovery lookups. // it get's a chance to compute new tasks on every iteration // of the main loop in server.run. @@ -318,14 +331,13 @@ func (t *dialTask) resolve(srv *server) bool { // dial performs the actual connection attempt. func (t *dialTask) dial(srv *server, dest *discover.Node) bool { - addr := &net.TCPAddr{IP: dest.IP, Port: int(dest.TCP)} - fd, err := srv.Dialer.Dial("tcp", addr.String()) + fd, err := srv.Dialer.Dial(dest) if err != nil { log.Trace("Dial error", "task", t, "err", err) return false } mfd := newMeteredConn(fd, false) - srv.setupConn(mfd, t.flags, dest) + srv.SetupConn(mfd, t.flags, dest) return true } diff --git a/p2p/dial_test.go b/p2p/dial_test.go index 55c0154f80..ad18ef9abe 100644 --- a/p2p/dial_test.go +++ b/p2p/dial_test.go @@ -597,8 +597,8 @@ func TestDialResolve(t *testing.T) { } // Now run the task, it should resolve the ID once. - config := Config{Dialer: &net.Dialer{Deadline: time.Now().Add(-5 * time.Minute)}} - srv := &server{ntab: table, Config: config} + config := Config{Dialer: TCPDialer{&net.Dialer{Deadline: time.Now().Add(-5 * time.Minute)}}} + srv := &Server{ntab: table, Config: config} tasks[0].Do(srv) if !reflect.DeepEqual(table.resolveCalls, []discover.NodeID{dest.ID}) { t.Fatalf("wrong resolve calls, got %v", table.resolveCalls) diff --git a/p2p/protocols/protocol.go b/p2p/protocols/protocol.go index 12fce6b93d..bbe4198455 100644 --- a/p2p/protocols/protocol.go +++ b/p2p/protocols/protocol.go @@ -30,13 +30,13 @@ Standard protocol supports: package protocols import ( + "context" "fmt" "reflect" - "time" + "sync" "github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/p2p" - "github.com/ethereum/go-ethereum/p2p/discover" ) // error codes used by this protocol scheme @@ -109,155 +109,104 @@ func errorf(code int, format string, params ...interface{}) *Error { return self } -// implements the code table spec -// listing the message codes and types etc -// and further metadata about the protocol -type CodeMap struct { - Name string // name of the protocol - Version uint // version - MaxMsgSize int // max length of message payload size - codepos int // the subsequent code - codes map[uint64]reflect.Type // index of codes to msg types - to create zero values - messages map[reflect.Type]uint64 // index of types to codes, for sending by type +// Spec is a protocol specification including its name and version as well as +// the types of messages which are exchanged +type Spec struct { + // Name is the name of the protocol, often a three-letter word + Name string + + // Version is the version number of the protocol + Version uint + + // MaxMsgSize is the maximum accepted length of the message payload + MaxMsgSize uint32 + + // Messages is a list of message types which this protocol uses, with + // each message type being sent with its array index as the code (so + // [&foo{}, &bar{}, &baz{}] would send foo, bar and baz with codes + // 0, 1 and 2 respectively) + Messages []interface{} + + initOnce sync.Once + codes map[reflect.Type]uint64 + types map[uint64]reflect.Type } -func (self *CodeMap) GetInterface(code uint64) (interface{}, bool) { - typ, found := self.codes[code] - if !found { +func (s *Spec) init() { + s.initOnce.Do(func() { + s.codes = make(map[reflect.Type]uint64, len(s.Messages)) + s.types = make(map[uint64]reflect.Type, len(s.Messages)) + for i, msg := range s.Messages { + code := uint64(i) + typ := reflect.TypeOf(msg) + if typ.Kind() == reflect.Ptr { + typ = typ.Elem() + } + s.codes[typ] = code + s.types[code] = typ + } + }) +} + +func (s *Spec) Length() uint64 { + return uint64(len(s.Messages)) +} + +func (s *Spec) GetCode(msg interface{}) (uint64, bool) { + s.init() + typ := reflect.TypeOf(msg) + if typ.Kind() == reflect.Ptr { + typ = typ.Elem() + } + code, ok := s.codes[typ] + return code, ok +} + +func (s *Spec) NewMsg(code uint64) (interface{}, bool) { + s.init() + typ, ok := s.types[code] + if !ok { return nil, false } - val := reflect.New(typ) - return val.Interface(), true -} - -func (self *CodeMap) GetCode(msg interface{}) (uint64, bool) { - code, found := self.messages[reflect.TypeOf(msg)] - return code, found -} - -// NewCodeMap construct the code to type map for the protocol -func NewCodeMap(name string, version uint, maxMsgSize int) *CodeMap { - return &CodeMap{ - Name: name, - Version: version, - MaxMsgSize: maxMsgSize, - messages: make(map[reflect.Type]uint64), - codes: make(map[uint64]reflect.Type), - } -} - -// Length returns the current highes codepos + 1 -func (self *CodeMap) Length() uint64 { - return uint64(self.codepos) -} - -// Register defines a new series of codes starting on series, incrementing -func (self *CodeMap) Register(series int, msgs ...interface{}) { - self.codepos = series - for _, msg := range msgs { - typ := reflect.TypeOf(msg) - _, found := self.messages[typ] - if found { - // ignore duplicates - continue - } - // next code assigned to message type typ - self.messages[typ] = uint64(self.codepos) - self.codes[uint64(self.codepos)] = typ - self.codepos++ - } -} - -func NewProtocol(protocolname string, protocolversion uint, run func(*Peer) error, ct *CodeMap, peerInfo func(id discover.NodeID) interface{}, nodeInfo func() interface{}) *p2p.Protocol { - - // PeerInfo is an optional helper method to retrieve protocol specific metadata - // about a certain peer in the network. If an info retrieval function is set, - // but returns nil, it is assumed that the protocol handshake is still running. - r := func(p *p2p.Peer, rw p2p.MsgReadWriter) error { - return run(NewPeer(p, ct, rw)) - - } - - return &p2p.Protocol{ - Name: protocolname, - Version: protocolversion, - Length: ct.Length(), - Run: r, - PeerInfo: peerInfo, - NodeInfo: nodeInfo, - } -} - -type Disconnect struct { - err error + return reflect.New(typ).Interface(), true } // A Peer represents a remote peer or protocol instance that is running on a peer connection with // a remote peer type Peer struct { - ct *CodeMap // CodeMap for the protocol - *p2p.Peer // the p2p.Peer object representing the remote - rw p2p.MsgReadWriter // p2p.MsgReadWriter to send messages to and read messages from - handlers map[reflect.Type][]func(interface{}) error // message type -> message handler callback(s) map + *p2p.Peer // the p2p.Peer object representing the remote + rw p2p.MsgReadWriter // p2p.MsgReadWriter to send messages to and read messages from + spec *Spec Errc chan error - ready chan bool // blocking send until handshake finishes + wErrc chan error // write error channel } // NewPeer returns a new peer // this constructor is called by the p2p.Protocol#Run function // the first two arguments are comming the arguments passed to p2p.Protocol.Run function // the third argument is the CodeMap describing the protocol messages and options -func NewPeer(p *p2p.Peer, ct *CodeMap, rw p2p.MsgReadWriter) *Peer { - ready := make(chan bool) - defer close(ready) +func NewPeer(p *p2p.Peer, rw p2p.MsgReadWriter, spec *Spec) *Peer { return &Peer{ - ct: ct, - Peer: p, - rw: rw, - Errc: make(chan error), - ready: ready, - handlers: make(map[reflect.Type][]func(interface{}) error), + Peer: p, + rw: rw, + spec: spec, + Errc: make(chan error), + wErrc: make(chan error), } } -// Register is called on the peer typically within the constructor of service instances running on peer connections -// These constructors are called by the p2p.Protocol#Run function -// It ties handler callbackss for specific message types -// A message type can have several handlers registered by the same or different protocol services -// Register is meant to be called once, deregistering is not currently supported therefore -// handlers are assumed to be static across handshake renegotiations -// i.e., a service instance either handles a message or not (irrespective of the handshake) -// it panics if the message type is not defined in the CodeMap -func (self *Peer) Register(msg interface{}, handler func(interface{}) error) uint64 { - typ := reflect.TypeOf(msg) - code, found := self.ct.messages[typ] - if !found { - panic(fmt.Sprintf("message type '%v' unknown ", typ)) - } - log.Trace(fmt.Sprintf("register handle for %v", typ)) - self.handlers[typ] = append(self.handlers[typ], handler) - return code -} - // Run starts the forever loop that handles incoming messages // called within the p2p.Protocol#Run function -func (self *Peer) Run() error { +func (self *Peer) Run(handler func(msg interface{}) error) error { go func() { for { - _, err := self.handleIncoming() - if err != nil { + if err := self.handleIncoming(handler); err != nil { self.Errc <- err return } } }() - err := <-self.Errc - d := &Disconnect{err} - for _, f := range self.handlers[reflect.TypeOf(d)] { - log.Trace(fmt.Sprintf("disconnect hook for %v", d)) - f(err) - } - return err + return <-self.Errc } // Drop disconnects a peer. @@ -275,26 +224,12 @@ func (self *Peer) Drop(err error) { // this low level call will be wrapped by libraries providing routed or broadcast sends // but often just used to forward and push messages to directly connected peers func (self *Peer) Send(msg interface{}) error { - <-self.ready - return self.send(msg) -} - -func (self *Peer) send(msg interface{}) error { - code, found := self.ct.GetCode(msg) + code, found := self.spec.GetCode(msg) if !found { return errorf(ErrInvalidMsgType, "%v", code) } log.Trace(fmt.Sprintf("=> msg #%d TO %v : %v", code, self.ID(), msg)) - - return p2p.Send(self.rw, uint64(code), msg) -} - -func (self *Peer) DisconnectHook(f func(error)) { - typ := reflect.TypeOf(&Disconnect{}) - self.handlers[typ] = append(self.handlers[typ], func(e interface{}) error { - f(e.(error)) - return nil - }) + return p2p.Send(self.rw, code, msg) } // handleIncoming(code) @@ -302,85 +237,72 @@ func (self *Peer) DisconnectHook(f func(error)) { // if this returns an error the loop returns and the peer is disconnected with the error // checks message size, out-of-range message codes, handles decoding with reflection, // call handlers as callback onside -func (self *Peer) handleIncoming() (interface{}, error) { +func (self *Peer) handleIncoming(handle func(msg interface{}) error) error { msg, err := self.rw.ReadMsg() if err != nil { - return nil, err + return err } log.Trace(fmt.Sprintf("<= %v", msg)) // make sure that the payload has been fully consumed defer msg.Discard() - if msg.Size > uint32(self.ct.MaxMsgSize) { - return nil, errorf(ErrMsgTooLong, "%v > %v", msg.Size, self.ct.MaxMsgSize) + if msg.Size > self.spec.MaxMsgSize { + return errorf(ErrMsgTooLong, "%v > %v", msg.Size, self.spec.MaxMsgSize) } - // check if the message code is correct - maxMsgCode := uint(len(self.ct.messages)) - if msg.Code >= uint64(maxMsgCode) { - return nil, errorf(ErrInvalidMsgCode, "%v (>=%v)", msg.Code, maxMsgCode) + val, ok := self.spec.NewMsg(msg.Code) + if !ok { + return errorf(ErrInvalidMsgCode, "%v", msg.Code) } - - // it is safe to be unsafe here - typ := self.ct.codes[msg.Code] - val := reflect.New(typ) - req := val.Elem() - req.Set(reflect.Zero(typ)) - if err := msg.Decode(val.Interface()); err != nil { - return nil, errorf(ErrDecode, "<= %v: %v", msg, err) + if err := msg.Decode(val); err != nil { + return errorf(ErrDecode, "<= %v: %v", msg, err) } - log.Trace(fmt.Sprintf("<= %v FROM %v %v %v", msg, self.ID(), req, typ)) + log.Trace(fmt.Sprintf("<= %v FROM %v %T %v", msg, self.ID(), val, val)) // call the registered handler callbacks // a registered callback take the decoded message as argument as an interface // which the handler is supposed to cast to the appropriate type // it is entirely safe not to check the cast in the handler since the handler is // chosen based on the proper type in the first place - handlers := self.handlers[typ] - if len(handlers) == 0 { - log.Trace(fmt.Sprintf("no handler (msg code %v)", msg.Code)) - // return nil, errorf(ErrNoHandler, "(msg code %v)", msg.Code) - } else { - for i, f := range handlers { - log.Trace(fmt.Sprintf("handler %v for %v", i, typ)) - err = f(req.Interface()) - if err != nil { - return nil, errorf(ErrHandler, "(msg code %v): %v", msg.Code, err) - } - } + if err := handle(val); err != nil { + return errorf(ErrHandler, "(msg code %v): %v", msg.Code, err) } - return req.Interface(), nil + return nil } // Handshake initiates a handshake on the peer connection // * the argument is the local handshake to be sent to the remote peer // * expects a remote handshake back of the same type // returns the remote hs and an error -func (self *Peer) Handshake(hs interface{}, handshakeTimeout time.Duration) (rhs interface{}, err error) { - typ := reflect.TypeOf(hs) - _, found := self.ct.messages[typ] - if !found { - return nil, errorf(ErrHandshake, "unknown handshake message type: %v", typ) +func (self *Peer) Handshake(ctx context.Context, hs interface{}) (interface{}, error) { + if _, ok := self.spec.GetCode(hs); !ok { + return nil, errorf(ErrHandshake, "unknown handshake message type: %T", hs) } - self.ready = make(chan bool) - received := make(chan bool) - defer close(self.ready) + errc := make(chan error, 2) go func() { - defer close(received) - // receiving and validating remote handshake, expect code - rhs, err = self.handleIncoming() - if err != nil { - err = errorf(ErrHandshake, "'%v': %v", self.ct.Name, err) + if err := self.Send(hs); err != nil { + errc <- errorf(ErrHandshake, "cannot send: %v", err) } }() - if e := self.send(hs); e != nil { - return nil, errorf(ErrHandshake, "cannot send: %v", e) - } - + hsc := make(chan interface{}) + go func() { + var rhs interface{} + err := self.handleIncoming(func(msg interface{}) error { + rhs = msg + return nil + }) + if err != nil { + errc <- err + return + } + hsc <- rhs + }() select { - case <-received: - case <-time.NewTimer(handshakeTimeout).C: - err = errorf(ErrHandshake, "timeout after %v", handshakeTimeout) + case rhs := <-hsc: + return rhs, nil + case <-ctx.Done(): + return nil, ctx.Err() + case err := <-errc: + return nil, err } - return rhs, err } diff --git a/p2p/protocols/protocol_test.go b/p2p/protocols/protocol_test.go index fc1238c5f7..e4f7a145a1 100644 --- a/p2p/protocols/protocol_test.go +++ b/p2p/protocols/protocol_test.go @@ -1,6 +1,8 @@ package protocols import ( + "context" + "errors" "fmt" "os" "testing" @@ -13,7 +15,7 @@ import ( ) func init() { - log.Root().SetHandler(log.LvlFilterHandler(log.LvlError, log.StreamHandler(os.Stderr, log.TerminalFormat(false)))) + log.Root().SetHandler(log.LvlFilterHandler(log.LvlTrace, log.StreamHandler(os.Stderr, log.TerminalFormat(false)))) } // handshake message type @@ -55,28 +57,26 @@ const networkId = "420" // newProtocol sets up a protocol // the run function here demonstrates a typical protocol using peerPool, handshake // and messages registered to handlers -func newProtocol(pp *p2ptest.TestPeerPool) adapters.RunProtocol { - ct := NewCodeMap("test", 42, 1024) - ct.Register(0, &protoHandshake{}, &hs0{}, &kill{}, &drop{}) +func newProtocol(pp *p2ptest.TestPeerPool) func(*p2p.Peer, p2p.MsgReadWriter) error { + spec := &Spec{ + Name: "test", + Version: 42, + MaxMsgSize: 10 * 1024, + Messages: []interface{}{ + protoHandshake{}, + hs0{}, + kill{}, + drop{}, + }, + } return func(p *p2p.Peer, rw p2p.MsgReadWriter) error { - peer := NewPeer(p, ct, rw) - - // demonstrates use of peerPool, killing another peer connection as a response to a message - peer.Register(&kill{}, func(msg interface{}) error { - id := msg.(*kill).C - pp.Get(id).Drop(fmt.Errorf("killed")) - log.Trace(fmt.Sprintf("id %v killed", id)) - return nil - }) - - // for testing we can trigger self induced disconnect upon receiving drop message - peer.Register(&drop{}, func(msg interface{}) error { - return fmt.Errorf("dropped") - }) + peer := NewPeer(p, rw, spec) // initiate one-off protohandshake and check validity - phs := &protoHandshake{ct.Version, networkId} - hs, err := peer.Handshake(phs, time.Second) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + phs := &protoHandshake{42, networkId} + hs, err := peer.Handshake(ctx, phs) if err != nil { return err } @@ -88,7 +88,7 @@ func newProtocol(pp *p2ptest.TestPeerPool) adapters.RunProtocol { lhs := &hs0{42} // module handshake demonstrating a simple repeatable exchange of same-type message - hs, err = peer.Handshake(lhs, time.Second) + hs, err = peer.Handshake(ctx, lhs) if err != nil { return err } @@ -97,19 +97,40 @@ func newProtocol(pp *p2ptest.TestPeerPool) adapters.RunProtocol { return fmt.Errorf("handshake mismatch remote %v > local %v", rmhs.C, lhs.C) } - peer.Register(lhs, func(msg interface{}) error { - rhs := msg.(*hs0) - if rhs.C > lhs.C { - return fmt.Errorf("handshake mismatch remote %v > local %v", rhs.C, lhs.C) + handle := func(msg interface{}) error { + switch msg := msg.(type) { + + case *protoHandshake: + return errors.New("duplicate handshake") + + case *hs0: + rhs := msg + if rhs.C > lhs.C { + return fmt.Errorf("handshake mismatch remote %v > local %v", rhs.C, lhs.C) + } + lhs.C += rhs.C + return peer.Send(lhs) + + case *kill: + // demonstrates use of peerPool, killing another peer connection as a response to a message + id := msg.C + pp.Get(id).Drop(errors.New("killed")) + log.Trace(fmt.Sprintf("id %v killed", id)) + return nil + + case *drop: + // for testing we can trigger self induced disconnect upon receiving drop message + return errors.New("dropped") + + default: + return fmt.Errorf("unknown message type: %T", msg) } - lhs.C += rhs.C - return peer.Send(lhs) - }) + } log.Trace(fmt.Sprintf("adding peer %v", peer)) pp.Add(peer) defer pp.Remove(peer) - err = peer.Run() + err = peer.Run(handle) log.Trace(fmt.Sprintf("peer %v protocol quitting: %v", peer, err)) return err diff --git a/p2p/server.go b/p2p/server.go index c641e3a628..14a987d875 100644 --- a/p2p/server.go +++ b/p2p/server.go @@ -131,7 +131,7 @@ type Config struct { // If Dialer is set to a non-nil value, the given Dialer // is used to dial outbound peer connections. - Dialer *net.Dialer `toml:"-"` + Dialer NodeDialer `toml:"-"` // If NoDial is true, the server will not dial any peers. NoDial bool `toml:",omitempty"` @@ -144,6 +144,7 @@ type Config struct { type Server interface { Start() error Stop() error + SetupConn(net.Conn, connFlag, *discover.Node) AddPeer(node *discover.Node) RemovePeer(node *discover.Node) SubscribeEvents(ch chan *PeerEvent) event.Subscription @@ -385,7 +386,7 @@ func (srv *server) Start() (err error) { srv.newTransport = newRLPX } if srv.Dialer == nil { - srv.Dialer = &net.Dialer{Timeout: defaultDialTimeout} + srv.Dialer = TCPDialer{&net.Dialer{Timeout: defaultDialTimeout}} } srv.quit = make(chan struct{}) srv.addpeer = make(chan *conn) @@ -697,7 +698,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(fd, inboundConn, nil) slots <- struct{}{} }() } @@ -706,7 +707,7 @@ 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(fd net.Conn, flags connFlag, dialDest *discover.Node) { // Prevent leftover pending conns from entering the handshake. srv.lock.Lock() running := srv.running diff --git a/p2p/simulations/adapters/inproc.go b/p2p/simulations/adapters/inproc.go index 6a851fda29..3226446927 100644 --- a/p2p/simulations/adapters/inproc.go +++ b/p2p/simulations/adapters/inproc.go @@ -17,14 +17,13 @@ package adapters import ( - "context" "errors" "fmt" + "math" "net" "sync" "github.com/ethereum/go-ethereum/event" - "github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/node" "github.com/ethereum/go-ethereum/p2p" "github.com/ethereum/go-ethereum/p2p/discover" @@ -73,15 +72,28 @@ func (s *SimAdapter) NewNode(config *NodeConfig) (Node, error) { node := &SimNode{ Id: id, + config: config, adapter: s, serviceFunc: serviceFunc, - peers: make(map[discover.NodeID]MsgReadWriteCloser), - dropPeers: make(chan struct{}), } s.nodes[id.NodeID] = node return node, nil } +func (s *SimAdapter) Dial(dest *discover.Node) (conn net.Conn, err error) { + node, ok := s.GetNode(dest.ID) + if !ok { + return nil, fmt.Errorf("unknown node: %s", dest.ID) + } + srv := node.Server() + if srv == nil { + return nil, fmt.Errorf("node not running: %s", dest.ID) + } + pipe1, pipe2 := net.Pipe() + go srv.SetupConn(pipe1, 0, nil) + return pipe2, nil +} + // GetNode returns the node with the given ID if it exists func (s *SimAdapter) GetNode(id discover.NodeID) (*SimNode, bool) { s.mtx.RLock() @@ -90,14 +102,6 @@ func (s *SimAdapter) GetNode(id discover.NodeID) (*SimNode, bool) { return node, ok } -// MsgReadWriteCloser wraps a MsgReadWriter with the addition of a Close method -// so we can simulate the closing of a p2p connection (which usually happens by -/// closing the underlying TCP connection) -type MsgReadWriteCloser interface { - p2p.MsgReadWriter - Close() error -} - // SimNode is an in-memory node which connects to other SimNodes using an // in-memory p2p.MsgReadWriter pipe, running an underlying service protocol // directly over that pipe. @@ -107,17 +111,13 @@ type MsgReadWriteCloser interface { type SimNode struct { lock sync.RWMutex Id *NodeId + config *NodeConfig adapter *SimAdapter - running node.Service serviceFunc ServiceFunc - peers map[discover.NodeID]MsgReadWriteCloser - peerFeed event.Feed + node *node.Node + running node.Service client *rpc.Client rpcMux *rpcMux - - // dropPeers is used to force peer disconnects when - // the node is stopped - dropPeers chan struct{} } // Addr returns the node's discovery address @@ -127,7 +127,7 @@ func (self *SimNode) Addr() []byte { // Node returns a discover.Node representing the SimNode func (self *SimNode) Node() *discover.Node { - return discover.NewNode(self.Id.NodeID, nil, 0, 0) + return discover.NewNode(self.Id.NodeID, net.IP{127, 0, 0, 1}, 30303, 30303) } // Client returns an rpc.Client which can be used to communicate with the @@ -154,83 +154,57 @@ func (self *SimNode) ServeRPC(conn net.Conn) error { return nil } -// Start initializes the service, starts the RPC handler and then starts -// the service -func (self *SimNode) Start(snapshot []byte) error { - service := self.serviceFunc(self.Id, snapshot) +// Snapshot creates a snapshot of the service state by calling the +// simulation_snapshot RPC method +func (self *SimNode) Snapshot() ([]byte, error) { + self.lock.Lock() + defer self.lock.Unlock() + if self.client == nil { + return nil, errors.New("RPC not started") + } + var snapshot []byte + return snapshot, self.client.Call(&snapshot, "simulation_snapshot") +} - // for simplicity, only support single protocol services (simulating - // multiple protocols on the same peer is extra effort, and we don't - // currently run any simulations which run multiple protocols) - if len(service.Protocols()) != 1 { - return errors.New("service must have a single protocol") +// Start starts the RPC handler and the underlying service +func (self *SimNode) Start(snapshot []byte) error { + self.lock.Lock() + defer self.lock.Unlock() + if self.node != nil { + return errors.New("node already started") } - self.dropPeers = make(chan struct{}) - if err := self.startRPC(service); err != nil { + newService := func(ctx *node.ServiceContext) (node.Service, error) { + service := self.serviceFunc(self.Id, snapshot) + self.running = service + return service, nil + } + + node, err := node.New(&node.Config{ + P2P: p2p.Config{ + PrivateKey: self.config.PrivateKey, + MaxPeers: math.MaxInt32, + NoDiscovery: true, + Dialer: self.adapter, + EnableMsgEvents: true, + }, + NoUSB: true, + }) + if err != nil { return err } - self.running = service - return service.Start(&simServer{self}) -} -// simServer wraps a SimNode but modifies the Start method signature so that -// it implements the p2p.Server interface (the Start method is never actually -// called when using the SimAdapter) -type simServer struct { - *SimNode -} - -func (s *simServer) Start() error { - return nil -} - -// Stop stops the RPC handler, stops the underlying service and disconnects -// any currently connected peers -func (self *SimNode) Stop() error { - self.stopRPC() - close(self.dropPeers) - return self.running.Stop() -} - -// Running returns whether or not the service is running -func (self *SimNode) Running() bool { - self.lock.Lock() - defer self.lock.Unlock() - return self.running != nil -} - -// Service returns the running node.Service -func (self *SimNode) Service() node.Service { - return self.running -} - -// startRPC starts an RPC server and connects to it using an in-process RPC -// client -func (self *SimNode) startRPC(service node.Service) error { - self.lock.Lock() - defer self.lock.Unlock() - if self.client != nil { - return nil - // return errors.New("RPC already started") + if err := node.Register(newService); err != nil { + return err } - // add SimAdminAPI so that the network can call the - // AddPeer, RemovePeer and PeerEvents RPC methods - apis := append(service.APIs(), []rpc.API{ - { - Namespace: "admin", - Version: "1.0", - Service: &SimAdminAPI{self}, - }, - }...) + if err := node.Start(); err != nil { + return err + } - // start the RPC handler - handler := rpc.NewServer() - for _, api := range apis { - if err := handler.RegisterName(api.Namespace, api.Service); err != nil { - return fmt.Errorf("error registering RPC: %s", err) - } + handler, err := node.RPCHandler() + if err != nil { + return err } // create an in-process RPC multiplexer @@ -241,197 +215,55 @@ func (self *SimNode) startRPC(service node.Service) error { // create an in-process RPC client self.client = self.rpcMux.Client() + self.node = node + return nil } -// stopRPC closes the node's RPC client -func (self *SimNode) stopRPC() { +func (self *SimNode) Stop() error { self.lock.Lock() defer self.lock.Unlock() - if self.client != nil { - self.client.Close() - self.client = nil - self.rpcMux = nil + if self.node == nil { + return nil } + if err := self.node.Stop(); err != nil { + return err + } + self.node = nil + return nil } -// RemovePeer removes the given node as a peer by looking up the corresponding -// p2p.MsgReadWriter pipe and closing it (which will cause both the local -// and peer Protocol.Run functions to exit) -func (self *SimNode) RemovePeer(peer *discover.Node) { +// Service returns the underlying node.Service +func (self *SimNode) Service() node.Service { self.lock.Lock() defer self.lock.Unlock() - peerRW, exists := self.peers[peer.ID] - if !exists { - return - } - peerRW.Close() - delete(self.peers, peer.ID) - log.Trace(fmt.Sprintf("dropped peer %v", peer.ID)) + return self.running } -// AddPeer adds the given node as a peer by creating a p2p.MsgReadWriter pipe -// and running both the local and peer's Protocol.Run function over the pipe -func (self *SimNode) AddPeer(peer *discover.Node) { +func (self *SimNode) Server() *p2p.Server { self.lock.Lock() defer self.lock.Unlock() - if _, exists := self.peers[peer.ID]; exists { - return + if self.node == nil { + return nil } - peerNode, exists := self.adapter.GetNode(peer.ID) - if !exists { - panic(fmt.Sprintf("unknown peer: %s", peer.ID)) - } - if !peerNode.Running() { - return - } - p1, p2 := p2p.MsgPipe() - localRW := p2p.NewMsgEventer(p1, &self.peerFeed, peer.ID) - peerRW := p2p.NewMsgEventer(p2, &self.peerFeed, self.Id.NodeID) - self.peers[peer.ID] = peerRW - peerNode.RunProtocol(self, peerRW) - self.RunProtocol(peerNode, localRW) + return self.node.Server() } -// SubscribeEvents subscribes the given channel to p2p peer events func (self *SimNode) SubscribeEvents(ch chan *p2p.PeerEvent) event.Subscription { - return self.peerFeed.Subscribe(ch) + srv := self.Server() + if srv == nil { + panic("node not running") + } + return srv.SubscribeEvents(ch) } -// PeerCount returns the number of currently connected peers -func (self *SimNode) PeerCount() int { - self.lock.Lock() - defer self.lock.Unlock() - return len(self.peers) -} - -// NodeInfo returns information about the node func (self *SimNode) NodeInfo() *p2p.NodeInfo { - self.lock.Lock() - defer self.lock.Unlock() - info := &p2p.NodeInfo{ - ID: self.Id.String(), - Enode: self.Node().String(), - Protocols: make(map[string]interface{}), - } - if self.running != nil { - for _, proto := range self.running.Protocols() { - nodeInfo := interface{}("unknown") - if query := proto.NodeInfo; query != nil { - nodeInfo = proto.NodeInfo() - } - info.Protocols[proto.Name] = nodeInfo + server := self.Server() + if server == nil { + return &p2p.NodeInfo{ + ID: self.Id.String(), + Enode: self.Node().String(), } } - return info -} - -// PeersInfo is a stub so that SimNode implements p2p.Server -func (self *SimNode) PeersInfo() (info []*p2p.PeerInfo) { - return nil -} - -// Snapshot creates a snapshot of the running service -func (self *SimNode) Snapshot() ([]byte, error) { - self.lock.Lock() - service := self.running - self.lock.Unlock() - if service == nil { - return nil, errors.New("service not running") - } - return SnapshotAPI{service}.Snapshot() -} - -// RunProtocol runs the underlying service's protocol with the peer using the -// given MsgReadWriteCloser, emitting peer add / drop events for peer event -// subscribers -func (self *SimNode) RunProtocol(peer *SimNode, rw MsgReadWriteCloser) { - // close the rw if the node is stopped to disconnect the peer - go func() { - <-self.dropPeers - log.Trace("dropping peer", "self.id", self.Id, "peer.id", peer.Id) - rw.Close() - }() - - id := peer.Id - log.Trace(fmt.Sprintf("protocol starting on peer %v (connection with %v)", self.Id, id)) - protocol := self.running.Protocols()[0] - p := p2p.NewPeer(id.NodeID, id.Label(), []p2p.Cap{}) - go func() { - // emit peer add event - self.peerFeed.Send(&p2p.PeerEvent{ - Type: p2p.PeerEventTypeAdd, - Peer: id.NodeID, - }) - - // run the protocol - err := protocol.Run(p, rw) - - // remove the peer - self.RemovePeer(peer.Node()) - log.Trace(fmt.Sprintf("protocol quit on peer %v (connection with %v broken: %v)", self.Id, id, err)) - - // emit peer drop event - self.peerFeed.Send(&p2p.PeerEvent{ - Type: p2p.PeerEventTypeDrop, - Peer: id.NodeID, - Error: err.Error(), - }) - }() -} - -// SimAdminAPI implements the AddPeer and RemovePeer RPC methods (API -// compatible with node.PrivateAdminAPI) -type SimAdminAPI struct { - *SimNode -} - -func (api *SimAdminAPI) AddPeer(url string) (bool, error) { - node, err := discover.ParseNode(url) - if err != nil { - return false, fmt.Errorf("invalid enode: %v", err) - } - api.SimNode.AddPeer(node) - return true, nil -} - -func (api *SimAdminAPI) RemovePeer(url string) (bool, error) { - node, err := discover.ParseNode(url) - if err != nil { - return false, fmt.Errorf("invalid enode: %v", err) - } - api.SimNode.RemovePeer(node) - return true, nil -} - -// PeerEvents creates an RPC subscription which receives peer events from the -// underlying p2p.Server -func (api *SimAdminAPI) PeerEvents(ctx context.Context) (*rpc.Subscription, error) { - notifier, supported := rpc.NotifierFromContext(ctx) - if !supported { - return &rpc.Subscription{}, rpc.ErrNotificationsUnsupported - } - - rpcSub := notifier.CreateSubscription() - - go func() { - events := make(chan *p2p.PeerEvent) - sub := api.SubscribeEvents(events) - defer sub.Unsubscribe() - - for { - select { - case event := <-events: - notifier.Notify(rpcSub.ID, event) - case <-sub.Err(): - return - case <-rpcSub.Err(): - return - case <-notifier.Closed(): - return - } - } - }() - - return rpcSub, nil + return server.NodeInfo() } diff --git a/p2p/testing/protocolsession.go b/p2p/testing/protocolsession.go index 1dfbb6ab49..d611cd45a5 100644 --- a/p2p/testing/protocolsession.go +++ b/p2p/testing/protocolsession.go @@ -13,7 +13,7 @@ import ( ) type ProtocolSession struct { - *adapters.SimNode + p2p.Server Ids []*adapters.NodeId adapter *adapters.SimAdapter diff --git a/p2p/testing/protocoltester.go b/p2p/testing/protocoltester.go index 16a23f120f..695cfd0292 100644 --- a/p2p/testing/protocoltester.go +++ b/p2p/testing/protocoltester.go @@ -20,10 +20,10 @@ type ProtocolTester struct { func NewProtocolTester(t *testing.T, id *adapters.NodeId, n int, run func(*p2p.Peer, p2p.MsgReadWriter) error) *ProtocolTester { services := map[string]adapters.ServiceFunc{ - "test": func(id *adapters.NodeId) node.Service { + "test": func(id *adapters.NodeId, _ []byte) node.Service { return &testNode{run} }, - "mock": func(id *adapters.NodeId) node.Service { + "mock": func(id *adapters.NodeId, _ []byte) node.Service { return newMockNode() }, } @@ -47,7 +47,7 @@ func NewProtocolTester(t *testing.T, id *adapters.NodeId, n int, run func(*p2p.P events := make(chan *p2p.PeerEvent, 1000) node.SubscribeEvents(events) ps := &ProtocolSession{ - SimNode: node, + Server: node.Server(), Ids: peerIDs, adapter: adapter, events: events, @@ -86,7 +86,10 @@ type testNode struct { } func (t *testNode) Protocols() []p2p.Protocol { - return []p2p.Protocol{{Run: t.run}} + return []p2p.Protocol{{ + Length: 100, + Run: t.run, + }} } func (t *testNode) APIs() []rpc.API { diff --git a/swarm/network/discovery.go b/swarm/network/discovery.go index 55a75ed858..b1830118d7 100644 --- a/swarm/network/discovery.go +++ b/swarm/network/discovery.go @@ -8,13 +8,6 @@ import ( // discovery bzz overlay extension doing peer relaying -// messages related to peer discovery -var DiscoveryMsgs = []interface{}{ - &getPeersMsg{}, - &peersMsg{}, - &subPeersMsg{}, -} - type discPeer struct { *bzzPeer overlay Overlay @@ -32,14 +25,26 @@ func NewDiscovery(p *bzzPeer, o Overlay) *discPeer { peers: make(map[string]bool), } self.seen(self) - - p.Register(&peersMsg{}, self.handlePeersMsg) - p.Register(&getPeersMsg{}, self.handleGetPeersMsg) - p.Register(&subPeersMsg{}, self.handleSubPeersMsg) - return self } +func (self *discPeer) HandleMsg(msg interface{}) error { + switch msg := msg.(type) { + + case *peersMsg: + return self.handlePeersMsg(msg) + + case *getPeersMsg: + return self.handleGetPeersMsg(msg) + + case *subPeersMsg: + return self.handleSubPeersMsg(msg) + + default: + return fmt.Errorf("unknown message type: %T", msg) + } +} + // NotifyPeer notifies the receiver remote end of a peer p or PO po. // callback for overlay driver func (self *discPeer) NotifyPeer(p OverlayPeer, po uint8) error { @@ -109,9 +114,8 @@ func (self subPeersMsg) String() string { return fmt.Sprintf("%T: request peers > PO%02d. ", self, self.ProxLimit) } -func (self *discPeer) handleSubPeersMsg(msg interface{}) error { - spm := msg.(*subPeersMsg) - self.proxLimit = spm.ProxLimit +func (self *discPeer) handleSubPeersMsg(msg *subPeersMsg) error { + self.proxLimit = msg.ProxLimit if !self.sentPeers { var peers []*bzzAddr self.overlay.EachConn(self.Over(), 255, func(p OverlayConn, po int, isproxbin bool) bool { @@ -138,17 +142,16 @@ func (self *discPeer) handleSubPeersMsg(msg interface{}) error { // handlePeersMsg called by the protocol when receiving peerset (for target address) // list of nodes ([]PeerAddr in peersMsg) is added to the overlay db using the // Register interface method -func (self *discPeer) handlePeersMsg(msg interface{}) error { +func (self *discPeer) handlePeersMsg(msg *peersMsg) error { // register all addresses - as := msg.(*peersMsg).Peers - if len(as) == 0 { + if len(msg.Peers) == 0 { log.Debug(fmt.Sprintf("whoops, no peers in incoming peersMsg from %v", self)) return nil } var c chan OverlayAddr go func() { - for _, a := range as { + for _, a := range msg.Peers { self.seen(a) c <- a } @@ -161,18 +164,17 @@ func (self *discPeer) handlePeersMsg(msg interface{}) error { // peers suggestions are retrieved from the overlay topology driver // using the EachConn interface iterator method // peers sent are remembered throughout a session and not sent twice -func (self *discPeer) handleGetPeersMsg(msg interface{}) error { +func (self *discPeer) handleGetPeersMsg(msg *getPeersMsg) error { var peers []*bzzAddr - req := msg.(*getPeersMsg) i := 0 - self.overlay.EachConn(self.Over(), int(req.Order), func(p OverlayConn, po int, isproxbin bool) bool { + self.overlay.EachConn(self.Over(), int(msg.Order), func(p OverlayConn, po int, isproxbin bool) bool { i++ // only send peers we have not sent before in this session a := ToAddr(p) if self.seen(a) { peers = append(peers, a) } - return len(peers) < int(req.Max) + return len(peers) < int(msg.Max) }) if len(peers) == 0 { log.Debug(fmt.Sprintf("no peers found for %v", self)) diff --git a/swarm/network/discovery_test.go b/swarm/network/discovery_test.go index 22a61a04ea..dfab1eef5a 100644 --- a/swarm/network/discovery_test.go +++ b/swarm/network/discovery_test.go @@ -16,22 +16,18 @@ import ( func TestDiscovery(t *testing.T) { addr := RandomAddr() to := NewKademlia(addr.OAddr, NewKadParams()) - ct := BzzCodeMap(DiscoveryMsgs...) - services := func(p *bzzPeer) error { + run := func(p *bzzPeer) error { dp := NewDiscovery(p, to) - to.On(dp) + to.On(p) + defer to.Off(p) log.Trace(fmt.Sprintf("kademlia on %v", p)) - p.DisconnectHook(func(err error) { - to.Off(p) - }) - return nil + return p.Run(dp.HandleMsg) } - s := newBzzBaseTester(t, 1, addr, ct, services) + s := newBzzBaseTester(t, 1, addr, DiscoveryProtocol, run) defer s.Stop() - s.runHandshakes() s.TestExchanges(p2ptest.Exchange{ Label: "outgoing SubPeersMsg", Expects: []p2ptest.Expect{ diff --git a/swarm/network/hive.go b/swarm/network/hive.go index 2ee8cd3b46..07c1bef6b3 100644 --- a/swarm/network/hive.go +++ b/swarm/network/hive.go @@ -34,7 +34,7 @@ it uses an Overlay Topology driver (e.g., generic kademlia nodetable) to find best peer list for any target this is used by the netstore to search for content in the swarm -It handles the bzz protocol getPeersMsg peersMsg exchange +It handles the hive protocol getPeersMsg peersMsg exchange and relay the peer request process to the Overlay module peer connections and disconnections are reported and registered @@ -57,23 +57,19 @@ type Overlay interface { BaseAddr() []byte } -// ReadWriter interface to persist known peers, uses disk for real nodes -type ReadWriter interface { - ReadAll(string) ([]byte, error) - WriteAll(string, []byte) error -} - // Hive implements the PeerPool interface type Hive struct { - *HiveParams // settings - Overlay // the overlay topology driver - RW ReadWriter // ReadWriter + *HiveParams // settings + Overlay // the overlay topology driver + store Store // bookkeeping lock sync.Mutex quit chan bool toggle chan bool more chan bool + + newTicker func() hiveTicker } // HiveParams holds the config options to hive @@ -81,7 +77,7 @@ type HiveParams struct { Discovery bool // if want discovery of not PeersBroadcastSetSize uint8 // how many peers to use when relaying MaxPeersPerRequest uint8 // max size for peer address batches - CallInterval uint // polling interval fir=== + KeepAliveInterval time.Duration } // NewHiveParams returns hive config with only the @@ -90,17 +86,18 @@ func NewHiveParams() *HiveParams { Discovery: true, PeersBroadcastSetSize: 2, MaxPeersPerRequest: 5, - CallInterval: 1000, + KeepAliveInterval: time.Second, } } // Hive constructor embeds both arguments // HiveParams: config parameters // Overlay: Topology Driver Interface -func NewHive(params *HiveParams, overlay Overlay) *Hive { +func NewHive(params *HiveParams, overlay Overlay, store Store) *Hive { return &Hive{ HiveParams: params, Overlay: overlay, + store: store, } } @@ -109,8 +106,8 @@ func NewHive(params *HiveParams, overlay Overlay) *Hive { // these are called on the p2p.Server which runs on the node // af() returns an arbitrary ticker channel // rw is a read writer for json configs -func (self *Hive) Start(server p2p.Server, af func() <-chan time.Time, rw ReadWriter) error { - if rw != nil { +func (self *Hive) Start(server p2p.Server) error { + if self.store != nil { if err := self.loadPeers(); err != nil { return err } @@ -120,7 +117,7 @@ func (self *Hive) Start(server p2p.Server, af func() <-chan time.Time, rw ReadWr self.quit = make(chan bool) log.Debug("hive started") // this loop is doing bootstrapping and maintains a healthy table - go self.keepAlive(af) + go self.keepAlive() go func() { // each iteration, ask kademlia about most preferred peer for more := range self.more { @@ -163,16 +160,18 @@ func (self *Hive) Start(server p2p.Server, af func() <-chan time.Time, rw ReadWr // Stop terminates the updateloop and saves the peers func (self *Hive) Stop() { - if self.RW != nil { + if self.store != nil { self.savePeers() } // closing toggle channel quits the updateloop close(self.quit) } -// default ticker, tickinterval is taken from KadParams.CallInterval -func (self *Hive) ticker() <-chan time.Time { - return time.NewTicker(time.Duration(self.CallInterval) * time.Millisecond).C +func (self *Hive) Run(peer *bzzPeer) error { + discPeer := NewDiscovery(peer, self) + self.On(discPeer) + defer self.Off(discPeer) + return peer.Run(discPeer.HandleMsg) } // Add is called at the end of a successful protocol handshake @@ -242,27 +241,48 @@ func ToAddr(pa OverlayPeer) *bzzAddr { return pa.(*bzzPeer).bzzAddr } +type hiveTicker interface { + Ch() <-chan time.Time + Stop() +} + +type timeTicker struct { + *time.Ticker +} + +func (t *timeTicker) Ch() <-chan time.Time { + return t.C +} + // keepAlive is a forever loop // in its awake state it periodically triggers connection attempts // by writing to self.more until Kademlia Table is saturated // wake state is toggled by writing to self.toggle // it goes to sleep mode if table is saturated // it restarts if the table becomes non-full again due to disconnections -func (self *Hive) keepAlive(af func() <-chan time.Time) { - log.Trace("keep alive loop started") - alarm := af() +func (self *Hive) keepAlive() { + if self.newTicker == nil { + self.newTicker = func() hiveTicker { + return &timeTicker{time.NewTicker(self.KeepAliveInterval)} + } + } + ticker := self.newTicker() + tick := ticker.Ch() for { select { - case <-alarm: + case <-tick: log.Trace("wake up: make hive alive") self.wake() case need := <-self.toggle: - if alarm == nil && need { - alarm = af() + if ticker == nil && need { + ticker = self.newTicker() + tick = ticker.Ch() } // if hive saturated, no more peers asked - if alarm != nil && !need { - alarm = nil + if ticker != nil && !need { + ticker.Stop() + ticker = nil + tick = nil } case <-self.quit: return @@ -272,8 +292,7 @@ func (self *Hive) keepAlive(af func() <-chan time.Time) { // loadPeers, savePeer implement persistence callback/ func (self *Hive) loadPeers() error { - rw := self.RW - data, err := rw.ReadAll("peers") + data, err := self.store.Load("peers") if err != nil { return err } @@ -310,7 +329,7 @@ func (self *Hive) savePeers() error { if err != nil { return fmt.Errorf("could not encode peers: %v", err) } - if err := self.RW.WriteAll("peers", data); err != nil { + if err := self.store.Save("peers", data); err != nil { return fmt.Errorf("could not save peers: %v", err) } return nil diff --git a/swarm/network/hive_test.go b/swarm/network/hive_test.go index aaa90e71a3..473ba81c0d 100644 --- a/swarm/network/hive_test.go +++ b/swarm/network/hive_test.go @@ -16,10 +16,13 @@ type testConnect struct { ticker chan time.Time } -func (self *testConnect) ping() <-chan time.Time { +func (self *testConnect) Ch() <-chan time.Time { return self.ticker } +func (self *testConnect) Stop() { +} + func (self *testConnect) connect(na string) error { self.mu.Lock() defer self.mu.Unlock() @@ -31,38 +34,10 @@ func (self *testConnect) connect(na string) error { func newHiveTester(t *testing.T, params *HiveParams) (*bzzTester, *Hive) { // setup addr := RandomAddr() // tested peers peer address - // to := NewTestOverlay(addr.Over()) // overlay topology drive - pp := NewHive(params, nil) // hive - // pp := NewHive(params, to) // hive - - ct := BzzCodeMap(DiscoveryMsgs...) // bzz protocol code map - services := func(p *bzzPeer) error { - pp.Add(p) - p.DisconnectHook(func(err error) { - pp.Remove(p) - }) - return nil - } - - return newBzzBaseTester(t, 1, addr, ct, services), pp -} - -func TestOverlayRegistration(t *testing.T) { - params := NewHiveParams() - params.Discovery = false - s, pp := newHiveTester(t, params) - defer s.Stop() - - id := s.Ids[0] - raddr := NewAddrFromNodeId(id) - - s.runHandshakes() - - // hive should have called the overlay - // if pp.Overlay.(*testOverlay).posMap[string(raddr.Over())] == nil { - // t.Fatalf("Overlay#On not called on new peer") - // } + to := NewKademlia(addr.OAddr, NewKadParams()) + pp := NewHive(params, to, nil) // hive + return newBzzBaseTester(t, 1, addr, DiscoveryProtocol, pp.Run), pp } func TestRegisterAndConnect(t *testing.T) { @@ -73,7 +48,12 @@ func TestRegisterAndConnect(t *testing.T) { id := s.Ids[0] raddr := NewAddrFromNodeId(id) - pp.Register(raddr) + ch := make(chan OverlayAddr) + go func() { + ch <- raddr + close(ch) + }() + pp.Register(ch) // start the hive and wait for the connection tc := &testConnect{ @@ -83,18 +63,17 @@ func TestRegisterAndConnect(t *testing.T) { }, ticker: make(chan time.Time), } - pp.Start(s, tc.ping, nil) + pp.newTicker = func() hiveTicker { return tc } + pp.Start(s) defer pp.Stop() tc.ticker <- time.Now() - s.runHandshakes() - // if pp.Overlay.(*testOverlay).posMap[string(raddr.Over())] == nil { // t.Fatalf("Overlay#On not called on new peer") // } // retrieve and broadcast - ord := order(raddr.Over()) + ord := raddr.Over()[0] / 32 o := 0 if ord == 0 { o = 1 diff --git a/swarm/network/kademlia_test.go b/swarm/network/kademlia_test.go index 40e1cbac97..c6bc875976 100644 --- a/swarm/network/kademlia_test.go +++ b/swarm/network/kademlia_test.go @@ -153,11 +153,14 @@ func (k *testKademlia) Off(offs ...string) *testKademlia { } func (k *testKademlia) Register(regs ...string) *testKademlia { - var ps []Addr - for _, s := range regs { - ps = append(ps, Addr(testKadPeerAddr(s))) - } - k.Kademlia.Register(ps...) + ch := make(chan OverlayAddr) + go func() { + defer close(ch) + for _, s := range regs { + ch <- testKadPeerAddr(s) + } + }() + k.Kademlia.Register(ch) return k } diff --git a/swarm/network/protocol.go b/swarm/network/protocol.go index 523031db17..e4ad09f705 100644 --- a/swarm/network/protocol.go +++ b/swarm/network/protocol.go @@ -17,7 +17,10 @@ package network import ( + "context" + "errors" "fmt" + "sync" "time" "github.com/ethereum/go-ethereum/crypto" @@ -26,15 +29,43 @@ import ( "github.com/ethereum/go-ethereum/p2p/discover" "github.com/ethereum/go-ethereum/p2p/protocols" "github.com/ethereum/go-ethereum/p2p/simulations/adapters" + "github.com/ethereum/go-ethereum/rpc" ) const ( - ProtocolName = "bzz" - Version = 0 NetworkId = 322 // BZZ in l33t ProtocolMaxMsgSize = 10 * 1024 * 1024 ) +var BzzProtocol = &protocols.Spec{ + Name: "bzz", + Version: 1, + MaxMsgSize: 10 * 1024 * 1024, + Messages: []interface{}{ + bzzHandshake{}, + }, +} + +var DiscoveryProtocol = &protocols.Spec{ + Name: "hive", + Version: 1, + MaxMsgSize: 10 * 1024 * 1024, + Messages: []interface{}{ + peersMsg{}, + getPeersMsg{}, + subPeersMsg{}, + }, +} + +var PssProtocol = &protocols.Spec{ + Name: "pss", + Version: 1, + MaxMsgSize: 10 * 1024 * 1024, + Messages: []interface{}{ + PssMsg{}, + }, +} + // the Addr interface that peerPool needs type Addr interface { OverlayPeer @@ -52,13 +83,145 @@ type Peer interface { // Conn interface represents an live peer connection type Conn interface { - ID() discover.NodeID // the key that uniquely identifies the Node for the peerPool - Handshake(interface{}, time.Duration) (interface{}, error) // can send messages - Send(interface{}) error // can send messages - Drop(error) // disconnect this peer - Register(interface{}, func(interface{}) error) uint64 // register message-handler callbacks - DisconnectHook(func(error)) // register message-handler callbacks - Run() error // the run function to run a protocol + ID() discover.NodeID // the key that uniquely identifies the Node for the peerPool + Handshake(context.Context, interface{}) (interface{}, error) // can send messages + Send(interface{}) error // can send messages + Drop(error) // disconnect this peer + Run(func(interface{}) error) error // the run function to run a protocol +} + +// TODO: implement store for exec nodes +type Store interface { + Load(string) ([]byte, error) + Save(string, []byte) error +} + +type BzzConfig struct { + OverlayAddr []byte + UnderlayAddr []byte + + KadParams *KadParams + HiveParams *HiveParams + PssParams *PssParams + + Store Store +} + +func NewBzz(config *BzzConfig) *Bzz { + kademlia := NewKademlia(config.OverlayAddr, config.KadParams) + bzz := &Bzz{ + Kademlia: kademlia, + Hive: NewHive(config.HiveParams, kademlia, config.Store), + localAddr: &bzzAddr{config.OverlayAddr, config.UnderlayAddr}, + handshakes: make(map[discover.NodeID]*bzzHandshake), + } + if config.PssParams != nil { + bzz.Pss = NewPss(kademlia, config.PssParams) + } + return bzz +} + +type Bzz struct { + Kademlia *Kademlia + Hive *Hive + Pss *Pss + + localAddr *bzzAddr + mtx sync.Mutex + handshakes map[discover.NodeID]*bzzHandshake +} + +func (b *Bzz) Protocols() []p2p.Protocol { + return []p2p.Protocol{ + { + Name: BzzProtocol.Name, + Version: BzzProtocol.Version, + Length: BzzProtocol.Length(), + Run: b.runHandshake, + }, + { + Name: DiscoveryProtocol.Name, + Version: DiscoveryProtocol.Version, + Length: DiscoveryProtocol.Length(), + Run: b.runProtocol(DiscoveryProtocol, b.Hive.Run), + NodeInfo: b.Hive.NodeInfo, + PeerInfo: b.Hive.PeerInfo, + }, + { + Name: PssProtocol.Name, + Version: PssProtocol.Version, + Length: PssProtocol.Length(), + Run: b.runProtocol(PssProtocol, b.Pss.Run), + }, + } +} + +func (b *Bzz) APIs() []rpc.API { + return []rpc.API{{ + Namespace: "hive", + Version: "1.0", + Service: b.Hive, + }} +} + +func (b *Bzz) Start(server p2p.Server) error { + return b.Hive.Start(server) +} + +func (b *Bzz) Stop() error { + b.Hive.Stop() + return nil +} + +func (b *Bzz) runHandshake(p *p2p.Peer, rw p2p.MsgReadWriter) error { + handshake := b.getHandshake(p.ID()) + + if err := handshake.Perform(p, rw); err != nil { + log.Error("handshake failed", "peer", p.ID(), "err", err) + return err + } + + // fail if we get another handshake + msg, err := rw.ReadMsg() + if err != nil { + return err + } + msg.Discard() + return errors.New("received multiple handshakes") +} + +func (b *Bzz) runProtocol(spec *protocols.Spec, run func(*bzzPeer) error) func(*p2p.Peer, p2p.MsgReadWriter) error { + return func(p *p2p.Peer, rw p2p.MsgReadWriter) error { + // wait for the bzz protocol to perform the handshake + handshake := b.getHandshake(p.ID()) + if err := handshake.Wait(); err != nil { + return err + } + + // the handshake has succeeded so run the service + peer := &bzzPeer{ + Conn: protocols.NewPeer(p, rw, spec), + localAddr: b.localAddr, + bzzAddr: handshake.peerAddr, + } + return run(peer) + } +} + +func (b *Bzz) getHandshake(peerID discover.NodeID) *bzzHandshake { + b.mtx.Lock() + defer b.mtx.Unlock() + handshake, ok := b.handshakes[peerID] + if !ok { + handshake = &bzzHandshake{ + Version: uint64(BzzProtocol.Version), + NetworkId: uint64(NetworkId), + Addr: b.localAddr, + done: make(chan struct{}), + } + b.handshakes[peerID] = handshake + } + return handshake } // bzzPeer is the bzz protocol view of a protocols.Peer (itself an extension of p2p.Peer) @@ -70,6 +233,13 @@ type bzzPeer struct { lastActive time.Time // time is updated whenever mutexes are releasing } +func newBzzPeer(conn Conn, over, under []byte) *bzzPeer { + return &bzzPeer{ + Conn: conn, + localAddr: &bzzAddr{over, under}, + } +} + // Off returns the overlay peer record for offline persistance func (self *bzzPeer) Off() OverlayAddr { return self.bzzAddr @@ -80,47 +250,6 @@ func (self *bzzPeer) LastActive() time.Time { return self.lastActive } -// BzzCodeMap compiles the message codes and message types bzz wire protocol. -// note each call to Register can start a new series (initial code is arg1) -// the initial offset for a series is arbitrary (to ensure u) -func BzzCodeMap(msgs ...interface{}) *protocols.CodeMap { - ct := protocols.NewCodeMap(ProtocolName, Version, ProtocolMaxMsgSize) - ct.Register(0, &bzzHandshake{}) - ct.Register(1, msgs...) - return ct -} - -// NewBzz is the protocol constructor -// returns p2p.Protocol that is to be offered by the node.Service -func NewBzz(over, under []byte, ct *protocols.CodeMap, services func(*bzzPeer) error, peerInfo func(id discover.NodeID) interface{}, nodeInfo func() interface{}) *p2p.Protocol { - run := func(p *protocols.Peer) error { - bee := &bzzPeer{ - Conn: p, - localAddr: &bzzAddr{over, under}, - } - // protocol handshake and its validation - // sets remote peer address - err := bee.bzzHandshake() - if err != nil { - log.Error(fmt.Sprintf("handshake error in peer %v: %v", bee.ID(), err)) - return err - } - - // mount external service models on the peer connection (swap, sync, hive) - if services != nil { - err = services(bee) - if err != nil { - log.Error(fmt.Sprintf("protocol service error for peer %v: %v", bee.ID(), err)) - return err - } - } - - return bee.Run() - } - - return protocols.NewProtocol(ProtocolName, Version, run, ct, peerInfo, nodeInfo) -} - /* Handshake @@ -132,12 +261,52 @@ type bzzHandshake struct { Version uint64 NetworkId uint64 Addr *bzzAddr + + // peerAddr is the address received in the peer handshake + peerAddr *bzzAddr + + done chan struct{} + err error } func (self *bzzHandshake) String() string { return fmt.Sprintf("Handshake: Version: %v, NetworkId: %v, Addr: %v", self.Version, self.NetworkId, self.Addr) } +const bzzHandshakeTimeout = time.Second + +func (self *bzzHandshake) Perform(p *p2p.Peer, rw p2p.MsgReadWriter) (err error) { + defer func() { + self.err = err + close(self.done) + }() + peer := protocols.NewPeer(p, rw, BzzProtocol) + ctx, cancel := context.WithTimeout(context.Background(), bzzHandshakeTimeout) + defer cancel() + hs, err := peer.Handshake(ctx, self) + if err != nil { + return err + } + rhs := hs.(*bzzHandshake) + if rhs.NetworkId != self.NetworkId { + return fmt.Errorf("network id mismatch %d (!= %d)", rhs.NetworkId, self.NetworkId) + } + if rhs.Version != self.Version { + return fmt.Errorf("version mismatch %d (!= %d)", rhs.Version, self.Version) + } + self.peerAddr = rhs.Addr + return nil +} + +func (self *bzzHandshake) Wait() error { + select { + case <-self.done: + return self.err + case <-time.After(bzzHandshakeTimeout): + return errors.New("timed out waiting for bzz handshake") + } +} + // bzzAddr implements the PeerAddr interface type bzzAddr struct { OAddr []byte @@ -149,6 +318,9 @@ func (self *bzzAddr) Address() []byte { return self.OAddr } +func (self *bzzAddr) Bytes() []byte { + return self.OAddr +} func (self *bzzAddr) Over() []byte { return self.OAddr } @@ -171,47 +343,6 @@ func (self *bzzAddr) String() string { return fmt.Sprintf("%x <%x>", self.OAddr, self.UAddr) } -// bzzHandshake negotiates the bzz master handshake -// and validates the response, returns error when -// mismatch/incompatibility is evident -func (self *bzzPeer) bzzHandshake() error { - - lhs := &bzzHandshake{ - Version: uint64(Version), - NetworkId: uint64(NetworkId), - Addr: self.localAddr, - } - - hs, err := self.Handshake(lhs, time.Second) - if err != nil { - log.Error(fmt.Sprintf("handshake failed: %v", err)) - return err - } - - rhs := hs.(*bzzHandshake) - self.bzzAddr = rhs.Addr - err = checkBzzHandshake(rhs) - if err != nil { - log.Error(fmt.Sprintf("handshake between %v and %v failed: %v", self.localAddr, self.bzzAddr, err)) - return err - } - return nil -} - -// checkBzzHandshake checks for the validity and compatibility of the remote handshake -func checkBzzHandshake(rhs *bzzHandshake) error { - - if NetworkId != rhs.NetworkId { - return fmt.Errorf("network id mismatch %d (!= %d)", rhs.NetworkId, NetworkId) - } - - if Version != rhs.Version { - return fmt.Errorf("version mismatch %d (!= %d)", rhs.Version, Version) - } - - return nil -} - // RandomAddr is a utility method generating an address from a public key func RandomAddr() *bzzAddr { key, err := crypto.GenerateKey() diff --git a/swarm/network/protocol_test.go b/swarm/network/protocol_test.go index cc4372a763..15875eb9ce 100644 --- a/swarm/network/protocol_test.go +++ b/swarm/network/protocol_test.go @@ -2,14 +2,43 @@ package network import ( "fmt" + "sync" "testing" "github.com/ethereum/go-ethereum/log" + "github.com/ethereum/go-ethereum/p2p" "github.com/ethereum/go-ethereum/p2p/protocols" "github.com/ethereum/go-ethereum/p2p/simulations/adapters" p2ptest "github.com/ethereum/go-ethereum/p2p/testing" ) +type testStore struct { + sync.Mutex + + values map[string][]byte +} + +func newTestStore() *testStore { + return &testStore{values: make(map[string][]byte)} +} + +func (t *testStore) Load(key string) ([]byte, error) { + t.Lock() + defer t.Unlock() + v, ok := t.values[key] + if !ok { + return nil, fmt.Errorf("key not found: %s", key) + } + return v, nil +} + +func (t *testStore) Save(key string, v []byte) error { + t.Lock() + defer t.Unlock() + t.values[key] = v + return nil +} + func bzzHandshakeExchange(lhs, rhs *bzzHandshake, id *adapters.NodeId) []p2ptest.Exchange { return []p2ptest.Exchange{ @@ -34,19 +63,21 @@ func bzzHandshakeExchange(lhs, rhs *bzzHandshake, id *adapters.NodeId) []p2ptest } } -func newBzzBaseTester(t *testing.T, n int, addr *bzzAddr, ct *protocols.CodeMap, services func(*bzzPeer) error) *bzzTester { - if ct == nil { - ct = BzzCodeMap() - } - +func newBzzBaseTester(t *testing.T, n int, addr *bzzAddr, spec *protocols.Spec, run func(*bzzPeer) error) *bzzTester { cs := make(map[string]chan bool) srv := func(p *bzzPeer) error { defer close(cs[p.ID().String()]) - return services(p) + return run(p) } - protocall := NewBzz(addr.Over(), addr.Under(), ct, srv, nil, nil).Run + protocall := func(p *p2p.Peer, rw p2p.MsgReadWriter) error { + return srv(&bzzPeer{ + Conn: protocols.NewPeer(p, rw, spec), + localAddr: addr, + bzzAddr: NewAddrFromNodeId(&adapters.NodeId{NodeID: p.ID()}), + }) + } s := p2ptest.NewProtocolTester(t, NewNodeIdFromAddr(addr), n, protocall) @@ -55,7 +86,7 @@ func newBzzBaseTester(t *testing.T, n int, addr *bzzAddr, ct *protocols.CodeMap, } return &bzzTester{ - addr: addr.Address(), + addr: addr, ProtocolTester: s, cs: cs, } @@ -63,26 +94,18 @@ func newBzzBaseTester(t *testing.T, n int, addr *bzzAddr, ct *protocols.CodeMap, type bzzTester struct { *p2ptest.ProtocolTester - addr []byte + addr *bzzAddr cs map[string]chan bool } -func newBzzTester(t *testing.T, n int, addr *bzzAddr, pp *p2ptest.TestPeerPool, ct *protocols.CodeMap, services func(Peer) error) *bzzTester { +func newBzzTester(t *testing.T, n int, addr *bzzAddr, pp *p2ptest.TestPeerPool, spec *protocols.Spec, services func(Peer) error) *bzzTester { extraservices := func(p *bzzPeer) error { pp.Add(p) - p.DisconnectHook(func(err error) { - pp.Remove(p) - }) - if services != nil { - err := services(p) - if err != nil { - return err - } - } - return nil + defer pp.Remove(p) + return services(p) } - return newBzzBaseTester(t, n, addr, ct, extraservices) + return newBzzBaseTester(t, n, addr, spec, extraservices) } // should test handshakes in one exchange? parallelisation @@ -113,7 +136,11 @@ func (s *bzzTester) runHandshakes(ids ...*adapters.NodeId) { } func correctBzzHandshake(addr *bzzAddr) *bzzHandshake { - return &bzzHandshake{0, 322, addr} + return &bzzHandshake{ + Version: 0, + NetworkId: 322, + Addr: addr, + } } func TestBzzHandshakeNetworkIdMismatch(t *testing.T) { @@ -125,7 +152,7 @@ func TestBzzHandshakeNetworkIdMismatch(t *testing.T) { id := s.Ids[0] s.testHandshake( correctBzzHandshake(addr), - &bzzHandshake{0, 321, NewAddrFromNodeId(id)}, + &bzzHandshake{Version: 0, NetworkId: 321, Addr: NewAddrFromNodeId(id)}, &p2ptest.Disconnect{Peer: id, Error: fmt.Errorf("network id mismatch 321 (!= 322)")}, ) } @@ -139,7 +166,7 @@ func TestBzzHandshakeVersionMismatch(t *testing.T) { id := s.Ids[0] s.testHandshake( correctBzzHandshake(addr), - &bzzHandshake{1, 322, NewAddrFromNodeId(id)}, + &bzzHandshake{Version: 1, NetworkId: 322, Addr: NewAddrFromNodeId(id)}, &p2ptest.Disconnect{Peer: id, Error: fmt.Errorf("version mismatch 1 (!= 0)")}, ) } @@ -153,7 +180,7 @@ func TestBzzHandshakeSuccess(t *testing.T) { id := s.Ids[0] s.testHandshake( correctBzzHandshake(addr), - &bzzHandshake{0, 322, NewAddrFromNodeId(id)}, + &bzzHandshake{Version: 0, NetworkId: 322, Addr: NewAddrFromNodeId(id)}, ) } @@ -215,7 +242,7 @@ func TestBzzPeerPoolNotAdd(t *testing.T) { defer s.Stop() id := s.Ids[0] - s.testHandshake(correctBzzHandshake(addr), &bzzHandshake{0, 321, NewAddrFromNodeId(id)}, &p2ptest.Disconnect{Peer: id, Error: fmt.Errorf("network id mismatch 321 (!= 322)")}) + s.testHandshake(correctBzzHandshake(addr), &bzzHandshake{Version: 0, NetworkId: 321, Addr: NewAddrFromNodeId(id)}, &p2ptest.Disconnect{Peer: id, Error: fmt.Errorf("network id mismatch 321 (!= 322)")}) if pp.Has(id) { t.Fatalf("peer %v incorrectly added: %v", id, pp) } diff --git a/swarm/network/pss.go b/swarm/network/pss.go index fa1265982a..d713c1d764 100644 --- a/swarm/network/pss.go +++ b/swarm/network/pss.go @@ -121,6 +121,7 @@ type pssDigest uint32 // - a message cache to spot messages that previously have been forwarded type Pss struct { Overlay // we can get the overlayaddress from this + //peerPool map[pot.Address]map[PssTopic]*PssReadWriter // keep track of all virtual p2p.Peers we are currently speaking to peerPool map[pot.Address]map[PssTopic]p2p.MsgReadWriter // keep track of all virtual p2p.Peers we are currently speaking to handlers map[PssTopic]func([]byte, *p2p.Peer, []byte) error // topic and version based pss payload handlers @@ -162,6 +163,20 @@ func NewPss(k Overlay, params *PssParams) *Pss { } } +func (p *Pss) Run(peer *bzzPeer) error { + return peer.Run(p.HandleMsg) +} + +func (p *Pss) HandleMsg(m interface{}) error { + msg, ok := m.(*PssMsg) + if !ok { + return fmt.Errorf("unknown pss protocol message type: %T", m) + } + _ = msg + // TODO: handle the message + return nil +} + // enables to set address of node, to avoid backwards forwarding // // currently not in use as forwarder address is not known in the handler function hooked to the pss dispatcher. @@ -380,7 +395,7 @@ type PssReadWriter struct { RecipientOAddr pot.Address LastActive time.Time rw chan p2p.Msg - ct *protocols.CodeMap + spec *protocols.Spec topic *PssTopic } @@ -396,7 +411,7 @@ func (prw PssReadWriter) ReadMsg() (p2p.Msg, error) { // Implements p2p.MsgWriter func (prw PssReadWriter) WriteMsg(msg p2p.Msg) error { log.Trace(fmt.Sprintf("pssrw writemsg: %v", msg)) - ifc, found := prw.ct.GetInterface(msg.Code) + ifc, found := prw.spec.NewMsg(msg.Code) if !found { return fmt.Errorf("Writemsg couldn't find matching interface for code %d", msg.Code) } @@ -417,20 +432,20 @@ func (prw PssReadWriter) injectMsg(msg p2p.Msg) error { } // Convenience object for passing messages in and out of the p2p layer -type PssProtocol struct { +type pssProtocol struct { *Pss virtualProtocol *p2p.Protocol topic *PssTopic - ct *protocols.CodeMap + spec *protocols.Spec } // Constructor -func NewPssProtocol(pss *Pss, topic *PssTopic, ct *protocols.CodeMap, targetprotocol *p2p.Protocol) *PssProtocol { - pp := &PssProtocol{ +func NewPssProtocol(pss *Pss, topic *PssTopic, spec *protocols.Spec, targetprotocol *p2p.Protocol) *pssProtocol { + pp := &pssProtocol{ Pss: pss, virtualProtocol: targetprotocol, topic: topic, - ct: ct, + spec: spec, } return pp } @@ -438,18 +453,18 @@ func NewPssProtocol(pss *Pss, topic *PssTopic, ct *protocols.CodeMap, targetprot // Retrieves a convenience method for passing an incoming message into the p2p layer // // If the implementer wishes to use the p2p.Protocol (or p2p/protocols) message handling, this handler can be directly registered as a handler for the PssMsg structure -func (self *PssProtocol) GetHandler() func([]byte, *p2p.Peer, []byte) error { +func (self *pssProtocol) GetHandler() func([]byte, *p2p.Peer, []byte) error { return self.handle } -func (self *PssProtocol) handle(msg []byte, p *p2p.Peer, senderAddr []byte) error { +func (self *pssProtocol) handle(msg []byte, p *p2p.Peer, senderAddr []byte) error { hashoaddr := pot.NewHashAddressFromBytes(senderAddr).Address if !self.isActive(hashoaddr, *self.topic) { rw := &PssReadWriter{ Pss: self.Pss, RecipientOAddr: hashoaddr, rw: make(chan p2p.Msg), - ct: self.ct, + spec: self.spec, topic: self.topic, } self.Pss.AddPeer(p, hashoaddr, self.virtualProtocol.Run, *self.topic, rw) diff --git a/swarm/network/pss_test.go b/swarm/network/pss_test.go index 68b3f2abfa..502558b0a0 100644 --- a/swarm/network/pss_test.go +++ b/swarm/network/pss_test.go @@ -1,25 +1,9 @@ package network import ( - "context" - "encoding/hex" - "fmt" - "math/rand" - "net" - "net/http" "os" - "strconv" - "testing" - "time" - "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/log" - "github.com/ethereum/go-ethereum/node" - "github.com/ethereum/go-ethereum/p2p" - "github.com/ethereum/go-ethereum/p2p/protocols" - "github.com/ethereum/go-ethereum/p2p/simulations" - "github.com/ethereum/go-ethereum/p2p/simulations/adapters" - "github.com/ethereum/go-ethereum/rpc" ) const ( @@ -32,6 +16,7 @@ func init() { log.Root().SetHandler(h) } +/* // example protocol implementation peer // message handlers are methods of this // channels allow receipt reporting from p2p.Protocol message handler @@ -66,10 +51,6 @@ func (n *pssTestNode) Add(peer *bzzPeer) error { return err } -func (n *pssTestNode) hiveKeepAlive() <-chan time.Time { - return time.Tick(time.Millisecond * 300) -} - func (n *pssTestNode) triggerCheck() { go func() { n.trigger <- n.id }() } @@ -98,8 +79,9 @@ type pssTestService struct { func newPssTestService(t *testing.T, handlefunc func(interface{}) error, testnode *pssTestNode) *pssTestService { hp := NewHiveParams() - //hp.CallInterval = 250 - testnode.Hive = NewHive(hp, testnode.Pss.Overlay) + hp.KeepAliveInterval = 300 + bzz := NewBzz(testnode.OverlayAddr(), testnode.UnderlayAddr(), newTestStore()) + testnode.Hive = NewHive(hp, testnode.Pss.Overlay, bzz) return &pssTestService{ //nid := adapters.NewNodeId(addr.UnderlayAddr()) msgFunc: handlefunc, @@ -108,7 +90,7 @@ func newPssTestService(t *testing.T, handlefunc func(interface{}) error, testnod } func (self *pssTestService) Start(server p2p.Server) error { - return self.node.Hive.Start(server, self.node.hiveKeepAlive, nil) + return self.node.Hive.Start(server) } func (self *pssTestService) Stop() error { @@ -117,24 +99,13 @@ func (self *pssTestService) Stop() error { } func (self *pssTestService) Protocols() []p2p.Protocol { - ct := BzzCodeMap() - ct.Register(0, &PssMsg{}) - for _, m := range DiscoveryMsgs { - ct.Register(1, m) - } - - srv := func(p *bzzPeer) error { - p.Register(&PssMsg{}, self.msgFunc) - self.node.Add(p) - p.DisconnectHook(func(err error) { - self.node.Remove(p) - }) - return nil - } - - proto := NewBzz(self.node.OverlayAddr(), self.node.UnderlayAddr(), ct, srv, nil, nil) - - return []p2p.Protocol{*proto} + bzz := NewBzz(self.node.OverlayAddr(), self.node.UnderlayAddr(), newTestStore()) + return append(self.node.Hive.Protocols(), p2p.Protocol{ + Name: PssProtocolName, + Version: PssProtocolVersion, + Length: PssProtocol.Length(), + Run: bzz.RunProtocol(PssProtocol, self.Run), + }) } func (self *pssTestService) APIs() []rpc.API { @@ -149,6 +120,12 @@ func (self *pssTestService) APIs() []rpc.API { return nil } +func (self *pssTestService) Run(peer *bzzPeer) error { + self.node.Add(peer) + defer self.node.Remove(peer) + return peer.Run(self.msgFunc) +} + func TestPssCache(t *testing.T) { var err error to, _ := hex.DecodeString("08090a0b0c0d0e0f1011121314150001020304050607161718191a1b1c1d1e1f") @@ -301,8 +278,9 @@ func testPssFullRandom(t *testing.T, numsends int, numnodes int, numfullnodes in expectnodesids := []*adapters.NodeId{} // the nodes to expect on (needed by checker) expectnodesresults := make(map[*adapters.NodeId][]int) // which messages expect actually got - vct := protocols.NewCodeMap(protocolName, protocolVersion, ProtocolMaxMsgSize) - vct.Register(0, &pssTestPayload{}) + vct := protocols.NewCodeMap(map[uint64]interface{}{ + 0: pssTestPayload{}, + }) topic, _ := MakeTopic(protocolName, protocolVersion) trigger := make(chan *adapters.NodeId) @@ -536,15 +514,15 @@ func TestPssFullLinearEcho(t *testing.T) { return err } - /*for i, id := range ids { - var peerId *adapters.NodeId - if i != 0 { - peerId = ids[i-1] - if err := net.Connect(id, peerId); err != nil { - return err - } - } - }*/ + // for i, id := range ids { + // var peerId *adapters.NodeId + // if i != 0 { + // peerId = ids[i-1] + // if err := net.Connect(id, peerId); err != nil { + // return err + // } + // } + // } return nil } check = func(ctx context.Context, id *adapters.NodeId) (bool, error) { @@ -1108,3 +1086,4 @@ func (ptp *pssTestPeer) SimpleHandlePssPayload(msg interface{}) error { return nil } +*/ diff --git a/swarm/network/simulations/discovery/discovery_test.go b/swarm/network/simulations/discovery/discovery_test.go index 22d97efec0..c436f20025 100644 --- a/swarm/network/simulations/discovery/discovery_test.go +++ b/swarm/network/simulations/discovery/discovery_test.go @@ -9,11 +9,10 @@ import ( "time" "github.com/ethereum/go-ethereum/log" - p2pnode "github.com/ethereum/go-ethereum/node" + "github.com/ethereum/go-ethereum/node" "github.com/ethereum/go-ethereum/p2p" "github.com/ethereum/go-ethereum/p2p/simulations" "github.com/ethereum/go-ethereum/p2p/simulations/adapters" - "github.com/ethereum/go-ethereum/rpc" "github.com/ethereum/go-ethereum/swarm/network" ) @@ -22,9 +21,7 @@ import ( const serviceName = "discovery" var services = adapters.Services{ - serviceName: func(id *adapters.NodeId, snapshot []byte) p2pnode.Service { - return newNode(id) - }, + serviceName: newService, } func init() { @@ -70,7 +67,7 @@ func testDiscoverySimulation(t *testing.T, adapter adapters.NodeAdapter) { for i := 0; i < nodeCount; i++ { node, err := net.NewNode() if err != nil { - t.Fatalf("error starting node %s: %s", node.ID().Label(), err) + t.Fatalf("error starting node: %s", err) } if err := net.Start(node.ID()); err != nil { t.Fatalf("error starting node %s: %s", node.ID().Label(), err) @@ -179,70 +176,24 @@ func triggerChecks(trigger chan *adapters.NodeId, net *simulations.Network, id * return nil } -type node struct { - *network.Hive +func newService(id *adapters.NodeId, snapshot []byte) node.Service { + addr := network.NewAddrFromNodeId(id) - protocol *p2p.Protocol -} - -func newNode(id *adapters.NodeId) *node { - addr := network.NewPeerAddrFromNodeId(id) - kademlia := newKademlia(addr.OverlayAddr()) - hive := newHive(kademlia) - codeMap := network.BzzCodeMap(network.DiscoveryMsgs...) - node := &node{Hive: hive} - services := func(peer network.Peer) error { - discoveryPeer := network.NewDiscovery(peer, kademlia) - node.Add(discoveryPeer) - peer.DisconnectHook(func(err error) { - node.Remove(discoveryPeer) - }) - return nil + config := &network.BzzConfig{ + OverlayAddr: addr.Over(), + UnderlayAddr: addr.Under(), + KadParams: network.NewKadParams(), + HiveParams: network.NewHiveParams(), } - node.protocol = network.Bzz(addr.OverlayAddr(), addr.UnderlayAddr(), codeMap, services, nil, nil) - return node -} - -func newKademlia(overlayAddr []byte) *network.Kademlia { - params := network.NewKadParams() - params.MinProxBinSize = 2 - params.MaxBinSize = 3 - params.MinBinSize = 1 - params.MaxRetries = 1000 - params.RetryExponent = 2 - params.RetryInterval = 1000000 - - return network.NewKademlia(overlayAddr, params) -} - -func newHive(kademlia *network.Kademlia) *network.Hive { - params := network.NewHiveParams() - params.CallInterval = 5000 - - return network.NewHive(params, kademlia) -} - -func (n *node) Protocols() []p2p.Protocol { - return []p2p.Protocol{*n.protocol} -} - -func (n *node) APIs() []rpc.API { - return []rpc.API{{ - Namespace: "hive", - Version: "1.0", - Service: n.Hive, - }} -} - -func (n *node) Start(server p2p.Server) error { - return n.Hive.Start(server, n.hiveKeepAlive) -} - -func (n *node) Stop() error { - n.Hive.Stop() - return nil -} - -func (n *node) hiveKeepAlive() <-chan time.Time { - return time.Tick(time.Second) + + config.KadParams.MinProxBinSize = 2 + config.KadParams.MaxBinSize = 3 + config.KadParams.MinBinSize = 1 + config.KadParams.MaxRetries = 1000 + config.KadParams.RetryExponent = 2 + config.KadParams.RetryInterval = 1000000 + + config.HiveParams.KeepAliveInterval = time.Second + + return network.NewBzz(config) }