p2p/enr: fix RLP encoding, python interop

This commit is contained in:
Felix Lange 2017-12-05 13:42:43 +01:00
parent 12ac6cd58b
commit 94f19f20c6
2 changed files with 112 additions and 116 deletions

View file

@ -21,6 +21,7 @@ import (
"bytes" "bytes"
"crypto/ecdsa" "crypto/ecdsa"
"errors" "errors"
"fmt"
"io" "io"
"math/big" "math/big"
"sort" "sort"
@ -28,12 +29,17 @@ import (
"github.com/btcsuite/btcd/btcec" "github.com/btcsuite/btcd/btcec"
"github.com/ethereum/go-ethereum/common/math" "github.com/ethereum/go-ethereum/common/math"
"github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/crypto"
"github.com/ethereum/go-ethereum/crypto/sha3"
"github.com/ethereum/go-ethereum/rlp" "github.com/ethereum/go-ethereum/rlp"
) )
var ( var (
errNoID = errors.New("unknown or unspecified identity scheme") errNoID = errors.New("unknown or unspecified identity scheme")
errInvalidSigsize = errors.New("invalid signature size") errInvalidSigsize = errors.New("invalid signature size")
errInvalidSig = errors.New("invalid signature")
errNotSorted = errors.New("record key/value pairs are not sorted by key")
errDuplicateKey = errors.New("record contains duplicate key")
errIncompletePair = errors.New("record contains incomplete k/v pair")
) )
// Key is implemented by known node record key types. // Key is implemented by known node record key types.
@ -45,17 +51,17 @@ type Key interface {
ENRKey() string ENRKey() string
} }
// pair is a key/value pair in a record.
type pair struct { type pair struct {
k string k string
v []byte v rlp.RawValue
} }
type Record struct { type Record struct {
seq uint32 // sequence number seq uint32 // sequence number
signature []byte // record's signature signature []byte // record's signature
raw []byte // RLP encoded record raw []byte // RLP encoded record
pairs []pair // list of all key/value pairs, sorted prior to RLP encoding pairs []pair // sorted list of all key/value pairs
signed bool // keeps track if record was modified after it was signed and encoded
} }
func (r Record) Seq() uint32 { func (r Record) Seq() uint32 {
@ -63,7 +69,7 @@ func (r Record) Seq() uint32 {
} }
func (r *Record) SetSeq(s uint32) { func (r *Record) SetSeq(s uint32) {
r.signed = false r.signature = nil
r.seq = s r.seq = s
} }
@ -78,7 +84,7 @@ func (r *Record) Load(k Key) (bool, error) {
} }
func (r *Record) Set(k Key) error { func (r *Record) Set(k Key) error {
r.signed = false r.signature = nil
blob, err := rlp.EncodeToBytes(k) blob, err := rlp.EncodeToBytes(k)
if err != nil { if err != nil {
return err return err
@ -110,101 +116,70 @@ func (r *Record) Set(k Key) error {
} }
func (r Record) EncodeRLP(w io.Writer) error { func (r Record) EncodeRLP(w io.Writer) error {
if !r.signed { if r.signature == nil {
return errors.New("record is not signed") return errors.New("record is not signed")
} }
_, err := w.Write(r.raw) _, err := w.Write(r.raw)
return err return err
} }
func (r *Record) DecodeRLP(s *rlp.Stream) error { func (r *Record) DecodeRLP(s *rlp.Stream) error {
var err error raw, err := s.Raw()
r.signature, err = s.Bytes()
if err != nil { if err != nil {
return err return err
} }
_, err = s.List() // Decode the RLP container.
if err != nil { dec := Record{raw: raw}
s = rlp.NewStream(bytes.NewReader(raw), 0)
if _, err := s.List(); err != nil {
return err
}
if err = s.Decode(&dec.signature); err != nil {
return err
}
if err = s.Decode(&dec.seq); err != nil {
return err
}
// The rest of the record contains sorted k/v pairs.
var prevkey string
for i := 0; ; i++ {
var kv pair
if err := s.Decode(&kv.k); err != nil {
if err == rlp.EOL {
break
}
return err
}
if err := s.Decode(&kv.v); err != nil {
if err == rlp.EOL {
return errIncompletePair
}
return err
}
if i > 0 {
if kv.k == prevkey {
return errDuplicateKey
}
if kv.k < prevkey {
return errNotSorted
}
}
dec.pairs = append(dec.pairs, kv)
prevkey = kv.k
}
if err := s.ListEnd(); err != nil {
return err return err
} }
if err := s.Decode(&r.seq); err != nil { // Verify signature.
if err = dec.verifySignature(); err != nil {
return err return err
} }
*r = dec
// read key/value pairs until we reach rlp.EOL
for _, _, err = s.Kind(); err == nil; _, _, err = s.Kind() {
key, err2 := s.Bytes()
if err2 != nil {
return err2
}
value, err2 := s.Bytes()
if err2 != nil {
return err2
}
r.pairs = append(r.pairs, pair{k: string(key), v: value})
}
if err != rlp.EOL {
return err
}
sigcontent, err := r.serialisedContent()
if err != nil {
return err
}
// update r.raw
blob, err := rlp.EncodeToBytes(r.signature)
if err != nil {
return err
}
r.raw = append(blob, sigcontent...)
err = r.verifySignature(sigcontent)
if err != nil {
return err
}
// mark record ready for encoding
r.signed = true
return nil return nil
} }
func (r Record) Equal(o Record) (bool, error) {
rr, err := r.serialisedContent()
if err != nil {
return false, err
}
oo, err := o.serialisedContent()
if err != nil {
return false, err
}
if bytes.Compare(rr, oo) != 0 {
return false, nil
}
if err := r.verifySignature(rr); err != nil {
return false, err
}
if err := o.verifySignature(oo); err != nil {
return false, err
}
return true, nil
}
func (r *Record) NodeAddr() ([]byte, error) { func (r *Record) NodeAddr() ([]byte, error) {
var secp256k1 Secp256k1 var secp256k1 Secp256k1
@ -221,73 +196,69 @@ func (r *Record) NodeAddr() ([]byte, error) {
} }
func (r *Record) Sign(privkey *ecdsa.PrivateKey) error { func (r *Record) Sign(privkey *ecdsa.PrivateKey) error {
r.seq = r.seq + 1
r.Set(ID(ID_SECP256k1_KECCAK))
pk := (*btcec.PublicKey)(&privkey.PublicKey) pk := (*btcec.PublicKey)(&privkey.PublicKey)
secp256k1 := Secp256k1(*pk) r.seq = r.seq + 1
r.Set(secp256k1) r.Set(ID(ID_SECP256k1_KECCAK))
r.Set(Secp256k1(*pk))
return r.signAndEncode(privkey) return r.signAndEncode(privkey)
} }
func (r *Record) serialisedContent() ([]byte, error) { func (r *Record) appendPairs(list []interface{}) []interface{} {
list := []interface{}{r.seq} list = append(list, r.seq)
for _, p := range r.pairs { for _, p := range r.pairs {
list = append(list, p.k, p.v) list = append(list, p.k, p.v)
} }
return list
return rlp.EncodeToBytes(list)
} }
func (r *Record) signAndEncode(privkey *ecdsa.PrivateKey) error { func (r *Record) signAndEncode(privkey *ecdsa.PrivateKey) error {
sigcontent, err := r.serialisedContent() // Put record elements into a flat list. Leave room for the signature.
list := make([]interface{}, 1, len(r.pairs)*2+2)
list = r.appendPairs(list)
// Sign the tail of the list.
h := sha3.NewKeccak256()
rlp.Encode(h, list[1:])
sig, err := (*btcec.PrivateKey)(privkey).Sign(h.Sum(nil))
if err != nil { if err != nil {
return err return err
} }
sig, err := (*btcec.PrivateKey)(privkey).Sign(crypto.Keccak256(sigcontent)) // Put signature in front.
if err != nil {
return err
}
r.signature = encodeCompactSignature(sig) r.signature = encodeCompactSignature(sig)
list[0] = r.signature
blob, err := rlp.EncodeToBytes(r.signature) r.raw, _ = rlp.EncodeToBytes(list)
if err != nil {
return err
}
r.raw = append(blob, sigcontent...)
// mark record ready for encoding
r.signed = true
return nil return nil
} }
func (r *Record) verifySignature(sigcontent []byte) error { func (r *Record) verifySignature() error {
// Get identity scheme, public key. // Get identity scheme, public key, signature.
var id ID var id ID
var secp256k1 Secp256k1 var secp256k1 Secp256k1
if _, err := r.Load(&id); err != nil { if ok, err := r.Load(&id); err != nil {
return err return err
} } else if !ok {
if id != ID_SECP256k1_KECCAK { return fmt.Errorf("can't verify signature: missing %q key", id.ENRKey())
} else if id != ID_SECP256k1_KECCAK {
return errNoID return errNoID
} }
if _, err := r.Load(&secp256k1); err != nil { if ok, err := r.Load(&secp256k1); err != nil {
return err return err
} else if !ok {
return fmt.Errorf("can't verify signature: missing %q key", secp256k1.ENRKey())
} }
// Verify the signature.
sig, err := parseCompactSignature(r.signature) sig, err := parseCompactSignature(r.signature)
if err != nil { if err != nil {
return err return err
} }
if !sig.Verify(crypto.Keccak256(sigcontent), (*btcec.PublicKey)(&secp256k1)) {
return errors.New("signature is not valid") // Verify the signature.
list := make([]interface{}, 0, len(r.pairs)*2+1)
list = r.appendPairs(list)
h := sha3.NewKeccak256()
rlp.Encode(h, list)
if !sig.Verify(h.Sum(nil), (*btcec.PublicKey)(&secp256k1)) {
return errInvalidSig
} }
return nil return nil
} }

View file

@ -201,10 +201,6 @@ func TestSignEncodeAndDecode(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
if ok, err := r.Equal(r2); err != nil || !ok {
t.Errorf("records not equal ; got\n%#v, expected\n%#v", r2, r)
}
if !reflect.DeepEqual(r, r2) { if !reflect.DeepEqual(r, r2) {
t.Errorf("records not deep equal ; got\n%#v, expected\n%#v", r2, r) t.Errorf("records not deep equal ; got\n%#v, expected\n%#v", r2, r)
} }
@ -243,3 +239,32 @@ func TestNodeAddress(t *testing.T) {
t.Errorf("got\n%#v, expected\n%#v", got, expected) t.Errorf("got\n%#v, expected\n%#v", got, expected)
} }
} }
func TestPythonInterop(t *testing.T) {
enc, _ := hex.DecodeString("f896b840638a54215d80a6713c8d523a6adc4e6e73652d859103a36b700851cb0e61b66b8ebfc1a610c57d732ec6e0a8f06a9a7a28df5051ece514702ff9cdff0b11f454018664697363763582765f82696490736563703235366b312d6b656363616b83697034847f00000189736563703235366b31a103ca634cae0d49acb401d8a4c6b6fe8c55b70d115bf400769cc1400f3258cd3138")
var r Record
if err := rlp.DecodeBytes(enc, &r); err != nil {
t.Fatalf("can't decode: %v", err)
}
var (
wantAddr, _ = hex.DecodeString("caaa1485d83b18b32ed9ad666026151bf0cae8a0a88c857ae2d4c5be2daa6726")
wantSeq = uint32(1)
wantIP = IP4(net.ParseIP("127.0.0.1").To4())
wantDiscport = DiscPort(30303)
)
if r.Seq() != wantSeq {
t.Errorf("wrong seq: got %d, want %d", r.Seq(), wantSeq)
}
if addr, _ := r.NodeAddr(); !bytes.Equal(addr, wantAddr) {
t.Errorf("wrong addr: got %x, want %x", addr, wantAddr)
}
want := map[Key]interface{}{new(IP4): &wantIP, new(DiscPort): &wantDiscport}
for k, v := range want {
if _, err := r.Load(k); err != nil {
t.Errorf("can't load %q: %v", k.ENRKey(), err)
} else if !reflect.DeepEqual(k, v) {
t.Errorf("wrong %q: got %v, want %v", k.ENRKey(), k, v)
}
}
}