p2p/discover: add packet kind method to avoid use of constants

This allows removing the packet type parameter from all methods dealing
with packet structs because the type byte can be retrieved with the
method instead.
This commit is contained in:
Felix Lange 2019-05-03 10:29:33 +02:00
parent b6c0234e0b
commit d2e5d97671
3 changed files with 65 additions and 50 deletions

View file

@ -126,14 +126,15 @@ type (
} }
) )
// packet is implemented by all v4 protocol messages. // packetV4 is implemented by all v4 protocol messages.
type packetV4 interface { type packetV4 interface {
// preverify checks whether the packet is valid and should be handled at all. // preverify checks whether the packet is valid and should be handled at all.
preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error
// handle handles the packet. // handle handles the packet.
handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []byte) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []byte)
// name returns the name of the packet for logging purposes. // packet name and type for logging purposes.
name() string name() string
kind() byte
} }
func makeEndpoint(addr *net.UDPAddr, tcpPort uint16) rpcEndpoint { func makeEndpoint(addr *net.UDPAddr, tcpPort uint16) rpcEndpoint {
@ -191,7 +192,7 @@ type UDPv4 struct {
closing chan struct{} closing chan struct{}
} }
// pending represents a pending reply. // replyMatcher represents a pending reply.
// //
// Some implementations of the protocol wish to send more than one // Some implementations of the protocol wish to send more than one
// reply packet to findnode. In general, any neighbors packet cannot // reply packet to findnode. In general, any neighbors packet cannot
@ -222,12 +223,11 @@ type replyMatcher struct {
type replyMatchFunc func(interface{}) (matched bool, requestDone bool) type replyMatchFunc func(interface{}) (matched bool, requestDone bool)
// reply is a reply packet from a certain node.
type reply struct { type reply struct {
from enode.ID from enode.ID
ip net.IP ip net.IP
ptype byte
data packetV4 data packetV4
// loop indicates whether there was // loop indicates whether there was
// a matching request by sending on this channel. // a matching request by sending on this channel.
matched chan<- bool matched chan<- bool
@ -426,7 +426,7 @@ func (t *UDPv4) sendPing(toid enode.ID, toaddr *net.UDPAddr, callback func()) <-
To: makeEndpoint(toaddr, 0), // TODO: maybe use known TCP port from DB To: makeEndpoint(toaddr, 0), // TODO: maybe use known TCP port from DB
Expiration: uint64(time.Now().Add(expiration).Unix()), Expiration: uint64(time.Now().Add(expiration).Unix()),
} }
packet, hash, err := t.encode(t.priv, p_pingV4, req) packet, hash, err := t.encode(t.priv, req)
if err != nil { if err != nil {
errc := make(chan error, 1) errc := make(chan error, 1)
errc <- err errc <- err
@ -475,7 +475,7 @@ func (t *UDPv4) findnode(toid enode.ID, toaddr *net.UDPAddr, target encPubkey) (
} }
return true, nreceived >= bucketSize return true, nreceived >= bucketSize
}) })
t.send(toaddr, toid, p_findnodeV4, &findnodeV4{ t.send(toaddr, toid, &findnodeV4{
Target: target, Target: target,
Expiration: uint64(time.Now().Add(expiration).Unix()), Expiration: uint64(time.Now().Add(expiration).Unix()),
}) })
@ -498,10 +498,10 @@ func (t *UDPv4) pending(id enode.ID, ip net.IP, ptype byte, callback replyMatchF
// handleReply dispatches a reply packet, invoking reply matchers. It returns // handleReply dispatches a reply packet, invoking reply matchers. It returns
// whether any matcher considered the packet acceptable. // whether any matcher considered the packet acceptable.
func (t *UDPv4) handleReply(from enode.ID, fromIP net.IP, ptype byte, req packetV4) bool { func (t *UDPv4) handleReply(from enode.ID, fromIP net.IP, req packetV4) bool {
matched := make(chan bool, 1) matched := make(chan bool, 1)
select { select {
case t.gotreply <- reply{from, fromIP, ptype, req, matched}: case t.gotreply <- reply{from, fromIP, req, matched}:
// loop will handle it // loop will handle it
return <-matched return <-matched
case <-t.closing: case <-t.closing:
@ -564,7 +564,7 @@ func (t *UDPv4) loop() {
var matched bool // whether any replyMatcher considered the reply acceptable. var matched bool // whether any replyMatcher considered the reply acceptable.
for el := plist.Front(); el != nil; el = el.Next() { for el := plist.Front(); el != nil; el = el.Next() {
p := el.Value.(*replyMatcher) p := el.Value.(*replyMatcher)
if p.from == r.from && p.ptype == r.ptype && p.ip.Equal(r.ip) { if p.from == r.from && p.ptype == r.data.kind() && p.ip.Equal(r.ip) {
ok, requestDone := p.callback(r.data) ok, requestDone := p.callback(r.data)
matched = matched || ok matched = matched || ok
// Remove the matcher if callback indicates that all replies have been received. // Remove the matcher if callback indicates that all replies have been received.
@ -634,8 +634,8 @@ func init() {
} }
} }
func (t *UDPv4) send(toaddr *net.UDPAddr, toid enode.ID, ptype byte, req packetV4) ([]byte, error) { func (t *UDPv4) send(toaddr *net.UDPAddr, toid enode.ID, req packetV4) ([]byte, error) {
packet, hash, err := t.encode(t.priv, ptype, req) packet, hash, err := t.encode(t.priv, req)
if err != nil { if err != nil {
return hash, err return hash, err
} }
@ -648,18 +648,19 @@ func (t *UDPv4) write(toaddr *net.UDPAddr, toid enode.ID, what string, packet []
return err return err
} }
func (t *UDPv4) encode(priv *ecdsa.PrivateKey, ptype byte, req interface{}) (packet, hash []byte, err error) { func (t *UDPv4) encode(priv *ecdsa.PrivateKey, req packetV4) (packet, hash []byte, err error) {
name := req.name()
b := new(bytes.Buffer) b := new(bytes.Buffer)
b.Write(headSpace) b.Write(headSpace)
b.WriteByte(ptype) b.WriteByte(req.kind())
if err := rlp.Encode(b, req); err != nil { if err := rlp.Encode(b, req); err != nil {
t.log.Error("Can't encode discv4 packet", "err", err) t.log.Error(fmt.Sprintf("Can't encode %s packet", name), "err", err)
return nil, nil, err return nil, nil, err
} }
packet = b.Bytes() packet = b.Bytes()
sig, err := crypto.Sign(crypto.Keccak256(packet[headSize:]), priv) sig, err := crypto.Sign(crypto.Keccak256(packet[headSize:]), priv)
if err != nil { if err != nil {
t.log.Error("Can't sign discv4 packet", "err", err) t.log.Error(fmt.Sprintf("Can't sign %s packet", name), "err", err)
return nil, nil, err return nil, nil, err
} }
copy(packet[macSize:], sig) copy(packet[macSize:], sig)
@ -752,6 +753,9 @@ func decodeV4(buf []byte) (packetV4, encPubkey, []byte, error) {
// Packet Handlers // Packet Handlers
func (req *pingV4) name() string { return "PING/v4" }
func (req *pingV4) kind() byte { return p_pingV4 }
func (req *pingV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error { func (req *pingV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error {
if expired(req.Expiration) { if expired(req.Expiration) {
return errExpired return errExpired
@ -766,7 +770,7 @@ func (req *pingV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromK
func (req *pingV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []byte) { func (req *pingV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []byte) {
// Reply. // Reply.
t.send(from, fromID, p_pongV4, &pongV4{ t.send(from, fromID, &pongV4{
To: makeEndpoint(from, req.From.TCP), To: makeEndpoint(from, req.From.TCP),
ReplyTok: mac, ReplyTok: mac,
Expiration: uint64(time.Now().Add(expiration).Unix()), Expiration: uint64(time.Now().Add(expiration).Unix()),
@ -787,13 +791,14 @@ func (req *pingV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []by
t.localNode.UDPEndpointStatement(from, &net.UDPAddr{IP: req.To.IP, Port: int(req.To.UDP)}) t.localNode.UDPEndpointStatement(from, &net.UDPAddr{IP: req.To.IP, Port: int(req.To.UDP)})
} }
func (req *pingV4) name() string { return "PING/v4" } func (req *pongV4) name() string { return "PONG/v4" }
func (req *pongV4) kind() byte { return p_pongV4 }
func (req *pongV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error { func (req *pongV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error {
if expired(req.Expiration) { if expired(req.Expiration) {
return errExpired return errExpired
} }
if !t.handleReply(fromID, from.IP, p_pongV4, req) { if !t.handleReply(fromID, from.IP, req) {
return errUnsolicitedReply return errUnsolicitedReply
} }
return nil return nil
@ -804,7 +809,8 @@ func (req *pongV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []by
t.db.UpdateLastPongReceived(fromID, from.IP, time.Now()) t.db.UpdateLastPongReceived(fromID, from.IP, time.Now())
} }
func (req *pongV4) name() string { return "PONG/v4" } func (req *findnodeV4) name() string { return "FINDNODE/v4" }
func (req *findnodeV4) kind() byte { return p_findnodeV4 }
func (req *findnodeV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error { func (req *findnodeV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error {
if expired(req.Expiration) { if expired(req.Expiration) {
@ -838,23 +844,24 @@ func (req *findnodeV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac
p.Nodes = append(p.Nodes, nodeToRPC(n)) p.Nodes = append(p.Nodes, nodeToRPC(n))
} }
if len(p.Nodes) == maxNeighbors { if len(p.Nodes) == maxNeighbors {
t.send(from, fromID, p_neighborsV4, &p) t.send(from, fromID, &p)
p.Nodes = p.Nodes[:0] p.Nodes = p.Nodes[:0]
sent = true sent = true
} }
} }
if len(p.Nodes) > 0 || !sent { if len(p.Nodes) > 0 || !sent {
t.send(from, fromID, p_neighborsV4, &p) t.send(from, fromID, &p)
} }
} }
func (req *findnodeV4) name() string { return "FINDNODE/v4" } func (req *neighborsV4) name() string { return "NEIGHBORS/v4" }
func (req *neighborsV4) kind() byte { return p_neighborsV4 }
func (req *neighborsV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error { func (req *neighborsV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID, fromKey encPubkey) error {
if expired(req.Expiration) { if expired(req.Expiration) {
return errExpired return errExpired
} }
if !t.handleReply(fromID, from.IP, p_neighborsV4, req) { if !t.handleReply(fromID, from.IP, req) {
return errUnsolicitedReply return errUnsolicitedReply
} }
return nil return nil
@ -863,8 +870,6 @@ func (req *neighborsV4) preverify(t *UDPv4, from *net.UDPAddr, fromID enode.ID,
func (req *neighborsV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []byte) { func (req *neighborsV4) handle(t *UDPv4, from *net.UDPAddr, fromID enode.ID, mac []byte) {
} }
func (req *neighborsV4) name() string { return "NEIGHBORS/v4" }
func expired(ts uint64) bool { func expired(ts uint64) bool {
return time.Unix(int64(ts), 0).Before(time.Now()) return time.Unix(int64(ts), 0).Before(time.Now())
} }

View file

@ -54,11 +54,11 @@ func TestUDPv4_Lookup(t *testing.T) {
n, key := lookupTestnet.nodeByAddr(to) n, key := lookupTestnet.nodeByAddr(to)
switch p.(type) { switch p.(type) {
case *pingV4: case *pingV4:
test.packetInFrom(nil, key, to, p_pongV4, &pongV4{Expiration: futureExp, ReplyTok: hash}) test.packetInFrom(nil, key, to, &pongV4{Expiration: futureExp, ReplyTok: hash})
case *findnodeV4: case *findnodeV4:
dist := enode.LogDist(n.ID(), lookupTestnet.target.id()) dist := enode.LogDist(n.ID(), lookupTestnet.target.id())
nodes := lookupTestnet.nodesAtDistance(dist - 1) nodes := lookupTestnet.nodesAtDistance(dist - 1)
test.packetInFrom(nil, key, to, p_neighborsV4, &neighborsV4{Expiration: futureExp, Nodes: nodes}) test.packetInFrom(nil, key, to, &neighborsV4{Expiration: futureExp, Nodes: nodes})
} }
}) })
} }

View file

@ -91,19 +91,19 @@ func (test *udpTest) close() {
} }
// handles a packet as if it had been sent to the transport. // handles a packet as if it had been sent to the transport.
func (test *udpTest) packetIn(wantError error, ptype byte, data packetV4) { func (test *udpTest) packetIn(wantError error, data packetV4) {
test.t.Helper() test.t.Helper()
test.packetInFrom(wantError, test.remotekey, test.remoteaddr, ptype, data) test.packetInFrom(wantError, test.remotekey, test.remoteaddr, data)
} }
// handles a packet as if it had been sent to the transport by the key/endpoint. // handles a packet as if it had been sent to the transport by the key/endpoint.
func (test *udpTest) packetInFrom(wantError error, key *ecdsa.PrivateKey, addr *net.UDPAddr, ptype byte, data packetV4) { func (test *udpTest) packetInFrom(wantError error, key *ecdsa.PrivateKey, addr *net.UDPAddr, data packetV4) {
test.t.Helper() test.t.Helper()
enc, _, err := test.udp.encode(key, ptype, data) enc, _, err := test.udp.encode(key, data)
if err != nil { if err != nil {
test.t.Errorf("packet (%d) encode error: %v", ptype, err) test.t.Errorf("%s encode error: %v", data.name(), err)
} }
test.sent = append(test.sent, enc) test.sent = append(test.sent, enc)
if err = test.udp.handlePacket(addr, enc); err != wantError { if err = test.udp.handlePacket(addr, enc); err != wantError {
@ -139,10 +139,10 @@ func TestUDPv4_packetErrors(t *testing.T) {
test := newUDPTest(t) test := newUDPTest(t)
defer test.close() defer test.close()
test.packetIn(errExpired, p_pingV4, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4}) test.packetIn(errExpired, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4})
test.packetIn(errUnsolicitedReply, p_pongV4, &pongV4{ReplyTok: []byte{}, Expiration: futureExp}) test.packetIn(errUnsolicitedReply, &pongV4{ReplyTok: []byte{}, Expiration: futureExp})
test.packetIn(errUnknownNode, p_findnodeV4, &findnodeV4{Expiration: futureExp}) test.packetIn(errUnknownNode, &findnodeV4{Expiration: futureExp})
test.packetIn(errUnsolicitedReply, p_neighborsV4, &neighborsV4{Expiration: futureExp}) test.packetIn(errUnsolicitedReply, &neighborsV4{Expiration: futureExp})
} }
func TestUDPv4_pingTimeout(t *testing.T) { func TestUDPv4_pingTimeout(t *testing.T) {
@ -158,6 +158,16 @@ func TestUDPv4_pingTimeout(t *testing.T) {
} }
} }
type testPacket byte
func (req testPacket) kind() byte { return byte(req) }
func (req testPacket) name() string { return "" }
func (req testPacket) preverify(*UDPv4, *net.UDPAddr, enode.ID, encPubkey) error {
return nil
}
func (req testPacket) handle(*UDPv4, *net.UDPAddr, enode.ID, []byte) {
}
func TestUDPv4_responseTimeouts(t *testing.T) { func TestUDPv4_responseTimeouts(t *testing.T) {
t.Parallel() t.Parallel()
test := newUDPTest(t) test := newUDPTest(t)
@ -192,7 +202,7 @@ func TestUDPv4_responseTimeouts(t *testing.T) {
p.errc = nilErr p.errc = nilErr
test.udp.addReplyMatcher <- p test.udp.addReplyMatcher <- p
time.AfterFunc(randomDuration(60*time.Millisecond), func() { time.AfterFunc(randomDuration(60*time.Millisecond), func() {
if !test.udp.handleReply(p.from, p.ip, p.ptype, nil) { if !test.udp.handleReply(p.from, p.ip, testPacket(p.ptype)) {
t.Logf("not matched: %v", p) t.Logf("not matched: %v", p)
} }
}) })
@ -277,7 +287,7 @@ func TestUDPv4_findnode(t *testing.T) {
// check that closest neighbors are returned. // check that closest neighbors are returned.
expected := test.table.closest(testTarget.id(), bucketSize, true) expected := test.table.closest(testTarget.id(), bucketSize, true)
test.packetIn(nil, p_findnodeV4, &findnodeV4{Target: testTarget, Expiration: futureExp}) test.packetIn(nil, &findnodeV4{Target: testTarget, Expiration: futureExp})
waitNeighbors := func(want []*node) { waitNeighbors := func(want []*node) {
test.waitPacketOut(func(p *neighborsV4, to *net.UDPAddr, hash []byte) { test.waitPacketOut(func(p *neighborsV4, to *net.UDPAddr, hash []byte) {
if len(p.Nodes) != len(want) { if len(p.Nodes) != len(want) {
@ -340,8 +350,8 @@ func TestUDPv4_findnodeMultiReply(t *testing.T) {
for i := range list { for i := range list {
rpclist[i] = nodeToRPC(list[i]) rpclist[i] = nodeToRPC(list[i])
} }
test.packetIn(nil, p_neighborsV4, &neighborsV4{Expiration: futureExp, Nodes: rpclist[:2]}) test.packetIn(nil, &neighborsV4{Expiration: futureExp, Nodes: rpclist[:2]})
test.packetIn(nil, p_neighborsV4, &neighborsV4{Expiration: futureExp, Nodes: rpclist[2:]}) test.packetIn(nil, &neighborsV4{Expiration: futureExp, Nodes: rpclist[2:]})
// check that the sent neighbors are all returned by findnode // check that the sent neighbors are all returned by findnode
select { select {
@ -364,22 +374,22 @@ func TestUDPv4_pingMatch(t *testing.T) {
randToken := make([]byte, 32) randToken := make([]byte, 32)
crand.Read(randToken) crand.Read(randToken)
test.packetIn(nil, p_pingV4, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4, Expiration: futureExp}) test.packetIn(nil, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4, Expiration: futureExp})
test.waitPacketOut(func(*pongV4, *net.UDPAddr, []byte) {}) test.waitPacketOut(func(*pongV4, *net.UDPAddr, []byte) {})
test.waitPacketOut(func(*pingV4, *net.UDPAddr, []byte) {}) test.waitPacketOut(func(*pingV4, *net.UDPAddr, []byte) {})
test.packetIn(errUnsolicitedReply, p_pongV4, &pongV4{ReplyTok: randToken, To: testLocalAnnounced, Expiration: futureExp}) test.packetIn(errUnsolicitedReply, &pongV4{ReplyTok: randToken, To: testLocalAnnounced, Expiration: futureExp})
} }
func TestUDPv4_pingMatchIP(t *testing.T) { func TestUDPv4_pingMatchIP(t *testing.T) {
test := newUDPTest(t) test := newUDPTest(t)
defer test.close() defer test.close()
test.packetIn(nil, p_pingV4, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4, Expiration: futureExp}) test.packetIn(nil, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4, Expiration: futureExp})
test.waitPacketOut(func(*pongV4, *net.UDPAddr, []byte) {}) test.waitPacketOut(func(*pongV4, *net.UDPAddr, []byte) {})
test.waitPacketOut(func(p *pingV4, to *net.UDPAddr, hash []byte) { test.waitPacketOut(func(p *pingV4, to *net.UDPAddr, hash []byte) {
wrongAddr := &net.UDPAddr{IP: net.IP{33, 44, 1, 2}, Port: 30000} wrongAddr := &net.UDPAddr{IP: net.IP{33, 44, 1, 2}, Port: 30000}
test.packetInFrom(errUnsolicitedReply, test.remotekey, wrongAddr, p_pongV4, &pongV4{ test.packetInFrom(errUnsolicitedReply, test.remotekey, wrongAddr, &pongV4{
ReplyTok: hash, ReplyTok: hash,
To: testLocalAnnounced, To: testLocalAnnounced,
Expiration: futureExp, Expiration: futureExp,
@ -394,7 +404,7 @@ func TestUDPv4_successfulPing(t *testing.T) {
defer test.close() defer test.close()
// The remote side sends a ping packet to initiate the exchange. // The remote side sends a ping packet to initiate the exchange.
go test.packetIn(nil, p_pingV4, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4, Expiration: futureExp}) go test.packetIn(nil, &pingV4{From: testRemote, To: testLocalAnnounced, Version: 4, Expiration: futureExp})
// the ping is replied to. // the ping is replied to.
test.waitPacketOut(func(p *pongV4, to *net.UDPAddr, hash []byte) { test.waitPacketOut(func(p *pongV4, to *net.UDPAddr, hash []byte) {
@ -427,7 +437,7 @@ func TestUDPv4_successfulPing(t *testing.T) {
if !reflect.DeepEqual(p.To, wantTo) { if !reflect.DeepEqual(p.To, wantTo) {
t.Errorf("got ping.To %v, want %v", p.To, wantTo) t.Errorf("got ping.To %v, want %v", p.To, wantTo)
} }
test.packetIn(nil, p_pongV4, &pongV4{ReplyTok: hash, Expiration: futureExp}) test.packetIn(nil, &pongV4{ReplyTok: hash, Expiration: futureExp})
}) })
// the node should be added to the table shortly after getting the // the node should be added to the table shortly after getting the