p2p/enr: update r.raw on DecodeRLP

This commit is contained in:
Anton Evangelatov 2017-12-01 19:20:51 +01:00 committed by Felix Lange
parent 0315eece6d
commit 74eda001aa
2 changed files with 30 additions and 13 deletions

View file

@ -135,7 +135,20 @@ func (r *Record) DecodeRLP(s *rlp.Stream) error {
return err return err
} }
err = r.verifySignature() 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 { if err != nil {
return err return err
} }
@ -161,11 +174,11 @@ func (r Record) Equal(o Record) (bool, error) {
return false, nil return false, nil
} }
if err := r.verifySignature(); err != nil { if err := r.verifySignature(rr); err != nil {
return false, err return false, err
} }
if err := o.verifySignature(); err != nil { if err := o.verifySignature(oo); err != nil {
return false, err return false, err
} }
@ -241,7 +254,7 @@ func (r *Record) signAndEncode(privkey *ecdsa.PrivateKey) error {
return nil return nil
} }
func (r *Record) verifySignature() error { func (r Record) verifySignature(sigcontent []byte) error {
var id ID var id ID
_, err := r.Load(&id) _, err := r.Load(&id)
if err != nil { if err != nil {
@ -254,18 +267,13 @@ func (r *Record) verifySignature() error {
} }
// get publickey from record // get publickey from record
var blob Secp256k1 var secp256k1 Secp256k1
_, err = r.Load(&blob) _, err = r.Load(&secp256k1)
if err != nil { if err != nil {
return err return err
} }
pk, err := btcec.ParsePubKey(blob, btcec.S256()) pk, err := btcec.ParsePubKey(secp256k1, btcec.S256())
if err != nil {
return err
}
sigcontent, err := r.serialisedContent()
if err != nil { if err != nil {
return err return err
} }

View file

@ -20,6 +20,7 @@ import (
"bytes" "bytes"
"encoding/hex" "encoding/hex"
"net" "net"
"reflect"
"testing" "testing"
"github.com/btcsuite/btcd/btcec" "github.com/btcsuite/btcd/btcec"
@ -182,10 +183,18 @@ func TestSignEncodeAndDecode(t *testing.T) {
t.Errorf("records not equal ; got\n%#v, expected\n%#v", r2, r) t.Errorf("records not equal ; got\n%#v, expected\n%#v", r2, r)
} }
_, err = rlp.EncodeToBytes(r2) if !reflect.DeepEqual(r, r2) {
t.Errorf("records not deep equal ; got\n%#v, expected\n%#v", r2, r)
}
blob2, err := rlp.EncodeToBytes(r2)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
if bytes.Compare(blob, blob2) != 0 {
t.Errorf("serialised records not equal ; got\n%#v, expected\n%#v", blob2, blob)
}
} }
func TestNodeAddress(t *testing.T) { func TestNodeAddress(t *testing.T) {