diff --git a/p2p/enr/discv5.go b/p2p/enr/discv5.go new file mode 100644 index 0000000000..ae99c56134 --- /dev/null +++ b/p2p/enr/discv5.go @@ -0,0 +1,25 @@ +package enr + +import ( + "io" + + "github.com/ethereum/go-ethereum/rlp" +) + +type DiscV5 uint32 + +func (DiscV5) ENRKey() string { + return "discv5" +} + +func (v DiscV5) EncodeRLP(w io.Writer) error { + port := uint32(v) + return rlp.Encode(w, port) +} + +func (v *DiscV5) DecodeRLP(s *rlp.Stream) error { + if err := s.Decode((*uint32)(v)); err != nil { + return err + } + return nil +} diff --git a/p2p/enr/enr.go b/p2p/enr/enr.go index a6b04e5fd0..b1ccf78d6d 100644 --- a/p2p/enr/enr.go +++ b/p2p/enr/enr.go @@ -20,9 +20,8 @@ package enr import ( "bytes" "crypto/ecdsa" - "encoding/binary" "errors" - "net" + "io" "sort" "github.com/ethereum/go-ethereum/crypto" @@ -31,203 +30,93 @@ import ( "github.com/btcsuite/btcd/btcec" ) +var () + // The maximum encoded size of a node record is 300 bytes. Implementations should reject records larger than this size. const ( RECORD_MAX_SIZE = 300 ID_SECP256k1_KECCAK = "secp256k1-keccak" // "secp256k1-keccak" identity scheme identifier ) -// Pseudo-const identifiers for pre-defined keys -var ( - ID = []byte(`id`) // name of identity scheme, e.g. "secp256k1-keccak" - SECP256K1 = []byte(`secp256k1`) // compressed secp256k1 public key - IP4 = []byte(`ip4`) // IPv4 address, 4 bytes - IP6 = []byte(`ip6`) // IPv6 address, 16 bytes - DISCV5 = []byte(`discv5`) // UDP port for discovery v5 -) +// Key is implemented by known node record key types. +// +// To define a new key that is to be included in a node record, +// create a Go type that satisfies this interface. The type should +// also implement rlp.Decoder if additional checks are needed on the value. +type Key interface { + ENRKey() string +} -type record struct { +type pair struct { k []byte v []byte } -type ENR struct { - seq uint32 // sequence number - signature []byte // record's signature - raw []byte // RLP encoded record - records []record // list of all key/value pairs, sorted prior to RLP encoding - dirty bool // keeps track if record was modified after it was signed and encoded +type Record struct { + seq uint32 // sequence number + signature []byte // record's signature + raw []byte // RLP encoded record + pairs []pair // list of all key/value pairs, sorted prior to RLP encoding + signed bool // keeps track if record was modified after it was signed and encoded } -func NewENR() *ENR { - return &ENR{ - dirty: true, - } +func (r Record) Seq() uint32 { + return r.seq } -func (e *ENR) GetID() (string, error) { - for _, r := range e.records { - if bytes.Compare(ID, r.k) == 0 { - return string(r.v), nil +func (r *Record) SetSeq(s uint32) { + r.signed = false + r.seq = s +} + +func (r *Record) Load(k Key) (bool, error) { + for _, p := range r.pairs { + if string(p.k) == k.ENRKey() { + err := rlp.DecodeBytes(p.v, k) + return true, err } } - return "", errors.New("id record does not exist") + return false, errors.New("record does not exist") } -func (e *ENR) SetID(id string) { - e.dirty = true - e.records = append(e.records, record{ID, []byte(id)}) -} - -func (e *ENR) GetSecp256k1() ([]byte, error) { - for _, r := range e.records { - if bytes.Compare(SECP256K1, r.k) == 0 { - return r.v, nil - } +func (r *Record) Set(k Key) error { + r.signed = false + blob, err := rlp.EncodeToBytes(k) + if err != nil { + return err } - - return nil, errors.New("secp256k1 record does not exist") -} - -func (e *ENR) SetSecp256k1(pk []byte) { - e.dirty = true - e.records = append(e.records, record{SECP256K1, pk}) -} - -func (e *ENR) GetIPv4() (net.IP, error) { - for _, r := range e.records { - if bytes.Compare(IP4, r.k) == 0 { - if len(r.v) != net.IPv4len { - return nil, errors.New("wrong ipv4 record length") - } - return net.IP(r.v), nil - } - } - - return nil, errors.New("ip4 record does not exist") -} - -func (e *ENR) SetIPv4(ip net.IP) error { - e.dirty = true - ipv4 := ip.To4() - if ipv4 == nil { - return errors.New("param is not a valid ipv4 address") - } - e.records = append(e.records, record{IP4, ipv4}) + r.pairs = append(r.pairs, pair{[]byte(k.ENRKey()), blob}) return nil } -func (e *ENR) GetIPv6() (net.IP, error) { - for _, r := range e.records { - if bytes.Compare(IP6, r.k) == 0 { - if len(r.v) != net.IPv6len { - return nil, errors.New("wrong ipv6 record length") - } - return net.IP(r.v), nil - } +func (r Record) EncodeRLP(w io.Writer) error { + if !r.signed { + return errors.New("record is not signed") } - return nil, errors.New("ip6 record does not exist") + _, err := w.Write(r.raw) + + return err } -func (e *ENR) SetIPv6(ip net.IP) error { - if len(ip) != net.IPv6len { - return errors.New("param length is not equal to 16 bytes") - } - e.dirty = true - e.records = append(e.records, record{IP6, ip}) - return nil -} +func (r *Record) DecodeRLP(s *rlp.Stream) error { + var err error -func (e *ENR) GetDiscv5() (uint32, error) { - for _, r := range e.records { - if bytes.Compare(DISCV5, r.k) == 0 { - buf := bytes.NewBuffer(r.v) - var port uint32 - err := binary.Read(buf, binary.BigEndian, &port) - if err != nil { - return 0, err - } - - return port, nil - } - } - - return 0, errors.New("secp256k1 record does not exist") -} - -func (e *ENR) SetDiscv5(port uint32) error { - e.dirty = true - buf := new(bytes.Buffer) - err := binary.Write(buf, binary.BigEndian, port) + r.signature, err = s.Bytes() if err != nil { return err } - e.records = append(e.records, record{DISCV5, buf.Bytes()}) - - return nil -} - -func (e *ENR) GetRaw(k []byte) ([]byte, error) { - for _, r := range e.records { - if bytes.Compare(r.k, k) == 0 { - return r.v, nil - } - } - - return nil, errors.New("record does not exist") -} - -func (e *ENR) SetRaw(k []byte, v []byte) { - e.dirty = true - e.records = append(e.records, record{k, v}) -} - -func (e *ENR) Encode() ([]byte, error) { - if e.dirty { - return nil, errors.New("record is not signed") - } - return e.raw, nil -} - -func (e *ENR) NodeAddress() ([]byte, error) { - pk, err := e.GetSecp256k1() - if err != nil { - return nil, err - } - - digest := crypto.Keccak256Hash(pk) - - return digest.Bytes(), nil -} - -func (e *ENR) Decode(data []byte) error { - if len(data) > RECORD_MAX_SIZE { - return errors.New("record is too big") - } - - s := rlp.NewStream(bytes.NewReader(data), RECORD_MAX_SIZE) - - signature, err := s.Bytes() - if err != nil { - return err - } - - // consume the list prefix _, err = s.List() if err != nil { return err } - seq, err := s.Uint() - if err != nil { + if err := s.Decode(&r.seq); err != nil { return err } - var records []record - // read key/value pairs until we reach rlp.EOL for _, _, err = s.Kind(); err == nil; _, _, err = s.Kind() { key, err2 := s.Bytes() @@ -240,115 +129,120 @@ func (e *ENR) Decode(data []byte) error { return err2 } - records = append(records, record{k: key, v: value}) + r.pairs = append(r.pairs, pair{k: key, v: value}) } if err != rlp.EOL { return err } - e.signature = signature - e.raw = data - e.seq = uint32(seq) - e.records = records - - err = e.verifySignature() - if err != nil { - return err - } - - e.dirty = false - - return nil -} - -func (e *ENR) Sign(privkey *ecdsa.PrivateKey) error { - e.seq = e.seq + 1 - - e.SetID(ID_SECP256k1_KECCAK) - - pk := (*btcec.PublicKey)(&privkey.PublicKey).SerializeCompressed() - e.SetSecp256k1(pk) - - var err error - e.signature, e.raw, err = e.SignAndEncode(privkey) + err = r.verifySignature() if err != nil { return err } // mark record ready for encoding - e.dirty = false + r.signed = true return nil } -func (e *ENR) SignAndEncode(privkey *ecdsa.PrivateKey) ([]byte, []byte, error) { - content, err := e.SerialisedContent() +func (r Record) Equal(o Record) (bool, error) { + rr, err := r.serialisedContent() if err != nil { - return nil, nil, err + return false, err } - digest := crypto.Keccak256Hash(content) - - signature, err := crypto.Sign(digest.Bytes(), privkey) + oo, err := o.serialisedContent() if err != nil { - return nil, nil, err + return false, err } - blob, err := rlp.EncodeToBytes(signature) - if err != nil { - return nil, nil, err + if bytes.Compare(rr, oo) != 0 { + return false, nil } - raw := append(blob, content...) + if err := r.verifySignature(); err != nil { + return false, err + } - return signature, raw, nil + if err := o.verifySignature(); err != nil { + return false, err + } + + return true, nil } -func (e *ENR) SerialisedContent() ([]byte, error) { - var buffer bytes.Buffer +func (r *Record) NodeAddr() ([]byte, error) { + var secp256k1 Secp256k1 - blob, err := rlp.EncodeToBytes(e.seq) + _, err := r.Load(&secp256k1) if err != nil { return nil, err } - _, err = buffer.Write(blob) - if err != nil { - return nil, err - } + digest := crypto.Keccak256Hash(secp256k1) - sort.Slice(e.records, func(i, j int) bool { - return bytes.Compare(e.records[i].k, e.records[j].k) < 0 + return digest.Bytes(), nil +} + +func (r *Record) Sign(privkey *ecdsa.PrivateKey) error { + r.seq = r.seq + 1 + + id := ID(ID_SECP256k1_KECCAK) + + r.Set(id) + + pk := (*btcec.PublicKey)(&privkey.PublicKey).SerializeCompressed() + secp256k1 := Secp256k1(pk) + r.Set(secp256k1) + + return r.signAndEncode(privkey) +} + +func (r *Record) serialisedContent() ([]byte, error) { + sort.Slice(r.pairs, func(i, j int) bool { + return string(r.pairs[i].k) < string(r.pairs[j].k) }) - for _, r := range e.records { - kk, err := rlp.EncodeToBytes(r.k) - if err != nil { - return nil, err - } + list := []interface{}{r.seq} - _, err = buffer.Write(kk) - if err != nil { - return nil, err - } - - vv, err := rlp.EncodeToBytes(r.v) - if err != nil { - return nil, err - } - - _, err = buffer.Write(vv) - if err != nil { - return nil, err - } + for _, p := range r.pairs { + list = append(list, p.k, p.v) } - return wrapList(buffer.Bytes()), nil + return rlp.EncodeToBytes(list) } -func (e *ENR) verifySignature() error { - id, err := e.GetID() +func (r *Record) signAndEncode(privkey *ecdsa.PrivateKey) error { + sigcontent, err := r.serialisedContent() + if err != nil { + return err + } + + digest := crypto.Keccak256Hash(sigcontent) + + r.signature, err = crypto.Sign(digest.Bytes(), privkey) + if err != nil { + return err + } + + blob, err := rlp.EncodeToBytes(r.signature) + if err != nil { + return err + } + + r.raw = append(blob, sigcontent...) + + // mark record ready for encoding + r.signed = true + + return nil +} + +func (r *Record) verifySignature() error { + var id ID + _, err := r.Load(&id) if err != nil { return err } @@ -359,7 +253,8 @@ func (e *ENR) verifySignature() error { } // get publickey from record - blob, err := e.GetSecp256k1() + var blob Secp256k1 + _, err = r.Load(&blob) if err != nil { return err } @@ -371,13 +266,14 @@ func (e *ENR) verifySignature() error { pubkey1 := pk.SerializeUncompressed() // get publickey from message and signature - content, err := e.SerialisedContent() + sigcontent, err := r.serialisedContent() if err != nil { return err } - digest := crypto.Keccak256Hash(content) - pubkey2, err := crypto.Ecrecover(digest.Bytes(), e.signature) + digest := crypto.Keccak256Hash(sigcontent) + + pubkey2, err := crypto.Ecrecover(digest.Bytes(), r.signature) if err != nil { return err } @@ -388,9 +284,3 @@ func (e *ENR) verifySignature() error { return nil } - -func wrapList(c []byte) []byte { - head := make([]byte, 9) - res := rlp.LengthPrefix(head, uint64(len(c))) - return append(res, c...) -} diff --git a/p2p/enr/enr_test.go b/p2p/enr/enr_test.go index a059e902dd..4c1df03f67 100644 --- a/p2p/enr/enr_test.go +++ b/p2p/enr/enr_test.go @@ -20,11 +20,11 @@ import ( "bytes" "encoding/hex" "net" - "reflect" "testing" "github.com/btcsuite/btcd/btcec" "github.com/ethereum/go-ethereum/crypto" + "github.com/ethereum/go-ethereum/rlp" ) const ( @@ -32,113 +32,119 @@ const ( ) func TestGetSetID(t *testing.T) { - id := "someid" - e := NewENR() - e.SetID(id) + id := ID("someid") + var r Record + r.Set(id) - got, err := e.GetID() + var id2 ID + + _, err := r.Load(&id2) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - if got != id { - t.Fatalf("got %#v, expected %#v", got, id) + if id != id2 { + t.Fatalf("got %#v, expected %#v", id2, id) } } func TestGetSetIP4(t *testing.T) { - ip := net.IP{192, 168, 0, 3} - e := NewENR() - e.SetIPv4(ip) + ip := IP4(net.IP{192, 168, 0, 3}) + var r Record + r.Set(ip) - got, err := e.GetIPv4() + var ip2 IP4 + + _, err := r.Load(&ip2) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - if !got.Equal(ip) { - t.Fatalf("got %#v, expected %#v", got, ip) + if bytes.Compare(ip, ip2) != 0 { + t.Fatalf("got %#v, expected %#v", ip2, ip) } } func TestGetSetIP6(t *testing.T) { - ip := net.IP{0x20, 0x01, 0x48, 0x60, 0, 0, 0x20, 0x01, 0, 0, 0, 0, 0, 0, 0x00, 0x68} - e := NewENR() - e.SetIPv6(ip) + ip := IP6(net.IP{0x20, 0x01, 0x48, 0x60, 0, 0, 0x20, 0x01, 0, 0, 0, 0, 0, 0, 0x00, 0x68}) + var r Record + r.Set(ip) - got, err := e.GetIPv6() + var ip2 IP6 + + _, err := r.Load(&ip2) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - if !got.Equal(ip) { - t.Fatalf("got %#v, expected %#v", got, ip) + if bytes.Compare(ip, ip2) != 0 { + t.Fatalf("got %#v, expected %#v", ip2, ip) } } func TestGetSetDiscv5(t *testing.T) { - port := uint32(30309) - e := NewENR() + port := DiscV5(30309) + var r Record + r.Set(port) - err := e.SetDiscv5(port) + var port2 DiscV5 + + _, err := r.Load(&port2) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - got, err := e.GetDiscv5() - if err != nil { - t.Fatalf("error: %#v", err) - } - - if got != port { - t.Fatalf("got %#v, expected %#v", got, port) + if port != port2 { + t.Fatalf("got %#v, expected %#v", port2, port) } } func TestGetSetSecp256k1(t *testing.T) { privkey, err := crypto.HexToECDSA(privkeyHex) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - e := NewENR() + var r Record - err = e.Sign(privkey) + err = r.Sign(privkey) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - got, err := e.GetSecp256k1() + var pk Secp256k1 + + _, err = r.Load(&pk) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } expected := (*btcec.PublicKey)(&privkey.PublicKey).SerializeCompressed() - if bytes.Compare(got, expected) != 0 { - t.Fatalf("got %#v, expected %#v", got, expected) + if bytes.Compare(pk, expected) != 0 { + t.Fatalf("got %#v, expected %#v", pk, expected) } } func TestDirty(t *testing.T) { privkey, err := crypto.HexToECDSA(privkeyHex) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - e := NewENR() + var r Record - err = e.Sign(privkey) + err = r.Sign(privkey) if err != nil { - t.Fatalf("error: %#v", err) + t.Fatal(err) } - if _, err := e.Encode(); err != nil { - t.Fatalf("error: %#v", err) + if _, err := rlp.EncodeToBytes(r); err != nil { + t.Fatal(err) } - e.SetRaw([]byte(`some key`), []byte(`some value`)) + r.SetSeq(3) - if _, err := e.Encode(); err == nil { + if _, err := rlp.EncodeToBytes(r); err == nil { t.Fatal("expected err, got nil") } } @@ -149,47 +155,37 @@ func TestSignEncodeAndDecode(t *testing.T) { t.Fatal(err) } - e := NewENR() - e.SetDiscv5(30303) - e.SetIPv4(net.ParseIP("127.0.0.1")) + var r Record + port := DiscV5(30303) + r.Set(port) - err = e.Sign(privkey) + ipv4 := IP4(net.ParseIP("127.0.0.1")) + r.Set(ipv4) + + err = r.Sign(privkey) if err != nil { t.Fatal(err) } - record, err := e.Encode() + blob, err := rlp.EncodeToBytes(r) if err != nil { t.Fatal(err) } - e2 := NewENR() - - err = e2.Decode(record) + var r2 Record + err = rlp.DecodeBytes(blob, &r2) if err != nil { t.Fatal(err) } - if !reflect.DeepEqual(e, e2) { - t.Errorf("got\n%#v, expected\n%#v", e2, e) + if ok, err := r.Equal(r2); err != nil || !ok { + t.Errorf("records not equal ; got\n%#v, expected\n%#v", r2, r) } - expectedRecord := "b8415571f9a36b1e26c366745894656dd1565033cdeda18d330c9a9bac67dfc3e786556b0490e509372fa9db5abf418accd895467e8ff047bbdc147789bef71a2cc401f8560186646973637635840000765f82696490736563703235366b312d6b656363616b83697034847f00000189736563703235366b31a103ca634cae0d49acb401d8a4c6b6fe8c55b70d115bf400769cc1400f3258cd3138" - - got := hex.EncodeToString(record) - if got != expectedRecord { - t.Errorf("got\n%#v, expected\n%#v", got, expectedRecord) - } - - blob, err := e2.Encode() + _, err = rlp.EncodeToBytes(r2) if err != nil { t.Fatal(err) } - - got = hex.EncodeToString(blob) - if got != expectedRecord { - t.Errorf("got\n%#v, expected\n%#v", got, expectedRecord) - } } func TestNodeAddress(t *testing.T) { @@ -198,14 +194,14 @@ func TestNodeAddress(t *testing.T) { t.Fatal(err) } - e := NewENR() + var r Record - err = e.Sign(privkey) + err = r.Sign(privkey) if err != nil { t.Fatal(err) } - addr, err := e.NodeAddress() + addr, err := r.NodeAddr() if err != nil { t.Fatal(err) } diff --git a/p2p/enr/id.go b/p2p/enr/id.go new file mode 100644 index 0000000000..8f7106f311 --- /dev/null +++ b/p2p/enr/id.go @@ -0,0 +1,25 @@ +package enr + +import ( + "io" + + "github.com/ethereum/go-ethereum/rlp" +) + +type ID string + +func (ID) ENRKey() string { + return "id" +} + +func (v ID) EncodeRLP(w io.Writer) error { + id := string(v) + return rlp.Encode(w, id) +} + +func (v *ID) DecodeRLP(s *rlp.Stream) error { + if err := s.Decode((*string)(v)); err != nil { + return err + } + return nil +} diff --git a/p2p/enr/ip4.go b/p2p/enr/ip4.go new file mode 100644 index 0000000000..530e609516 --- /dev/null +++ b/p2p/enr/ip4.go @@ -0,0 +1,35 @@ +package enr + +import ( + "fmt" + "io" + "net" + + "github.com/ethereum/go-ethereum/rlp" +) + +// IP4 represents an 4-byte IPv4 address in a node record. +type IP4 net.IP + +// ENRKey returns the node record key for an IPv4 address. +func (IP4) ENRKey() string { + return "ip4" +} + +func (v IP4) EncodeRLP(w io.Writer) error { + ip4 := net.IP(v).To4() + if ip4 == nil { + return fmt.Errorf("invalid IPv4 address: %v", v) + } + return rlp.Encode(w, ip4) +} + +func (v *IP4) DecodeRLP(s *rlp.Stream) error { + if err := s.Decode((*net.IP)(v)); err != nil { + return err + } + if len(*v) != 4 { + return fmt.Errorf("invalid IPv4 address, want 4 bytes: %v", *v) + } + return nil +} diff --git a/p2p/enr/ip6.go b/p2p/enr/ip6.go new file mode 100644 index 0000000000..5eea18d5fa --- /dev/null +++ b/p2p/enr/ip6.go @@ -0,0 +1,32 @@ +package enr + +import ( + "fmt" + "io" + "net" + + "github.com/ethereum/go-ethereum/rlp" +) + +// IP6 represents an 16-byte IPv6 address in a node record. +type IP6 net.IP + +// ENRKey returns the node record key for an IPv6 address. +func (IP6) ENRKey() string { + return "ip6" +} + +func (v IP6) EncodeRLP(w io.Writer) error { + ip6 := net.IP(v) + return rlp.Encode(w, ip6) +} + +func (v *IP6) DecodeRLP(s *rlp.Stream) error { + if err := s.Decode((*net.IP)(v)); err != nil { + return err + } + if len(*v) != 16 { + return fmt.Errorf("invalid IPv6 address, want 16 bytes: %v", *v) + } + return nil +} diff --git a/p2p/enr/secp256k1.go b/p2p/enr/secp256k1.go new file mode 100644 index 0000000000..1b4e61d65a --- /dev/null +++ b/p2p/enr/secp256k1.go @@ -0,0 +1,25 @@ +package enr + +import ( + "io" + + "github.com/ethereum/go-ethereum/rlp" +) + +type Secp256k1 []byte + +func (Secp256k1) ENRKey() string { + return "secp256k1" +} + +func (v Secp256k1) EncodeRLP(w io.Writer) error { + blob := []byte(v) + return rlp.Encode(w, blob) +} + +func (v *Secp256k1) DecodeRLP(s *rlp.Stream) error { + if err := s.Decode((*[]byte)(v)); err != nil { + return err + } + return nil +}