p2p/discv5: implement UDP TALK request in discv5

This commit is contained in:
Zsolt Felfoldi 2020-01-03 20:38:33 +01:00
parent 01d92531ee
commit 9907044bf5
4 changed files with 204 additions and 15 deletions

View file

@ -22,12 +22,14 @@ import (
"errors"
"fmt"
"net"
"sync"
"time"
"github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/common/mclock"
"github.com/ethereum/go-ethereum/crypto"
"github.com/ethereum/go-ethereum/log"
"github.com/ethereum/go-ethereum/p2p/enode"
"github.com/ethereum/go-ethereum/p2p/netutil"
"github.com/ethereum/go-ethereum/rlp"
"golang.org/x/crypto/sha3"
@ -64,12 +66,17 @@ type Network struct {
refreshResp chan (<-chan struct{}) // ...and get the channel to block on from this one
read chan ingressPacket // ingress packets arrive here
timeout chan timeoutEvent
queryReq chan *findnodeQuery // lookups submit findnode queries on this channel
queryReq chan deferredQuery // lookups submit findnode queries on this channel
tableOpReq chan func()
tableOpResp chan struct{}
topicRegisterReq chan topicRegisterReq
topicSearchReq chan topicSearchReq
talkRequestSubLock sync.RWMutex
talkResponseSubLock sync.Mutex
talkRequestSubs map[string]TalkRequestHandler
talkResponseSubs map[string]TalkResponseHandler
// State of the main loop.
tab *Table
topictab *topicTable
@ -79,6 +86,21 @@ type Network struct {
timeoutTimers map[timeoutEvent]*time.Timer
}
type (
// TalkRequestHandler processes an incoming talkRequest. If ok is true then response is
// sent back, along with the recommended delay feedback. If there is no urgent need to
// communicate (for example, in case of regular polling) then the sender should wait the
// given amount of seconds before sending the next request. If delay recommendation is
// disobeyed too many times by a sender then the handler can stop responding until the
// last recommended delay has elapsed.
TalkRequestHandler func(id enode.ID, addr *net.UDPAddr, request interface{}) (response interface{}, delay uint, ok bool)
// TalkResponseHandler is registered by the sender of each talkRequest. If the request
// is canceled the handler is still called with a nil parameter. If too many responses
// from a peer are considered invalid by their handlers then the peer goes into
// contested state.
TalkResponseHandler func(response interface{}, delay uint) (valid bool)
)
// transport is implemented by the UDP transport.
// it is an interface so we can test without opening lots of UDP
// sockets and without generating a private key.
@ -101,6 +123,14 @@ type findnodeQuery struct {
reply chan<- []*Node
}
type talkQuery struct {
remote *Node
talkID string
payload interface{}
key string
handler TalkResponseHandler
}
type topicRegisterReq struct {
add bool
topic Topic
@ -151,10 +181,12 @@ func newNetwork(conn transport, ourPubkey ecdsa.PublicKey, dbPath string, netres
timeoutTimers: make(map[timeoutEvent]*time.Timer),
tableOpReq: make(chan func()),
tableOpResp: make(chan struct{}),
queryReq: make(chan *findnodeQuery),
queryReq: make(chan deferredQuery),
topicRegisterReq: make(chan topicRegisterReq),
topicSearchReq: make(chan topicSearchReq),
nodes: make(map[NodeID]*Node),
talkRequestSubs: make(map[string]TalkRequestHandler),
talkResponseSubs: make(map[string]TalkResponseHandler),
}
go net.loop()
return net, nil
@ -410,7 +442,6 @@ loop:
// Ingress packet handling.
case pkt := <-net.read:
//fmt.Println("read", pkt.ev)
log.Trace("<-net.read")
n := net.internNode(&pkt)
prestate := n.state
@ -446,7 +477,7 @@ loop:
case q := <-net.queryReq:
log.Trace("<-net.queryReq")
if !q.start(net) {
q.remote.deferQuery(q)
q.deferQuery()
}
// Interacting with the table.
@ -767,14 +798,21 @@ type nodeNetGuts struct {
// State machine fields. Access to these fields
// is restricted to the Network.loop goroutine.
state *nodeState
pingEcho []byte // hash of last ping sent by us
pingTopics []Topic // topic set sent by us in last ping
deferredQueries []*findnodeQuery // queries that can't be sent yet
pendingNeighbours *findnodeQuery // current query, waiting for reply
pingEcho []byte // hash of last ping sent by us
pingTopics []Topic // topic set sent by us in last ping
deferredQueries []deferredQuery // queries that can't be sent yet
pendingNeighbours *findnodeQuery // current query, waiting for reply
queryTimeouts int
talkFailures int
}
func (n *nodeNetGuts) deferQuery(q *findnodeQuery) {
type deferredQuery interface {
start(net *Network) bool
cancel()
deferQuery()
}
func (n *nodeNetGuts) deferQuery(q deferredQuery) {
n.deferredQueries = append(n.deferredQueries, q)
}
@ -810,6 +848,14 @@ func (q *findnodeQuery) start(net *Network) bool {
return false
}
func (q *findnodeQuery) cancel() {
q.reply <- nil
}
func (q *findnodeQuery) deferQuery() {
q.remote.deferQuery(q)
}
// Node Events (the input to the state machine).
type nodeEvent uint
@ -828,6 +874,8 @@ const (
topicRegisterPacket
topicQueryPacket
topicNodesPacket
talkRequestPacket
talkResponsePacket
// Non-packet events.
// Event values in this category are allocated outside
@ -835,6 +883,7 @@ const (
pongTimeout nodeEvent = iota + 256
pingTimeout
neighboursTimeout
talkTimeout
)
// Node State Machine.
@ -868,7 +917,7 @@ func init() {
n.pingEcho = nil
// Abort active queries.
for _, q := range n.deferredQueries {
q.reply <- nil
q.cancel()
}
n.deferredQueries = nil
if n.pendingNeighbours != nil {
@ -1021,7 +1070,7 @@ func (net *Network) handle(n *Node, ev nodeEvent, pkt *ingressPacket) error {
//fmt.Println("handle", n.addr().String(), n.state, ev)
if pkt != nil {
if err := net.checkPacket(n, ev, pkt); err != nil {
//fmt.Println("check err:", err)
//fmt.Println("check err:", err, pkt)
return err
}
// Start the background expiration goroutine after the first
@ -1036,6 +1085,7 @@ func (net *Network) handle(n *Node, ev nodeEvent, pkt *ingressPacket) error {
if n.state == nil {
n.state = unknown //???
}
//fmt.Println("old state:", n.state)
next, err := n.state.handle(net, n, ev, pkt)
net.transition(n, next)
//fmt.Println("new state:", n.state)
@ -1196,7 +1246,51 @@ func (net *Network) handleQueryEvent(n *Node, ev nodeEvent, pkt *ingressPacket)
}
}
return n.state, nil
case talkRequestPacket:
p := pkt.data.(*talkRequest)
net.talkRequestSubLock.RLock()
subFn := net.talkRequestSubs[string(p.TalkID)]
net.talkRequestSubLock.RUnlock()
if subFn != nil {
resp, delay, ok := subFn(enode.ID(n.sha), n.addr(), p.Payload)
if ok {
net.conn.send(n, talkResponsePacket, talkResponse{ReplyTok: pkt.hash, Delay: delay, Payload: resp})
} else {
n.talkFailures++
}
} else {
n.talkFailures++
}
if n.talkFailures > maxTalkFailures && n.state == known {
return contested, errors.New("too many talk failures")
}
return n.state, nil
case talkResponsePacket:
p := pkt.data.(*talkResponse)
net.talkResponseSubLock.Lock()
key := string(n.sha[:]) + string(p.ReplyTok)
subFn := net.talkResponseSubs[key]
if subFn != nil {
delete(net.talkResponseSubs, key)
}
net.talkResponseSubLock.Unlock()
if subFn == nil || !subFn(p.Payload, p.Delay) {
n.talkFailures++
if n.talkFailures > maxTalkFailures && n.state == known {
return contested, errors.New("too many talk failures")
}
}
return n.state, nil
case talkTimeout:
if n.pendingNeighbours != nil {
n.pendingNeighbours.reply <- nil
n.pendingNeighbours = nil
}
n.queryTimeouts++
if n.queryTimeouts > maxFindnodeFailures && n.state == known {
return contested, errors.New("too many timeouts")
}
return n.state, nil
default:
return n.state, errInvalidEvent
}
@ -1260,3 +1354,82 @@ func (net *Network) handleNeighboursPacket(n *Node, pkt *ingressPacket) error {
n.startNextQuery(net)
return nil
}
// RegisterTalkHandler assigns a handler callback to the given talk ID
func (net *Network) RegisterTalkHandler(talkID string, handler TalkRequestHandler) {
net.talkRequestSubLock.Lock()
net.talkRequestSubs[talkID] = handler
net.talkRequestSubLock.Unlock()
}
// RemoveTalkHandler removes the handler assigned to the given talk ID
func (net *Network) RemoveTalkHandler(talkID string) {
net.talkRequestSubLock.Lock()
delete(net.talkRequestSubs, talkID)
net.talkRequestSubLock.Unlock()
}
func (q *talkQuery) start(net *Network) bool {
if q.remote == net.tab.self {
return false
}
if q.remote.state == known {
net.talkResponseSubLock.Lock()
hash := net.conn.send(q.remote, talkRequestPacket, talkRequest{TalkID: []byte(q.talkID), Payload: q.payload})
q.key = string(q.remote.sha[:]) + string(hash[:])
net.talkResponseSubs[q.key] = q.handler
net.talkResponseSubLock.Unlock()
return true
}
// If the node is not known yet, it won't accept queries.
// Initiate the transition to known.
// The request will be sent later when the node reaches known state.
if q.remote.state == unknown {
net.transition(q.remote, verifyinit)
}
return false
}
func (q *talkQuery) cancel() {
q.handler(nil, 0)
}
func (q *talkQuery) deferQuery() {
q.remote.deferQuery(q)
}
// SendTalkRequest sends a talRequest and registers a response handler. It returns
// a cancel function that removes the response handler and calls it with a nil parameter
// if the response has not arrived yet.
func (net *Network) SendTalkRequest(to *enode.Node, talkID string, payload interface{}, handler TalkResponseHandler) (cancel func() bool) {
var nodeID NodeID
copy(nodeID[:], crypto.FromECDSAPub(to.Pubkey())[1:])
node := net.nodes[nodeID]
if node == nil {
node = NewNode(nodeID, to.IP(), uint16(to.UDP()), uint16(to.TCP()))
node.state = unknown
net.nodes[nodeID] = node
}
q := &talkQuery{remote: node, talkID: talkID, payload: payload, handler: handler}
if node.state == nil || !node.state.canQuery {
net.ping(node, node.addr())
}
select {
case net.queryReq <- q:
case <-net.closed:
return nil
}
return func() bool {
net.talkResponseSubLock.Lock()
cancel := q.key != "" && net.talkResponseSubs[q.key] != nil
if cancel {
delete(net.talkResponseSubs, q.key)
}
net.talkResponseSubLock.Unlock()
if cancel {
handler(nil, 0)
}
return cancel
}
}

View file

@ -35,6 +35,7 @@ const (
nBuckets = hashBits + 1 // Number of buckets
maxFindnodeFailures = 5
maxTalkFailures = 5
)
type Table struct {

View file

@ -119,6 +119,17 @@ type (
Nodes []rpcNode
}
talkRequest struct {
TalkID []byte
Payload interface{}
}
talkResponse struct {
ReplyTok []byte
Delay uint
Payload interface{}
}
rpcNode struct {
IP net.IP // len 4 for IPv4 or 16 for IPv6
UDP uint16 // for discovery protocol
@ -420,6 +431,10 @@ func decodePacket(buffer []byte, pkt *ingressPacket) error {
pkt.data = new(topicQuery)
case topicNodesPacket:
pkt.data = new(topicNodes)
case talkRequestPacket:
pkt.data = new(talkRequest)
case talkResponsePacket:
pkt.data = new(talkResponse)
default:
return fmt.Errorf("unknown packet type: %d", sigdata[0])
}

View file

@ -217,7 +217,7 @@ func (n ID) MarshalText() ([]byte, error) {
// UnmarshalText implements the encoding.TextUnmarshaler interface.
func (n *ID) UnmarshalText(text []byte) error {
id, err := parseID(string(text))
id, err := ParseID(string(text))
if err != nil {
return err
}
@ -229,14 +229,14 @@ func (n *ID) UnmarshalText(text []byte) error {
// The string may be prefixed with 0x.
// It panics if the string is not a valid ID.
func HexID(in string) ID {
id, err := parseID(in)
id, err := ParseID(in)
if err != nil {
panic(err)
}
return id
}
func parseID(in string) (ID, error) {
func ParseID(in string) (ID, error) {
var id ID
b, err := hex.DecodeString(strings.TrimPrefix(in, "0x"))
if err != nil {