crypto: abstracted Key and added Equals.
Juan Batiz-Benet committed
Sep 27, 2014 at 00:18 UTC
484d6004f75fb23f26a2e110f67f0d06c577cbd4
4 files changed
+75
-18
crypto/key.go
+23
-6
@@ -23,7 +23,17 @@ const (
23
RSA = iota
24
)
25
26
+type Key interface {
27
+ // Bytes returns a serialized, storeable representation of this key
28
+ Bytes() ([]byte, error)
29
+
30
+ // Equals checks whether two PubKeys are the same
31
+ Equals(Key) bool
32
+}
33
+
34
type PrivKey interface {
35
+ Key
36
+
37
// Cryptographically sign the given bytes
38
Sign([]byte) ([]byte, error)
39
@@ -32,17 +42,13 @@ type PrivKey interface {
42
43
// Generate a secret string of bytes
44
GenSecret() []byte
35
-
36
- // Bytes returns a serialized, storeable representation of this key
37
- Bytes() ([]byte, error)
45
}
46
47
type PubKey interface {
48
+ Key
49
+
50
// Verify that 'sig' is the signed hash of 'data'
51
Verify(data []byte, sig []byte) (bool, error)
43
-
44
- // Bytes returns a serialized, storeable representation of this key
45
- Bytes() ([]byte, error)
52
}
53
54
// Given a public key, generates the shared key.
@@ -229,3 +235,14 @@ func UnmarshalPrivateKey(data []byte) (PrivKey, error) {
235
return nil, ErrBadKeyType
236
}
237
}
238
+
239
+// KeyEqual checks whether two
240
+func KeyEqual(k1, k2 Key) bool {
241
+ if k1 == k2 {
242
+ return true
243
+ }
244
+
245
+ b1, err1 := k1.Bytes()
246
+ b2, err2 := k2.Bytes()
247
+ return bytes.Equal(b1, b2) && err1 == err2
248
+}
crypto/key_test.go
+41
-1
@@ -3,12 +3,14 @@ package crypto
3
import "testing"
4
5
func TestRsaKeys(t *testing.T) {
6
- sk, _, err := GenerateKeyPair(RSA, 512)
6
+ sk, pk, err := GenerateKeyPair(RSA, 512)
7
if err != nil {
8
t.Fatal(err)
9
}
10
testKeySignature(t, sk)
11
testKeyEncoding(t, sk)
12
+ testKeyEquals(t, sk)
13
+ testKeyEquals(t, pk)
14
}
15
16
func testKeySignature(t *testing.T, sk PrivKey) {
@@ -52,3 +54,41 @@ func testKeyEncoding(t *testing.T, sk PrivKey) {
54
t.Fatal(err)
55
}
56
}
57
+
58
+func testKeyEquals(t *testing.T, k Key) {
59
+ kb, err := k.Bytes()
60
+ if err != nil {
61
+ t.Fatal(err)
62
+ }
63
+
64
+ if !KeyEqual(k, k) {
65
+ t.Fatal("Key not equal to itself.")
66
+ }
67
+
68
+ if !KeyEqual(k, testkey(kb)) {
69
+ t.Fatal("Key not equal to key with same bytes.")
70
+ }
71
+
72
+ sk, pk, err := GenerateKeyPair(RSA, 512)
73
+ if err != nil {
74
+ t.Fatal(err)
75
+ }
76
+
77
+ if KeyEqual(k, sk) {
78
+ t.Fatal("Keys should not equal.")
79
+ }
80
+
81
+ if KeyEqual(k, pk) {
82
+ t.Fatal("Keys should not equal.")
83
+ }
84
+}
85
+
86
+type testkey []byte
87
+
88
+func (pk testkey) Bytes() ([]byte, error) {
89
+ return pk, nil
90
+}
91
+
92
+func (pk testkey) Equals(k Key) bool {
93
+ return KeyEqual(pk, k)
94
+}
crypto/rsa.go
+10
@@ -41,6 +41,11 @@ func (pk *RsaPublicKey) Bytes() ([]byte, error) {
41
return proto.Marshal(pbmes)
42
}
43
44
+// Equals checks whether this key is equal to another
45
+func (pk *RsaPublicKey) Equals(k Key) bool {
46
+ return KeyEqual(pk, k)
47
+}
48
+
49
func (sk *RsaPrivateKey) GenSecret() []byte {
50
buf := make([]byte, 16)
51
rand.Read(buf)
@@ -65,6 +70,11 @@ func (sk *RsaPrivateKey) Bytes() ([]byte, error) {
70
return proto.Marshal(pbmes)
71
}
72
73
+// Equals checks whether this key is equal to another
74
+func (sk *RsaPrivateKey) Equals(k Key) bool {
75
+ return KeyEqual(sk, k)
76
+}
77
+
78
func UnmarshalRsaPrivateKey(b []byte) (*RsaPrivateKey, error) {
79
sk, err := x509.ParsePKCS1PrivateKey(b)
80
if err != nil {
crypto/spipe/handshake.go
+1
-11
@@ -379,17 +379,7 @@ func getOrConstructPeer(peers peer.Peerstore, rpk ci.PubKey) (*peer.Peer, error)
379
// did have pubkey, let's verify it's really the same.
380
// this shouldn't ever happen, given we hashed, etc, but it could mean
381
// expected code (or protocol) invariants violated.
382
-
383
- lb, err1 := npeer.PubKey.Bytes()
384
- if err1 != nil {
385
- return nil, err1
386
- }
387
- rb, err2 := rpk.Bytes()
388
- if err2 != nil {
389
- return nil, err2
390
- }
391
-
392
- if !bytes.Equal(lb, rb) {
382
+ if !npeer.PubKey.Equals(rpk) {
383
return nil, fmt.Errorf("WARNING: PubKey mismatch: %v", npeer.ID.Pretty())
384
}
385
return npeer, nil