p2p/protocols: Refactor

Signed-off-by: Lewis Marshall <lewis@lmars.net>
This commit is contained in:
Lewis Marshall 2017-05-14 00:40:32 -07:00
parent c870e2be26
commit 91c198778c
18 changed files with 737 additions and 844 deletions

View file

@ -47,6 +47,19 @@ const (
maxResolveDelay = time.Hour 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. // dialstate schedules dials and discovery lookups.
// it get's a chance to compute new tasks on every iteration // it get's a chance to compute new tasks on every iteration
// of the main loop in server.run. // of the main loop in server.run.
@ -318,14 +331,13 @@ func (t *dialTask) resolve(srv *server) bool {
// dial performs the actual connection attempt. // dial performs the actual connection attempt.
func (t *dialTask) dial(srv *server, dest *discover.Node) bool { 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(dest)
fd, err := srv.Dialer.Dial("tcp", addr.String())
if err != nil { if err != nil {
log.Trace("Dial error", "task", t, "err", err) log.Trace("Dial error", "task", t, "err", err)
return false return false
} }
mfd := newMeteredConn(fd, false) mfd := newMeteredConn(fd, false)
srv.setupConn(mfd, t.flags, dest) srv.SetupConn(mfd, t.flags, dest)
return true return true
} }

View file

@ -597,8 +597,8 @@ func TestDialResolve(t *testing.T) {
} }
// Now run the task, it should resolve the ID once. // Now run the task, it should resolve the ID once.
config := Config{Dialer: &net.Dialer{Deadline: time.Now().Add(-5 * time.Minute)}} config := Config{Dialer: TCPDialer{&net.Dialer{Deadline: time.Now().Add(-5 * time.Minute)}}}
srv := &server{ntab: table, Config: config} srv := &Server{ntab: table, Config: config}
tasks[0].Do(srv) tasks[0].Do(srv)
if !reflect.DeepEqual(table.resolveCalls, []discover.NodeID{dest.ID}) { if !reflect.DeepEqual(table.resolveCalls, []discover.NodeID{dest.ID}) {
t.Fatalf("wrong resolve calls, got %v", table.resolveCalls) t.Fatalf("wrong resolve calls, got %v", table.resolveCalls)

View file

@ -30,13 +30,13 @@ Standard protocol supports:
package protocols package protocols
import ( import (
"context"
"fmt" "fmt"
"reflect" "reflect"
"time" "sync"
"github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/log"
"github.com/ethereum/go-ethereum/p2p" "github.com/ethereum/go-ethereum/p2p"
"github.com/ethereum/go-ethereum/p2p/discover"
) )
// error codes used by this protocol scheme // error codes used by this protocol scheme
@ -109,155 +109,104 @@ func errorf(code int, format string, params ...interface{}) *Error {
return self return self
} }
// implements the code table spec // Spec is a protocol specification including its name and version as well as
// listing the message codes and types etc // the types of messages which are exchanged
// and further metadata about the protocol type Spec struct {
type CodeMap struct { // Name is the name of the protocol, often a three-letter word
Name string // name of the protocol Name string
Version uint // version
MaxMsgSize int // max length of message payload size // Version is the version number of the protocol
codepos int // the subsequent code Version uint
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 // 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) { func (s *Spec) init() {
typ, found := self.codes[code] s.initOnce.Do(func() {
if !found { 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 return nil, false
} }
val := reflect.New(typ) return reflect.New(typ).Interface(), true
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
} }
// A Peer represents a remote peer or protocol instance that is running on a peer connection with // A Peer represents a remote peer or protocol instance that is running on a peer connection with
// a remote peer // a remote peer
type Peer struct { type Peer struct {
ct *CodeMap // CodeMap for the protocol
*p2p.Peer // the p2p.Peer object representing the remote *p2p.Peer // the p2p.Peer object representing the remote
rw p2p.MsgReadWriter // p2p.MsgReadWriter to send messages to and read messages from 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 spec *Spec
Errc chan error Errc chan error
ready chan bool // blocking send until handshake finishes wErrc chan error // write error channel
} }
// NewPeer returns a new peer // NewPeer returns a new peer
// this constructor is called by the p2p.Protocol#Run function // 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 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 // the third argument is the CodeMap describing the protocol messages and options
func NewPeer(p *p2p.Peer, ct *CodeMap, rw p2p.MsgReadWriter) *Peer { func NewPeer(p *p2p.Peer, rw p2p.MsgReadWriter, spec *Spec) *Peer {
ready := make(chan bool)
defer close(ready)
return &Peer{ return &Peer{
ct: ct,
Peer: p, Peer: p,
rw: rw, rw: rw,
spec: spec,
Errc: make(chan error), Errc: make(chan error),
ready: ready, wErrc: make(chan error),
handlers: make(map[reflect.Type][]func(interface{}) 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 // Run starts the forever loop that handles incoming messages
// called within the p2p.Protocol#Run function // called within the p2p.Protocol#Run function
func (self *Peer) Run() error { func (self *Peer) Run(handler func(msg interface{}) error) error {
go func() { go func() {
for { for {
_, err := self.handleIncoming() if err := self.handleIncoming(handler); err != nil {
if err != nil {
self.Errc <- err self.Errc <- err
return return
} }
} }
}() }()
err := <-self.Errc return <-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
} }
// Drop disconnects a peer. // 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 // 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 // but often just used to forward and push messages to directly connected peers
func (self *Peer) Send(msg interface{}) error { func (self *Peer) Send(msg interface{}) error {
<-self.ready code, found := self.spec.GetCode(msg)
return self.send(msg)
}
func (self *Peer) send(msg interface{}) error {
code, found := self.ct.GetCode(msg)
if !found { if !found {
return errorf(ErrInvalidMsgType, "%v", code) return errorf(ErrInvalidMsgType, "%v", code)
} }
log.Trace(fmt.Sprintf("=> msg #%d TO %v : %v", code, self.ID(), msg)) log.Trace(fmt.Sprintf("=> msg #%d TO %v : %v", code, self.ID(), msg))
return p2p.Send(self.rw, code, 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
})
} }
// handleIncoming(code) // 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 // 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, // checks message size, out-of-range message codes, handles decoding with reflection,
// call handlers as callback onside // 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() msg, err := self.rw.ReadMsg()
if err != nil { if err != nil {
return nil, err return err
} }
log.Trace(fmt.Sprintf("<= %v", msg)) log.Trace(fmt.Sprintf("<= %v", msg))
// make sure that the payload has been fully consumed // make sure that the payload has been fully consumed
defer msg.Discard() defer msg.Discard()
if msg.Size > uint32(self.ct.MaxMsgSize) { if msg.Size > self.spec.MaxMsgSize {
return nil, errorf(ErrMsgTooLong, "%v > %v", msg.Size, self.ct.MaxMsgSize) return errorf(ErrMsgTooLong, "%v > %v", msg.Size, self.spec.MaxMsgSize)
} }
// check if the message code is correct val, ok := self.spec.NewMsg(msg.Code)
maxMsgCode := uint(len(self.ct.messages)) if !ok {
if msg.Code >= uint64(maxMsgCode) { return errorf(ErrInvalidMsgCode, "%v", msg.Code)
return nil, errorf(ErrInvalidMsgCode, "%v (>=%v)", msg.Code, maxMsgCode)
} }
if err := msg.Decode(val); err != nil {
// it is safe to be unsafe here return errorf(ErrDecode, "<= %v: %v", msg, err)
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)
} }
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 // call the registered handler callbacks
// a registered callback take the decoded message as argument as an interface // a registered callback take the decoded message as argument as an interface
// which the handler is supposed to cast to the appropriate type // 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 // 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 // chosen based on the proper type in the first place
handlers := self.handlers[typ] if err := handle(val); err != nil {
if len(handlers) == 0 { return errorf(ErrHandler, "(msg code %v): %v", msg.Code, err)
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)
} }
} return nil
}
return req.Interface(), nil
} }
// Handshake initiates a handshake on the peer connection // Handshake initiates a handshake on the peer connection
// * the argument is the local handshake to be sent to the remote peer // * the argument is the local handshake to be sent to the remote peer
// * expects a remote handshake back of the same type // * expects a remote handshake back of the same type
// returns the remote hs and an error // returns the remote hs and an error
func (self *Peer) Handshake(hs interface{}, handshakeTimeout time.Duration) (rhs interface{}, err error) { func (self *Peer) Handshake(ctx context.Context, hs interface{}) (interface{}, error) {
typ := reflect.TypeOf(hs) if _, ok := self.spec.GetCode(hs); !ok {
_, found := self.ct.messages[typ] return nil, errorf(ErrHandshake, "unknown handshake message type: %T", hs)
if !found {
return nil, errorf(ErrHandshake, "unknown handshake message type: %v", typ)
} }
self.ready = make(chan bool) errc := make(chan error, 2)
received := make(chan bool)
defer close(self.ready)
go func() { go func() {
defer close(received) if err := self.Send(hs); err != nil {
// receiving and validating remote handshake, expect code errc <- errorf(ErrHandshake, "cannot send: %v", err)
rhs, err = self.handleIncoming()
if err != nil {
err = errorf(ErrHandshake, "'%v': %v", self.ct.Name, err)
} }
}() }()
if e := self.send(hs); e != nil { hsc := make(chan interface{})
return nil, errorf(ErrHandshake, "cannot send: %v", e) 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 { select {
case <-received: case rhs := <-hsc:
case <-time.NewTimer(handshakeTimeout).C: return rhs, nil
err = errorf(ErrHandshake, "timeout after %v", handshakeTimeout) case <-ctx.Done():
return nil, ctx.Err()
case err := <-errc:
return nil, err
} }
return rhs, err
} }

View file

@ -1,6 +1,8 @@
package protocols package protocols
import ( import (
"context"
"errors"
"fmt" "fmt"
"os" "os"
"testing" "testing"
@ -13,7 +15,7 @@ import (
) )
func init() { 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 // handshake message type
@ -55,28 +57,26 @@ const networkId = "420"
// newProtocol sets up a protocol // newProtocol sets up a protocol
// the run function here demonstrates a typical protocol using peerPool, handshake // the run function here demonstrates a typical protocol using peerPool, handshake
// and messages registered to handlers // and messages registered to handlers
func newProtocol(pp *p2ptest.TestPeerPool) adapters.RunProtocol { func newProtocol(pp *p2ptest.TestPeerPool) func(*p2p.Peer, p2p.MsgReadWriter) error {
ct := NewCodeMap("test", 42, 1024) spec := &Spec{
ct.Register(0, &protoHandshake{}, &hs0{}, &kill{}, &drop{}) Name: "test",
Version: 42,
MaxMsgSize: 10 * 1024,
Messages: []interface{}{
protoHandshake{},
hs0{},
kill{},
drop{},
},
}
return func(p *p2p.Peer, rw p2p.MsgReadWriter) error { return func(p *p2p.Peer, rw p2p.MsgReadWriter) error {
peer := NewPeer(p, ct, rw) peer := NewPeer(p, rw, spec)
// 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")
})
// initiate one-off protohandshake and check validity // initiate one-off protohandshake and check validity
phs := &protoHandshake{ct.Version, networkId} ctx, cancel := context.WithTimeout(context.Background(), time.Second)
hs, err := peer.Handshake(phs, time.Second) defer cancel()
phs := &protoHandshake{42, networkId}
hs, err := peer.Handshake(ctx, phs)
if err != nil { if err != nil {
return err return err
} }
@ -88,7 +88,7 @@ func newProtocol(pp *p2ptest.TestPeerPool) adapters.RunProtocol {
lhs := &hs0{42} lhs := &hs0{42}
// module handshake demonstrating a simple repeatable exchange of same-type message // 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 { if err != nil {
return err 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) return fmt.Errorf("handshake mismatch remote %v > local %v", rmhs.C, lhs.C)
} }
peer.Register(lhs, func(msg interface{}) error { handle := func(msg interface{}) error {
rhs := msg.(*hs0) switch msg := msg.(type) {
case *protoHandshake:
return errors.New("duplicate handshake")
case *hs0:
rhs := msg
if rhs.C > lhs.C { if rhs.C > lhs.C {
return fmt.Errorf("handshake mismatch remote %v > local %v", rhs.C, lhs.C) return fmt.Errorf("handshake mismatch remote %v > local %v", rhs.C, lhs.C)
} }
lhs.C += rhs.C lhs.C += rhs.C
return peer.Send(lhs) 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)
}
}
log.Trace(fmt.Sprintf("adding peer %v", peer)) log.Trace(fmt.Sprintf("adding peer %v", peer))
pp.Add(peer) pp.Add(peer)
defer pp.Remove(peer) defer pp.Remove(peer)
err = peer.Run() err = peer.Run(handle)
log.Trace(fmt.Sprintf("peer %v protocol quitting: %v", peer, err)) log.Trace(fmt.Sprintf("peer %v protocol quitting: %v", peer, err))
return err return err

View file

@ -131,7 +131,7 @@ type Config struct {
// If Dialer is set to a non-nil value, the given Dialer // If Dialer is set to a non-nil value, the given Dialer
// is used to dial outbound peer connections. // 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. // If NoDial is true, the server will not dial any peers.
NoDial bool `toml:",omitempty"` NoDial bool `toml:",omitempty"`
@ -144,6 +144,7 @@ type Config struct {
type Server interface { type Server interface {
Start() error Start() error
Stop() error Stop() error
SetupConn(net.Conn, connFlag, *discover.Node)
AddPeer(node *discover.Node) AddPeer(node *discover.Node)
RemovePeer(node *discover.Node) RemovePeer(node *discover.Node)
SubscribeEvents(ch chan *PeerEvent) event.Subscription SubscribeEvents(ch chan *PeerEvent) event.Subscription
@ -385,7 +386,7 @@ func (srv *server) Start() (err error) {
srv.newTransport = newRLPX srv.newTransport = newRLPX
} }
if srv.Dialer == nil { if srv.Dialer == nil {
srv.Dialer = &net.Dialer{Timeout: defaultDialTimeout} srv.Dialer = TCPDialer{&net.Dialer{Timeout: defaultDialTimeout}}
} }
srv.quit = make(chan struct{}) srv.quit = make(chan struct{})
srv.addpeer = make(chan *conn) 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 // Spawn the handler. It will give the slot back when the connection
// has been established. // has been established.
go func() { go func() {
srv.setupConn(fd, inboundConn, nil) srv.SetupConn(fd, inboundConn, nil)
slots <- struct{}{} slots <- struct{}{}
}() }()
} }
@ -706,7 +707,7 @@ func (srv *server) listenLoop() {
// setupConn runs the handshakes and attempts to add the connection // setupConn runs the handshakes and attempts to add the connection
// as a peer. It returns when the connection has been added as a peer // as a peer. It returns when the connection has been added as a peer
// or the handshakes have failed. // 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. // Prevent leftover pending conns from entering the handshake.
srv.lock.Lock() srv.lock.Lock()
running := srv.running running := srv.running

View file

@ -17,14 +17,13 @@
package adapters package adapters
import ( import (
"context"
"errors" "errors"
"fmt" "fmt"
"math"
"net" "net"
"sync" "sync"
"github.com/ethereum/go-ethereum/event" "github.com/ethereum/go-ethereum/event"
"github.com/ethereum/go-ethereum/log"
"github.com/ethereum/go-ethereum/node" "github.com/ethereum/go-ethereum/node"
"github.com/ethereum/go-ethereum/p2p" "github.com/ethereum/go-ethereum/p2p"
"github.com/ethereum/go-ethereum/p2p/discover" "github.com/ethereum/go-ethereum/p2p/discover"
@ -73,15 +72,28 @@ func (s *SimAdapter) NewNode(config *NodeConfig) (Node, error) {
node := &SimNode{ node := &SimNode{
Id: id, Id: id,
config: config,
adapter: s, adapter: s,
serviceFunc: serviceFunc, serviceFunc: serviceFunc,
peers: make(map[discover.NodeID]MsgReadWriteCloser),
dropPeers: make(chan struct{}),
} }
s.nodes[id.NodeID] = node s.nodes[id.NodeID] = node
return node, nil 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 // GetNode returns the node with the given ID if it exists
func (s *SimAdapter) GetNode(id discover.NodeID) (*SimNode, bool) { func (s *SimAdapter) GetNode(id discover.NodeID) (*SimNode, bool) {
s.mtx.RLock() s.mtx.RLock()
@ -90,14 +102,6 @@ func (s *SimAdapter) GetNode(id discover.NodeID) (*SimNode, bool) {
return node, ok 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 // SimNode is an in-memory node which connects to other SimNodes using an
// in-memory p2p.MsgReadWriter pipe, running an underlying service protocol // in-memory p2p.MsgReadWriter pipe, running an underlying service protocol
// directly over that pipe. // directly over that pipe.
@ -107,17 +111,13 @@ type MsgReadWriteCloser interface {
type SimNode struct { type SimNode struct {
lock sync.RWMutex lock sync.RWMutex
Id *NodeId Id *NodeId
config *NodeConfig
adapter *SimAdapter adapter *SimAdapter
running node.Service
serviceFunc ServiceFunc serviceFunc ServiceFunc
peers map[discover.NodeID]MsgReadWriteCloser node *node.Node
peerFeed event.Feed running node.Service
client *rpc.Client client *rpc.Client
rpcMux *rpcMux rpcMux *rpcMux
// dropPeers is used to force peer disconnects when
// the node is stopped
dropPeers chan struct{}
} }
// Addr returns the node's discovery address // Addr returns the node's discovery address
@ -127,7 +127,7 @@ func (self *SimNode) Addr() []byte {
// Node returns a discover.Node representing the SimNode // Node returns a discover.Node representing the SimNode
func (self *SimNode) Node() *discover.Node { 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 // 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 return nil
} }
// Start initializes the service, starts the RPC handler and then starts // Snapshot creates a snapshot of the service state by calling the
// the service // simulation_snapshot RPC method
func (self *SimNode) Start(snapshot []byte) error { func (self *SimNode) Snapshot() ([]byte, error) {
service := self.serviceFunc(self.Id, snapshot) 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 // Start starts the RPC handler and the underlying service
// multiple protocols on the same peer is extra effort, and we don't func (self *SimNode) Start(snapshot []byte) error {
// currently run any simulations which run multiple protocols) self.lock.Lock()
if len(service.Protocols()) != 1 { defer self.lock.Unlock()
return errors.New("service must have a single protocol") if self.node != nil {
return errors.New("node already started")
} }
self.dropPeers = make(chan struct{}) newService := func(ctx *node.ServiceContext) (node.Service, error) {
if err := self.startRPC(service); err != nil { 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 return err
} }
self.running = service
return service.Start(&simServer{self})
}
// simServer wraps a SimNode but modifies the Start method signature so that if err := node.Register(newService); err != nil {
// it implements the p2p.Server interface (the Start method is never actually return err
// 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")
} }
// add SimAdminAPI so that the network can call the if err := node.Start(); err != nil {
// AddPeer, RemovePeer and PeerEvents RPC methods return err
apis := append(service.APIs(), []rpc.API{
{
Namespace: "admin",
Version: "1.0",
Service: &SimAdminAPI{self},
},
}...)
// 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 // create an in-process RPC multiplexer
@ -241,197 +215,55 @@ func (self *SimNode) startRPC(service node.Service) error {
// create an in-process RPC client // create an in-process RPC client
self.client = self.rpcMux.Client() self.client = self.rpcMux.Client()
self.node = node
return nil return nil
} }
// stopRPC closes the node's RPC client func (self *SimNode) Stop() error {
func (self *SimNode) stopRPC() {
self.lock.Lock() self.lock.Lock()
defer self.lock.Unlock() defer self.lock.Unlock()
if self.client != nil { if self.node == nil {
self.client.Close() return nil
self.client = nil
self.rpcMux = 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 // Service returns the underlying node.Service
// p2p.MsgReadWriter pipe and closing it (which will cause both the local func (self *SimNode) Service() node.Service {
// and peer Protocol.Run functions to exit)
func (self *SimNode) RemovePeer(peer *discover.Node) {
self.lock.Lock() self.lock.Lock()
defer self.lock.Unlock() defer self.lock.Unlock()
peerRW, exists := self.peers[peer.ID] return self.running
if !exists {
return
}
peerRW.Close()
delete(self.peers, peer.ID)
log.Trace(fmt.Sprintf("dropped peer %v", peer.ID))
} }
// AddPeer adds the given node as a peer by creating a p2p.MsgReadWriter pipe func (self *SimNode) Server() *p2p.Server {
// and running both the local and peer's Protocol.Run function over the pipe
func (self *SimNode) AddPeer(peer *discover.Node) {
self.lock.Lock() self.lock.Lock()
defer self.lock.Unlock() defer self.lock.Unlock()
if _, exists := self.peers[peer.ID]; exists { if self.node == nil {
return return nil
} }
peerNode, exists := self.adapter.GetNode(peer.ID) return self.node.Server()
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)
} }
// SubscribeEvents subscribes the given channel to p2p peer events
func (self *SimNode) SubscribeEvents(ch chan *p2p.PeerEvent) event.Subscription { 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 { func (self *SimNode) NodeInfo() *p2p.NodeInfo {
self.lock.Lock() server := self.Server()
defer self.lock.Unlock() if server == nil {
info := &p2p.NodeInfo{ return &p2p.NodeInfo{
ID: self.Id.String(), ID: self.Id.String(),
Enode: self.Node().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
} }
} }
return info return server.NodeInfo()
}
// 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
} }

View file

@ -13,7 +13,7 @@ import (
) )
type ProtocolSession struct { type ProtocolSession struct {
*adapters.SimNode p2p.Server
Ids []*adapters.NodeId Ids []*adapters.NodeId
adapter *adapters.SimAdapter adapter *adapters.SimAdapter

View file

@ -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 { func NewProtocolTester(t *testing.T, id *adapters.NodeId, n int, run func(*p2p.Peer, p2p.MsgReadWriter) error) *ProtocolTester {
services := map[string]adapters.ServiceFunc{ services := map[string]adapters.ServiceFunc{
"test": func(id *adapters.NodeId) node.Service { "test": func(id *adapters.NodeId, _ []byte) node.Service {
return &testNode{run} return &testNode{run}
}, },
"mock": func(id *adapters.NodeId) node.Service { "mock": func(id *adapters.NodeId, _ []byte) node.Service {
return newMockNode() 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) events := make(chan *p2p.PeerEvent, 1000)
node.SubscribeEvents(events) node.SubscribeEvents(events)
ps := &ProtocolSession{ ps := &ProtocolSession{
SimNode: node, Server: node.Server(),
Ids: peerIDs, Ids: peerIDs,
adapter: adapter, adapter: adapter,
events: events, events: events,
@ -86,7 +86,10 @@ type testNode struct {
} }
func (t *testNode) Protocols() []p2p.Protocol { 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 { func (t *testNode) APIs() []rpc.API {

View file

@ -8,13 +8,6 @@ import (
// discovery bzz overlay extension doing peer relaying // discovery bzz overlay extension doing peer relaying
// messages related to peer discovery
var DiscoveryMsgs = []interface{}{
&getPeersMsg{},
&peersMsg{},
&subPeersMsg{},
}
type discPeer struct { type discPeer struct {
*bzzPeer *bzzPeer
overlay Overlay overlay Overlay
@ -32,14 +25,26 @@ func NewDiscovery(p *bzzPeer, o Overlay) *discPeer {
peers: make(map[string]bool), peers: make(map[string]bool),
} }
self.seen(self) self.seen(self)
p.Register(&peersMsg{}, self.handlePeersMsg)
p.Register(&getPeersMsg{}, self.handleGetPeersMsg)
p.Register(&subPeersMsg{}, self.handleSubPeersMsg)
return self 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. // NotifyPeer notifies the receiver remote end of a peer p or PO po.
// callback for overlay driver // callback for overlay driver
func (self *discPeer) NotifyPeer(p OverlayPeer, po uint8) error { 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) return fmt.Sprintf("%T: request peers > PO%02d. ", self, self.ProxLimit)
} }
func (self *discPeer) handleSubPeersMsg(msg interface{}) error { func (self *discPeer) handleSubPeersMsg(msg *subPeersMsg) error {
spm := msg.(*subPeersMsg) self.proxLimit = msg.ProxLimit
self.proxLimit = spm.ProxLimit
if !self.sentPeers { if !self.sentPeers {
var peers []*bzzAddr var peers []*bzzAddr
self.overlay.EachConn(self.Over(), 255, func(p OverlayConn, po int, isproxbin bool) bool { 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) // 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 // list of nodes ([]PeerAddr in peersMsg) is added to the overlay db using the
// Register interface method // Register interface method
func (self *discPeer) handlePeersMsg(msg interface{}) error { func (self *discPeer) handlePeersMsg(msg *peersMsg) error {
// register all addresses // register all addresses
as := msg.(*peersMsg).Peers if len(msg.Peers) == 0 {
if len(as) == 0 {
log.Debug(fmt.Sprintf("whoops, no peers in incoming peersMsg from %v", self)) log.Debug(fmt.Sprintf("whoops, no peers in incoming peersMsg from %v", self))
return nil return nil
} }
var c chan OverlayAddr var c chan OverlayAddr
go func() { go func() {
for _, a := range as { for _, a := range msg.Peers {
self.seen(a) self.seen(a)
c <- a c <- a
} }
@ -161,18 +164,17 @@ func (self *discPeer) handlePeersMsg(msg interface{}) error {
// peers suggestions are retrieved from the overlay topology driver // peers suggestions are retrieved from the overlay topology driver
// using the EachConn interface iterator method // using the EachConn interface iterator method
// peers sent are remembered throughout a session and not sent twice // 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 var peers []*bzzAddr
req := msg.(*getPeersMsg)
i := 0 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++ i++
// only send peers we have not sent before in this session // only send peers we have not sent before in this session
a := ToAddr(p) a := ToAddr(p)
if self.seen(a) { if self.seen(a) {
peers = append(peers, a) peers = append(peers, a)
} }
return len(peers) < int(req.Max) return len(peers) < int(msg.Max)
}) })
if len(peers) == 0 { if len(peers) == 0 {
log.Debug(fmt.Sprintf("no peers found for %v", self)) log.Debug(fmt.Sprintf("no peers found for %v", self))

View file

@ -16,22 +16,18 @@ import (
func TestDiscovery(t *testing.T) { func TestDiscovery(t *testing.T) {
addr := RandomAddr() addr := RandomAddr()
to := NewKademlia(addr.OAddr, NewKadParams()) to := NewKademlia(addr.OAddr, NewKadParams())
ct := BzzCodeMap(DiscoveryMsgs...)
services := func(p *bzzPeer) error { run := func(p *bzzPeer) error {
dp := NewDiscovery(p, to) dp := NewDiscovery(p, to)
to.On(dp) to.On(p)
defer to.Off(p)
log.Trace(fmt.Sprintf("kademlia on %v", p)) log.Trace(fmt.Sprintf("kademlia on %v", p))
p.DisconnectHook(func(err error) { return p.Run(dp.HandleMsg)
to.Off(p)
})
return nil
} }
s := newBzzBaseTester(t, 1, addr, ct, services) s := newBzzBaseTester(t, 1, addr, DiscoveryProtocol, run)
defer s.Stop() defer s.Stop()
s.runHandshakes()
s.TestExchanges(p2ptest.Exchange{ s.TestExchanges(p2ptest.Exchange{
Label: "outgoing SubPeersMsg", Label: "outgoing SubPeersMsg",
Expects: []p2ptest.Expect{ Expects: []p2ptest.Expect{

View file

@ -34,7 +34,7 @@ it uses an Overlay Topology driver (e.g., generic kademlia nodetable)
to find best peer list for any target to find best peer list for any target
this is used by the netstore to search for content in the swarm 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 and relay the peer request process to the Overlay module
peer connections and disconnections are reported and registered peer connections and disconnections are reported and registered
@ -57,23 +57,19 @@ type Overlay interface {
BaseAddr() []byte 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 // Hive implements the PeerPool interface
type Hive struct { type Hive struct {
*HiveParams // settings *HiveParams // settings
Overlay // the overlay topology driver Overlay // the overlay topology driver
RW ReadWriter // ReadWriter store Store
// bookkeeping // bookkeeping
lock sync.Mutex lock sync.Mutex
quit chan bool quit chan bool
toggle chan bool toggle chan bool
more chan bool more chan bool
newTicker func() hiveTicker
} }
// HiveParams holds the config options to hive // HiveParams holds the config options to hive
@ -81,7 +77,7 @@ type HiveParams struct {
Discovery bool // if want discovery of not Discovery bool // if want discovery of not
PeersBroadcastSetSize uint8 // how many peers to use when relaying PeersBroadcastSetSize uint8 // how many peers to use when relaying
MaxPeersPerRequest uint8 // max size for peer address batches MaxPeersPerRequest uint8 // max size for peer address batches
CallInterval uint // polling interval fir=== KeepAliveInterval time.Duration
} }
// NewHiveParams returns hive config with only the // NewHiveParams returns hive config with only the
@ -90,17 +86,18 @@ func NewHiveParams() *HiveParams {
Discovery: true, Discovery: true,
PeersBroadcastSetSize: 2, PeersBroadcastSetSize: 2,
MaxPeersPerRequest: 5, MaxPeersPerRequest: 5,
CallInterval: 1000, KeepAliveInterval: time.Second,
} }
} }
// Hive constructor embeds both arguments // Hive constructor embeds both arguments
// HiveParams: config parameters // HiveParams: config parameters
// Overlay: Topology Driver Interface // Overlay: Topology Driver Interface
func NewHive(params *HiveParams, overlay Overlay) *Hive { func NewHive(params *HiveParams, overlay Overlay, store Store) *Hive {
return &Hive{ return &Hive{
HiveParams: params, HiveParams: params,
Overlay: overlay, 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 // these are called on the p2p.Server which runs on the node
// af() returns an arbitrary ticker channel // af() returns an arbitrary ticker channel
// rw is a read writer for json configs // rw is a read writer for json configs
func (self *Hive) Start(server p2p.Server, af func() <-chan time.Time, rw ReadWriter) error { func (self *Hive) Start(server p2p.Server) error {
if rw != nil { if self.store != nil {
if err := self.loadPeers(); err != nil { if err := self.loadPeers(); err != nil {
return err 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) self.quit = make(chan bool)
log.Debug("hive started") log.Debug("hive started")
// this loop is doing bootstrapping and maintains a healthy table // this loop is doing bootstrapping and maintains a healthy table
go self.keepAlive(af) go self.keepAlive()
go func() { go func() {
// each iteration, ask kademlia about most preferred peer // each iteration, ask kademlia about most preferred peer
for more := range self.more { 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 // Stop terminates the updateloop and saves the peers
func (self *Hive) Stop() { func (self *Hive) Stop() {
if self.RW != nil { if self.store != nil {
self.savePeers() self.savePeers()
} }
// closing toggle channel quits the updateloop // closing toggle channel quits the updateloop
close(self.quit) close(self.quit)
} }
// default ticker, tickinterval is taken from KadParams.CallInterval func (self *Hive) Run(peer *bzzPeer) error {
func (self *Hive) ticker() <-chan time.Time { discPeer := NewDiscovery(peer, self)
return time.NewTicker(time.Duration(self.CallInterval) * time.Millisecond).C self.On(discPeer)
defer self.Off(discPeer)
return peer.Run(discPeer.HandleMsg)
} }
// Add is called at the end of a successful protocol handshake // Add is called at the end of a successful protocol handshake
@ -242,27 +241,48 @@ func ToAddr(pa OverlayPeer) *bzzAddr {
return pa.(*bzzPeer).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 // keepAlive is a forever loop
// in its awake state it periodically triggers connection attempts // in its awake state it periodically triggers connection attempts
// by writing to self.more until Kademlia Table is saturated // by writing to self.more until Kademlia Table is saturated
// wake state is toggled by writing to self.toggle // wake state is toggled by writing to self.toggle
// it goes to sleep mode if table is saturated // it goes to sleep mode if table is saturated
// it restarts if the table becomes non-full again due to disconnections // it restarts if the table becomes non-full again due to disconnections
func (self *Hive) keepAlive(af func() <-chan time.Time) { func (self *Hive) keepAlive() {
log.Trace("keep alive loop started") if self.newTicker == nil {
alarm := af() self.newTicker = func() hiveTicker {
return &timeTicker{time.NewTicker(self.KeepAliveInterval)}
}
}
ticker := self.newTicker()
tick := ticker.Ch()
for { for {
select { select {
case <-alarm: case <-tick:
log.Trace("wake up: make hive alive") log.Trace("wake up: make hive alive")
self.wake() self.wake()
case need := <-self.toggle: case need := <-self.toggle:
if alarm == nil && need { if ticker == nil && need {
alarm = af() ticker = self.newTicker()
tick = ticker.Ch()
} }
// if hive saturated, no more peers asked // if hive saturated, no more peers asked
if alarm != nil && !need { if ticker != nil && !need {
alarm = nil ticker.Stop()
ticker = nil
tick = nil
} }
case <-self.quit: case <-self.quit:
return return
@ -272,8 +292,7 @@ func (self *Hive) keepAlive(af func() <-chan time.Time) {
// loadPeers, savePeer implement persistence callback/ // loadPeers, savePeer implement persistence callback/
func (self *Hive) loadPeers() error { func (self *Hive) loadPeers() error {
rw := self.RW data, err := self.store.Load("peers")
data, err := rw.ReadAll("peers")
if err != nil { if err != nil {
return err return err
} }
@ -310,7 +329,7 @@ func (self *Hive) savePeers() error {
if err != nil { if err != nil {
return fmt.Errorf("could not encode peers: %v", err) 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 fmt.Errorf("could not save peers: %v", err)
} }
return nil return nil

View file

@ -16,10 +16,13 @@ type testConnect struct {
ticker chan time.Time ticker chan time.Time
} }
func (self *testConnect) ping() <-chan time.Time { func (self *testConnect) Ch() <-chan time.Time {
return self.ticker return self.ticker
} }
func (self *testConnect) Stop() {
}
func (self *testConnect) connect(na string) error { func (self *testConnect) connect(na string) error {
self.mu.Lock() self.mu.Lock()
defer self.mu.Unlock() defer self.mu.Unlock()
@ -31,38 +34,10 @@ func (self *testConnect) connect(na string) error {
func newHiveTester(t *testing.T, params *HiveParams) (*bzzTester, *Hive) { func newHiveTester(t *testing.T, params *HiveParams) (*bzzTester, *Hive) {
// setup // setup
addr := RandomAddr() // tested peers peer address addr := RandomAddr() // tested peers peer address
// to := NewTestOverlay(addr.Over()) // overlay topology drive to := NewKademlia(addr.OAddr, NewKadParams())
pp := NewHive(params, nil) // hive pp := NewHive(params, to, 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")
// }
return newBzzBaseTester(t, 1, addr, DiscoveryProtocol, pp.Run), pp
} }
func TestRegisterAndConnect(t *testing.T) { func TestRegisterAndConnect(t *testing.T) {
@ -73,7 +48,12 @@ func TestRegisterAndConnect(t *testing.T) {
id := s.Ids[0] id := s.Ids[0]
raddr := NewAddrFromNodeId(id) 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 // start the hive and wait for the connection
tc := &testConnect{ tc := &testConnect{
@ -83,18 +63,17 @@ func TestRegisterAndConnect(t *testing.T) {
}, },
ticker: make(chan time.Time), ticker: make(chan time.Time),
} }
pp.Start(s, tc.ping, nil) pp.newTicker = func() hiveTicker { return tc }
pp.Start(s)
defer pp.Stop() defer pp.Stop()
tc.ticker <- time.Now() tc.ticker <- time.Now()
s.runHandshakes()
// if pp.Overlay.(*testOverlay).posMap[string(raddr.Over())] == nil { // if pp.Overlay.(*testOverlay).posMap[string(raddr.Over())] == nil {
// t.Fatalf("Overlay#On not called on new peer") // t.Fatalf("Overlay#On not called on new peer")
// } // }
// retrieve and broadcast // retrieve and broadcast
ord := order(raddr.Over()) ord := raddr.Over()[0] / 32
o := 0 o := 0
if ord == 0 { if ord == 0 {
o = 1 o = 1

View file

@ -153,11 +153,14 @@ func (k *testKademlia) Off(offs ...string) *testKademlia {
} }
func (k *testKademlia) Register(regs ...string) *testKademlia { func (k *testKademlia) Register(regs ...string) *testKademlia {
var ps []Addr ch := make(chan OverlayAddr)
go func() {
defer close(ch)
for _, s := range regs { for _, s := range regs {
ps = append(ps, Addr(testKadPeerAddr(s))) ch <- testKadPeerAddr(s)
} }
k.Kademlia.Register(ps...) }()
k.Kademlia.Register(ch)
return k return k
} }

View file

@ -17,7 +17,10 @@
package network package network
import ( import (
"context"
"errors"
"fmt" "fmt"
"sync"
"time" "time"
"github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/crypto"
@ -26,15 +29,43 @@ import (
"github.com/ethereum/go-ethereum/p2p/discover" "github.com/ethereum/go-ethereum/p2p/discover"
"github.com/ethereum/go-ethereum/p2p/protocols" "github.com/ethereum/go-ethereum/p2p/protocols"
"github.com/ethereum/go-ethereum/p2p/simulations/adapters" "github.com/ethereum/go-ethereum/p2p/simulations/adapters"
"github.com/ethereum/go-ethereum/rpc"
) )
const ( const (
ProtocolName = "bzz"
Version = 0
NetworkId = 322 // BZZ in l33t NetworkId = 322 // BZZ in l33t
ProtocolMaxMsgSize = 10 * 1024 * 1024 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 // the Addr interface that peerPool needs
type Addr interface { type Addr interface {
OverlayPeer OverlayPeer
@ -53,12 +84,144 @@ type Peer interface {
// Conn interface represents an live peer connection // Conn interface represents an live peer connection
type Conn interface { type Conn interface {
ID() discover.NodeID // the key that uniquely identifies the Node for the peerPool ID() discover.NodeID // the key that uniquely identifies the Node for the peerPool
Handshake(interface{}, time.Duration) (interface{}, error) // can send messages Handshake(context.Context, interface{}) (interface{}, error) // can send messages
Send(interface{}) error // can send messages Send(interface{}) error // can send messages
Drop(error) // disconnect this peer Drop(error) // disconnect this peer
Register(interface{}, func(interface{}) error) uint64 // register message-handler callbacks Run(func(interface{}) error) error // the run function to run a protocol
DisconnectHook(func(error)) // register message-handler callbacks }
Run() 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) // 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 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 // Off returns the overlay peer record for offline persistance
func (self *bzzPeer) Off() OverlayAddr { func (self *bzzPeer) Off() OverlayAddr {
return self.bzzAddr return self.bzzAddr
@ -80,47 +250,6 @@ func (self *bzzPeer) LastActive() time.Time {
return self.lastActive 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 Handshake
@ -132,12 +261,52 @@ type bzzHandshake struct {
Version uint64 Version uint64
NetworkId uint64 NetworkId uint64
Addr *bzzAddr Addr *bzzAddr
// peerAddr is the address received in the peer handshake
peerAddr *bzzAddr
done chan struct{}
err error
} }
func (self *bzzHandshake) String() string { func (self *bzzHandshake) String() string {
return fmt.Sprintf("Handshake: Version: %v, NetworkId: %v, Addr: %v", self.Version, self.NetworkId, self.Addr) 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 // bzzAddr implements the PeerAddr interface
type bzzAddr struct { type bzzAddr struct {
OAddr []byte OAddr []byte
@ -149,6 +318,9 @@ func (self *bzzAddr) Address() []byte {
return self.OAddr return self.OAddr
} }
func (self *bzzAddr) Bytes() []byte {
return self.OAddr
}
func (self *bzzAddr) Over() []byte { func (self *bzzAddr) Over() []byte {
return self.OAddr return self.OAddr
} }
@ -171,47 +343,6 @@ func (self *bzzAddr) String() string {
return fmt.Sprintf("%x <%x>", self.OAddr, self.UAddr) 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 // RandomAddr is a utility method generating an address from a public key
func RandomAddr() *bzzAddr { func RandomAddr() *bzzAddr {
key, err := crypto.GenerateKey() key, err := crypto.GenerateKey()

View file

@ -2,14 +2,43 @@ package network
import ( import (
"fmt" "fmt"
"sync"
"testing" "testing"
"github.com/ethereum/go-ethereum/log" "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/protocols"
"github.com/ethereum/go-ethereum/p2p/simulations/adapters" "github.com/ethereum/go-ethereum/p2p/simulations/adapters"
p2ptest "github.com/ethereum/go-ethereum/p2p/testing" 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 { func bzzHandshakeExchange(lhs, rhs *bzzHandshake, id *adapters.NodeId) []p2ptest.Exchange {
return []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 { func newBzzBaseTester(t *testing.T, n int, addr *bzzAddr, spec *protocols.Spec, run func(*bzzPeer) error) *bzzTester {
if ct == nil {
ct = BzzCodeMap()
}
cs := make(map[string]chan bool) cs := make(map[string]chan bool)
srv := func(p *bzzPeer) error { srv := func(p *bzzPeer) error {
defer close(cs[p.ID().String()]) 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) 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{ return &bzzTester{
addr: addr.Address(), addr: addr,
ProtocolTester: s, ProtocolTester: s,
cs: cs, cs: cs,
} }
@ -63,26 +94,18 @@ func newBzzBaseTester(t *testing.T, n int, addr *bzzAddr, ct *protocols.CodeMap,
type bzzTester struct { type bzzTester struct {
*p2ptest.ProtocolTester *p2ptest.ProtocolTester
addr []byte addr *bzzAddr
cs map[string]chan bool 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 { extraservices := func(p *bzzPeer) error {
pp.Add(p) pp.Add(p)
p.DisconnectHook(func(err error) { defer pp.Remove(p)
pp.Remove(p) return services(p)
})
if services != nil {
err := services(p)
if err != nil {
return err
} }
} return newBzzBaseTester(t, n, addr, spec, extraservices)
return nil
}
return newBzzBaseTester(t, n, addr, ct, extraservices)
} }
// should test handshakes in one exchange? parallelisation // should test handshakes in one exchange? parallelisation
@ -113,7 +136,11 @@ func (s *bzzTester) runHandshakes(ids ...*adapters.NodeId) {
} }
func correctBzzHandshake(addr *bzzAddr) *bzzHandshake { func correctBzzHandshake(addr *bzzAddr) *bzzHandshake {
return &bzzHandshake{0, 322, addr} return &bzzHandshake{
Version: 0,
NetworkId: 322,
Addr: addr,
}
} }
func TestBzzHandshakeNetworkIdMismatch(t *testing.T) { func TestBzzHandshakeNetworkIdMismatch(t *testing.T) {
@ -125,7 +152,7 @@ func TestBzzHandshakeNetworkIdMismatch(t *testing.T) {
id := s.Ids[0] id := s.Ids[0]
s.testHandshake( s.testHandshake(
correctBzzHandshake(addr), 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)")}, &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] id := s.Ids[0]
s.testHandshake( s.testHandshake(
correctBzzHandshake(addr), 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)")}, &p2ptest.Disconnect{Peer: id, Error: fmt.Errorf("version mismatch 1 (!= 0)")},
) )
} }
@ -153,7 +180,7 @@ func TestBzzHandshakeSuccess(t *testing.T) {
id := s.Ids[0] id := s.Ids[0]
s.testHandshake( s.testHandshake(
correctBzzHandshake(addr), 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() defer s.Stop()
id := s.Ids[0] 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) { if pp.Has(id) {
t.Fatalf("peer %v incorrectly added: %v", id, pp) t.Fatalf("peer %v incorrectly added: %v", id, pp)
} }

View file

@ -121,6 +121,7 @@ type pssDigest uint32
// - a message cache to spot messages that previously have been forwarded // - a message cache to spot messages that previously have been forwarded
type Pss struct { type Pss struct {
Overlay // we can get the overlayaddress from this 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]*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 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 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 // 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. // 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 RecipientOAddr pot.Address
LastActive time.Time LastActive time.Time
rw chan p2p.Msg rw chan p2p.Msg
ct *protocols.CodeMap spec *protocols.Spec
topic *PssTopic topic *PssTopic
} }
@ -396,7 +411,7 @@ func (prw PssReadWriter) ReadMsg() (p2p.Msg, error) {
// Implements p2p.MsgWriter // Implements p2p.MsgWriter
func (prw PssReadWriter) WriteMsg(msg p2p.Msg) error { func (prw PssReadWriter) WriteMsg(msg p2p.Msg) error {
log.Trace(fmt.Sprintf("pssrw writemsg: %v", msg)) log.Trace(fmt.Sprintf("pssrw writemsg: %v", msg))
ifc, found := prw.ct.GetInterface(msg.Code) ifc, found := prw.spec.NewMsg(msg.Code)
if !found { if !found {
return fmt.Errorf("Writemsg couldn't find matching interface for code %d", msg.Code) 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 // Convenience object for passing messages in and out of the p2p layer
type PssProtocol struct { type pssProtocol struct {
*Pss *Pss
virtualProtocol *p2p.Protocol virtualProtocol *p2p.Protocol
topic *PssTopic topic *PssTopic
ct *protocols.CodeMap spec *protocols.Spec
} }
// Constructor // Constructor
func NewPssProtocol(pss *Pss, topic *PssTopic, ct *protocols.CodeMap, targetprotocol *p2p.Protocol) *PssProtocol { func NewPssProtocol(pss *Pss, topic *PssTopic, spec *protocols.Spec, targetprotocol *p2p.Protocol) *pssProtocol {
pp := &PssProtocol{ pp := &pssProtocol{
Pss: pss, Pss: pss,
virtualProtocol: targetprotocol, virtualProtocol: targetprotocol,
topic: topic, topic: topic,
ct: ct, spec: spec,
} }
return pp 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 // 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 // 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 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 hashoaddr := pot.NewHashAddressFromBytes(senderAddr).Address
if !self.isActive(hashoaddr, *self.topic) { if !self.isActive(hashoaddr, *self.topic) {
rw := &PssReadWriter{ rw := &PssReadWriter{
Pss: self.Pss, Pss: self.Pss,
RecipientOAddr: hashoaddr, RecipientOAddr: hashoaddr,
rw: make(chan p2p.Msg), rw: make(chan p2p.Msg),
ct: self.ct, spec: self.spec,
topic: self.topic, topic: self.topic,
} }
self.Pss.AddPeer(p, hashoaddr, self.virtualProtocol.Run, *self.topic, rw) self.Pss.AddPeer(p, hashoaddr, self.virtualProtocol.Run, *self.topic, rw)

View file

@ -1,25 +1,9 @@
package network package network
import ( import (
"context"
"encoding/hex"
"fmt"
"math/rand"
"net"
"net/http"
"os" "os"
"strconv"
"testing"
"time"
"github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/log" "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 ( const (
@ -32,6 +16,7 @@ func init() {
log.Root().SetHandler(h) log.Root().SetHandler(h)
} }
/*
// example protocol implementation peer // example protocol implementation peer
// message handlers are methods of this // message handlers are methods of this
// channels allow receipt reporting from p2p.Protocol message handler // channels allow receipt reporting from p2p.Protocol message handler
@ -66,10 +51,6 @@ func (n *pssTestNode) Add(peer *bzzPeer) error {
return err return err
} }
func (n *pssTestNode) hiveKeepAlive() <-chan time.Time {
return time.Tick(time.Millisecond * 300)
}
func (n *pssTestNode) triggerCheck() { func (n *pssTestNode) triggerCheck() {
go func() { n.trigger <- n.id }() 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 { func newPssTestService(t *testing.T, handlefunc func(interface{}) error, testnode *pssTestNode) *pssTestService {
hp := NewHiveParams() hp := NewHiveParams()
//hp.CallInterval = 250 hp.KeepAliveInterval = 300
testnode.Hive = NewHive(hp, testnode.Pss.Overlay) bzz := NewBzz(testnode.OverlayAddr(), testnode.UnderlayAddr(), newTestStore())
testnode.Hive = NewHive(hp, testnode.Pss.Overlay, bzz)
return &pssTestService{ return &pssTestService{
//nid := adapters.NewNodeId(addr.UnderlayAddr()) //nid := adapters.NewNodeId(addr.UnderlayAddr())
msgFunc: handlefunc, msgFunc: handlefunc,
@ -108,7 +90,7 @@ func newPssTestService(t *testing.T, handlefunc func(interface{}) error, testnod
} }
func (self *pssTestService) Start(server p2p.Server) error { 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 { func (self *pssTestService) Stop() error {
@ -117,24 +99,13 @@ func (self *pssTestService) Stop() error {
} }
func (self *pssTestService) Protocols() []p2p.Protocol { func (self *pssTestService) Protocols() []p2p.Protocol {
ct := BzzCodeMap() bzz := NewBzz(self.node.OverlayAddr(), self.node.UnderlayAddr(), newTestStore())
ct.Register(0, &PssMsg{}) return append(self.node.Hive.Protocols(), p2p.Protocol{
for _, m := range DiscoveryMsgs { Name: PssProtocolName,
ct.Register(1, m) Version: PssProtocolVersion,
} Length: PssProtocol.Length(),
Run: bzz.RunProtocol(PssProtocol, self.Run),
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}
} }
func (self *pssTestService) APIs() []rpc.API { func (self *pssTestService) APIs() []rpc.API {
@ -149,6 +120,12 @@ func (self *pssTestService) APIs() []rpc.API {
return nil 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) { func TestPssCache(t *testing.T) {
var err error var err error
to, _ := hex.DecodeString("08090a0b0c0d0e0f1011121314150001020304050607161718191a1b1c1d1e1f") 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) expectnodesids := []*adapters.NodeId{} // the nodes to expect on (needed by checker)
expectnodesresults := make(map[*adapters.NodeId][]int) // which messages expect actually got expectnodesresults := make(map[*adapters.NodeId][]int) // which messages expect actually got
vct := protocols.NewCodeMap(protocolName, protocolVersion, ProtocolMaxMsgSize) vct := protocols.NewCodeMap(map[uint64]interface{}{
vct.Register(0, &pssTestPayload{}) 0: pssTestPayload{},
})
topic, _ := MakeTopic(protocolName, protocolVersion) topic, _ := MakeTopic(protocolName, protocolVersion)
trigger := make(chan *adapters.NodeId) trigger := make(chan *adapters.NodeId)
@ -536,15 +514,15 @@ func TestPssFullLinearEcho(t *testing.T) {
return err return err
} }
/*for i, id := range ids { // for i, id := range ids {
var peerId *adapters.NodeId // var peerId *adapters.NodeId
if i != 0 { // if i != 0 {
peerId = ids[i-1] // peerId = ids[i-1]
if err := net.Connect(id, peerId); err != nil { // if err := net.Connect(id, peerId); err != nil {
return err // return err
} // }
} // }
}*/ // }
return nil return nil
} }
check = func(ctx context.Context, id *adapters.NodeId) (bool, error) { check = func(ctx context.Context, id *adapters.NodeId) (bool, error) {
@ -1108,3 +1086,4 @@ func (ptp *pssTestPeer) SimpleHandlePssPayload(msg interface{}) error {
return nil return nil
} }
*/

View file

@ -9,11 +9,10 @@ import (
"time" "time"
"github.com/ethereum/go-ethereum/log" "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"
"github.com/ethereum/go-ethereum/p2p/simulations" "github.com/ethereum/go-ethereum/p2p/simulations"
"github.com/ethereum/go-ethereum/p2p/simulations/adapters" "github.com/ethereum/go-ethereum/p2p/simulations/adapters"
"github.com/ethereum/go-ethereum/rpc"
"github.com/ethereum/go-ethereum/swarm/network" "github.com/ethereum/go-ethereum/swarm/network"
) )
@ -22,9 +21,7 @@ import (
const serviceName = "discovery" const serviceName = "discovery"
var services = adapters.Services{ var services = adapters.Services{
serviceName: func(id *adapters.NodeId, snapshot []byte) p2pnode.Service { serviceName: newService,
return newNode(id)
},
} }
func init() { func init() {
@ -70,7 +67,7 @@ func testDiscoverySimulation(t *testing.T, adapter adapters.NodeAdapter) {
for i := 0; i < nodeCount; i++ { for i := 0; i < nodeCount; i++ {
node, err := net.NewNode() node, err := net.NewNode()
if err != nil { 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 { if err := net.Start(node.ID()); err != nil {
t.Fatalf("error starting node %s: %s", node.ID().Label(), err) 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 return nil
} }
type node struct { func newService(id *adapters.NodeId, snapshot []byte) node.Service {
*network.Hive addr := network.NewAddrFromNodeId(id)
protocol *p2p.Protocol config := &network.BzzConfig{
} OverlayAddr: addr.Over(),
UnderlayAddr: addr.Under(),
func newNode(id *adapters.NodeId) *node { KadParams: network.NewKadParams(),
addr := network.NewPeerAddrFromNodeId(id) HiveParams: network.NewHiveParams(),
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
} }
node.protocol = network.Bzz(addr.OverlayAddr(), addr.UnderlayAddr(), codeMap, services, nil, nil)
return node config.KadParams.MinProxBinSize = 2
} config.KadParams.MaxBinSize = 3
config.KadParams.MinBinSize = 1
func newKademlia(overlayAddr []byte) *network.Kademlia { config.KadParams.MaxRetries = 1000
params := network.NewKadParams() config.KadParams.RetryExponent = 2
params.MinProxBinSize = 2 config.KadParams.RetryInterval = 1000000
params.MaxBinSize = 3
params.MinBinSize = 1 config.HiveParams.KeepAliveInterval = time.Second
params.MaxRetries = 1000
params.RetryExponent = 2 return network.NewBzz(config)
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)
} }