From 03f723754152d3e4822203602d146f36ccea80d4 Mon Sep 17 00:00:00 2001 From: Filippo Valsorda Date: Wed, 27 Nov 2019 22:49:40 -0400 Subject: [PATCH] internal/age: refactor to use the new golang.org/x/crypto/curve25519 API --- cmd/age/encrypted_keys.go | 14 +--- internal/age/age_test.go | 17 ++--- internal/age/recipients_test.go | 42 ++++++------ internal/age/ssh.go | 118 ++++++++++++++------------------ internal/age/x25519.go | 4 +- internal/curve25519/x25519.go | 86 ----------------------- 6 files changed, 84 insertions(+), 197 deletions(-) delete mode 100644 internal/curve25519/x25519.go diff --git a/cmd/age/encrypted_keys.go b/cmd/age/encrypted_keys.go index fc6b26f..74752e5 100644 --- a/cmd/age/encrypted_keys.go +++ b/cmd/age/encrypted_keys.go @@ -7,10 +7,8 @@ package main import ( - "bytes" "crypto/ed25519" "crypto/rsa" - "crypto/sha256" "fmt" "os" @@ -86,18 +84,8 @@ func (i *EncryptedSSHIdentity) Matches(block *format.Recipient) error { if len(block.Args) != 1 { return fmt.Errorf("invalid %v recipient block", i.Type()) } - hash, err := format.DecodeString(block.Args[0]) - if err != nil { - return fmt.Errorf("failed to parse %v recipient: %v", i.Type(), err) - } - if len(hash) != 4 { - return fmt.Errorf("invalid %v recipient block", i.Type()) - } - sH := sha256.New() - sH.Write(i.pubKey.Marshal()) - hh := sH.Sum(nil) - if !bytes.Equal(hh[:4], hash) { + if block.Args[0] != age.SSHFingerprint(i.pubKey) { return age.ErrIncorrectIdentity } return nil diff --git a/internal/age/age_test.go b/internal/age/age_test.go index 0697a05..7c9a834 100644 --- a/internal/age/age_test.go +++ b/internal/age/age_test.go @@ -20,21 +20,22 @@ import ( const helloWorld = "Hello, Twitch!" func TestEncryptDecryptX25519(t *testing.T) { - var secretKeyA, publicKeyA, secretKeyB, publicKeyB [32]byte - if _, err := rand.Read(secretKeyA[:]); err != nil { + secretKeyA := make([]byte, curve25519.ScalarSize) + secretKeyB := make([]byte, curve25519.ScalarSize) + if _, err := rand.Read(secretKeyA); err != nil { t.Fatal(err) } - if _, err := rand.Read(secretKeyB[:]); err != nil { + if _, err := rand.Read(secretKeyB); err != nil { t.Fatal(err) } - curve25519.ScalarBaseMult(&publicKeyA, &secretKeyA) - curve25519.ScalarBaseMult(&publicKeyB, &secretKeyB) + publicKeyA, _ := curve25519.X25519(secretKeyA, curve25519.Basepoint) + publicKeyB, _ := curve25519.X25519(secretKeyB, curve25519.Basepoint) - rA, err := age.NewX25519Recipient(publicKeyA[:]) + rA, err := age.NewX25519Recipient(publicKeyA) if err != nil { t.Fatal(err) } - rB, err := age.NewX25519Recipient(publicKeyB[:]) + rB, err := age.NewX25519Recipient(publicKeyB) if err != nil { t.Fatal(err) } @@ -52,7 +53,7 @@ func TestEncryptDecryptX25519(t *testing.T) { t.Logf("%s", buf.Bytes()) - i, err := age.NewX25519Identity(secretKeyB[:]) + i, err := age.NewX25519Identity(secretKeyB) if err != nil { t.Fatal(err) } diff --git a/internal/age/recipients_test.go b/internal/age/recipients_test.go index 0be0d17..1ba6f8f 100644 --- a/internal/age/recipients_test.go +++ b/internal/age/recipients_test.go @@ -19,17 +19,17 @@ import ( ) func TestX25519RoundTrip(t *testing.T) { - var secretKey, publicKey [32]byte - if _, err := rand.Read(secretKey[:]); err != nil { + secretKey := make([]byte, curve25519.ScalarSize) + if _, err := rand.Read(secretKey); err != nil { t.Fatal(err) } - curve25519.ScalarBaseMult(&publicKey, &secretKey) + publicKey, _ := curve25519.X25519(secretKey, curve25519.Basepoint) - r, err := age.NewX25519Recipient(publicKey[:]) + r, err := age.NewX25519Recipient(publicKey) if err != nil { t.Fatal(err) } - i, err := age.NewX25519Identity(secretKey[:]) + i, err := age.NewX25519Identity(secretKey) if err != nil { t.Fatal(err) } @@ -39,10 +39,10 @@ func TestX25519RoundTrip(t *testing.T) { } fileKey := make([]byte, 16) - if _, err := rand.Read(fileKey[:]); err != nil { + if _, err := rand.Read(fileKey); err != nil { t.Fatal(err) } - block, err := r.Wrap(fileKey[:]) + block, err := r.Wrap(fileKey) if err != nil { t.Fatal(err) } @@ -55,8 +55,8 @@ func TestX25519RoundTrip(t *testing.T) { t.Fatal(err) } - if !bytes.Equal(fileKey[:], out) { - t.Errorf("invalid output: %x, expected %x", out, fileKey[:]) + if !bytes.Equal(fileKey, out) { + t.Errorf("invalid output: %x, expected %x", out, fileKey) } } @@ -78,10 +78,10 @@ func TestScryptRoundTrip(t *testing.T) { } fileKey := make([]byte, 16) - if _, err := rand.Read(fileKey[:]); err != nil { + if _, err := rand.Read(fileKey); err != nil { t.Fatal(err) } - block, err := r.Wrap(fileKey[:]) + block, err := r.Wrap(fileKey) if err != nil { t.Fatal(err) } @@ -94,8 +94,8 @@ func TestScryptRoundTrip(t *testing.T) { t.Fatal(err) } - if !bytes.Equal(fileKey[:], out) { - t.Errorf("invalid output: %x, expected %x", out, fileKey[:]) + if !bytes.Equal(fileKey, out) { + t.Errorf("invalid output: %x, expected %x", out, fileKey) } } @@ -123,10 +123,10 @@ func TestSSHRSARoundTrip(t *testing.T) { } fileKey := make([]byte, 16) - if _, err := rand.Read(fileKey[:]); err != nil { + if _, err := rand.Read(fileKey); err != nil { t.Fatal(err) } - block, err := r.Wrap(fileKey[:]) + block, err := r.Wrap(fileKey) if err != nil { t.Fatal(err) } @@ -139,8 +139,8 @@ func TestSSHRSARoundTrip(t *testing.T) { t.Fatal(err) } - if !bytes.Equal(fileKey[:], out) { - t.Errorf("invalid output: %x, expected %x", out, fileKey[:]) + if !bytes.Equal(fileKey, out) { + t.Errorf("invalid output: %x, expected %x", out, fileKey) } } @@ -168,10 +168,10 @@ func TestSSHEd25519RoundTrip(t *testing.T) { } fileKey := make([]byte, 16) - if _, err := rand.Read(fileKey[:]); err != nil { + if _, err := rand.Read(fileKey); err != nil { t.Fatal(err) } - block, err := r.Wrap(fileKey[:]) + block, err := r.Wrap(fileKey) if err != nil { t.Fatal(err) } @@ -184,7 +184,7 @@ func TestSSHEd25519RoundTrip(t *testing.T) { t.Fatal(err) } - if !bytes.Equal(fileKey[:], out) { - t.Errorf("invalid output: %x, expected %x", out, fileKey[:]) + if !bytes.Equal(fileKey, out) { + t.Errorf("invalid output: %x, expected %x", out, fileKey) } } diff --git a/internal/age/ssh.go b/internal/age/ssh.go index 6ae0818..f71dbea 100644 --- a/internal/age/ssh.go +++ b/internal/age/ssh.go @@ -7,7 +7,6 @@ package age import ( - "bytes" "crypto/ed25519" "crypto/rand" "crypto/rsa" @@ -25,6 +24,13 @@ import ( "golang.org/x/crypto/ssh" ) +func SSHFingerprint(pk ssh.PublicKey) string { + h := sha256.New() + h.Write(pk.Marshal()) + hh := h.Sum(nil) + return format.EncodeToString(hh[:4]) +} + const oaepLabel = "age-tool.com ssh-rsa" type SSHRSARecipient struct { @@ -57,13 +63,9 @@ func NewSSHRSARecipient(pk ssh.PublicKey) (*SSHRSARecipient, error) { } func (r *SSHRSARecipient) Wrap(fileKey []byte) (*format.Recipient, error) { - h := sha256.New() - h.Write(r.sshKey.Marshal()) - hh := h.Sum(nil) - l := &format.Recipient{ Type: "ssh-rsa", - Args: []string{format.EncodeToString(hh[:4])}, + Args: []string{SSHFingerprint(r.sshKey)}, } wrappedKey, err := rsa.EncryptOAEP(sha256.New(), rand.Reader, @@ -103,18 +105,8 @@ func (i *SSHRSAIdentity) Unwrap(block *format.Recipient) ([]byte, error) { if len(block.Args) != 1 { return nil, errors.New("invalid ssh-rsa recipient block") } - hash, err := format.DecodeString(block.Args[0]) - if err != nil { - return nil, fmt.Errorf("failed to parse ssh-rsa recipient: %v", err) - } - if len(hash) != 4 { - return nil, errors.New("invalid ssh-rsa recipient block") - } - h := sha256.New() - h.Write(i.sshKey.Marshal()) - hh := h.Sum(nil) - if !bytes.Equal(hh[:4], hash) { + if block.Args[0] != SSHFingerprint(i.sshKey) { return nil, ErrIncorrectIdentity } @@ -128,7 +120,7 @@ func (i *SSHRSAIdentity) Unwrap(block *format.Recipient) ([]byte, error) { type SSHEd25519Recipient struct { sshKey ssh.PublicKey - theirPublicKey [32]byte + theirPublicKey []byte } var _ Recipient = &SSHEd25519Recipient{} @@ -145,8 +137,7 @@ func NewSSHEd25519Recipient(pk ssh.PublicKey) (*SSHEd25519Recipient, error) { if pk, ok := pk.(ssh.CryptoPublicKey); ok { if pk, ok := pk.CryptoPublicKey().(ed25519.PublicKey); ok { - pubKey := ed25519PublicKeyToCurve25519(pk) - copy(r.theirPublicKey[:], pubKey) + r.theirPublicKey = ed25519PublicKeyToCurve25519(pk) } else { return nil, errors.New("unexpected public key type") } @@ -200,7 +191,7 @@ func ed25519PublicKeyToCurve25519(pk ed25519.PublicKey) []byte { u := y.Mul(y.Add(y, big.NewInt(1)), denom) u.Mod(u, curve25519P) - out := make([]byte, 32) + out := make([]byte, curve25519.PointSize) uBytes := u.Bytes() for i, b := range uBytes { out[len(uBytes)-i-1] = b @@ -212,35 +203,37 @@ func ed25519PublicKeyToCurve25519(pk ed25519.PublicKey) []byte { const ed25519Label = "age-tool.com ssh-ed25519" func (r *SSHEd25519Recipient) Wrap(fileKey []byte) (*format.Recipient, error) { - // TODO: DRY this up with the X25519 implementation. - var ephemeral, ourPublicKey [32]byte - if _, err := rand.Read(ephemeral[:]); err != nil { + ephemeral := make([]byte, curve25519.ScalarSize) + if _, err := rand.Read(ephemeral); err != nil { + return nil, err + } + ourPublicKey, err := curve25519.X25519(ephemeral, curve25519.Basepoint) + if err != nil { return nil, err } - curve25519.ScalarBaseMult(&ourPublicKey, &ephemeral) - var sharedSecret, tweak [32]byte + sharedSecret, err := curve25519.X25519(ephemeral, r.theirPublicKey) + if err != nil { + return nil, err + } + + tweak := make([]byte, curve25519.ScalarSize) tH := hkdf.New(sha256.New, nil, r.sshKey.Marshal(), []byte(ed25519Label)) - if _, err := io.ReadFull(tH, tweak[:]); err != nil { + if _, err := io.ReadFull(tH, tweak); err != nil { return nil, err } - curve25519.ScalarMult(&sharedSecret, &ephemeral, &r.theirPublicKey) - curve25519.ScalarMult(&sharedSecret, &tweak, &sharedSecret) - - sH := sha256.New() - sH.Write(r.sshKey.Marshal()) - hh := sH.Sum(nil) + sharedSecret, _ = curve25519.X25519(tweak, sharedSecret) l := &format.Recipient{ Type: "ssh-ed25519", - Args: []string{format.EncodeToString(hh[:4]), + Args: []string{SSHFingerprint(r.sshKey), format.EncodeToString(ourPublicKey[:])}, } - salt := make([]byte, 0, 32*2) - salt = append(salt, ourPublicKey[:]...) - salt = append(salt, r.theirPublicKey[:]...) - h := hkdf.New(sha256.New, sharedSecret[:], salt, []byte(ed25519Label)) + salt := make([]byte, 0, len(ourPublicKey)+len(r.theirPublicKey)) + salt = append(salt, ourPublicKey...) + salt = append(salt, r.theirPublicKey...) + h := hkdf.New(sha256.New, sharedSecret, salt, []byte(ed25519Label)) wrappingKey := make([]byte, chacha20poly1305.KeySize) if _, err := io.ReadFull(h, wrappingKey); err != nil { return nil, err @@ -256,7 +249,7 @@ func (r *SSHEd25519Recipient) Wrap(fileKey []byte) (*format.Recipient, error) { } type SSHEd25519Identity struct { - secretKey, ourPublicKey [32]byte + secretKey, ourPublicKey []byte sshKey ssh.PublicKey } @@ -270,11 +263,10 @@ func NewSSHEd25519Identity(key ed25519.PrivateKey) (*SSHEd25519Identity, error) return nil, err } i := &SSHEd25519Identity{ - sshKey: s.PublicKey(), + sshKey: s.PublicKey(), + secretKey: ed25519PrivateKeyToCurve25519(key), } - secretKey := ed25519PrivateKeyToCurve25519(key) - copy(i.secretKey[:], secretKey) - curve25519.ScalarBaseMult(&i.ourPublicKey, &i.secretKey) + i.ourPublicKey, _ = curve25519.X25519(i.secretKey, curve25519.Basepoint) return i, nil } @@ -296,54 +288,46 @@ func ParseSSHIdentity(pemBytes []byte) (Identity, error) { func ed25519PrivateKeyToCurve25519(pk ed25519.PrivateKey) []byte { h := sha512.New() - h.Write(pk[:32]) + h.Write(pk.Seed()) out := h.Sum(nil) - return out[:32] + return out[:curve25519.ScalarSize] } func (i *SSHEd25519Identity) Unwrap(block *format.Recipient) ([]byte, error) { - // TODO: DRY this up with the X25519 implementation. if block.Type != "ssh-ed25519" { return nil, ErrIncorrectIdentity } if len(block.Args) != 2 { return nil, errors.New("invalid ssh-ed25519 recipient block") } - hash, err := format.DecodeString(block.Args[0]) - if err != nil { - return nil, fmt.Errorf("failed to parse ssh-ed25519 recipient: %v", err) - } - if len(hash) != 4 { - return nil, errors.New("invalid ssh-ed25519 recipient block") - } publicKey, err := format.DecodeString(block.Args[1]) if err != nil { return nil, fmt.Errorf("failed to parse ssh-ed25519 recipient: %v", err) } - if len(publicKey) != 32 { + if len(publicKey) != curve25519.PointSize { return nil, errors.New("invalid ssh-ed25519 recipient block") } - sH := sha256.New() - sH.Write(i.sshKey.Marshal()) - hh := sH.Sum(nil) - if !bytes.Equal(hh[:4], hash) { + if block.Args[0] != SSHFingerprint(i.sshKey) { return nil, ErrIncorrectIdentity } - var sharedSecret, theirPublicKey, tweak [32]byte - copy(theirPublicKey[:], publicKey) + sharedSecret, err := curve25519.X25519(i.secretKey, publicKey) + if err != nil { + return nil, fmt.Errorf("invalid X25519 recipient: %v", err) + } + + tweak := make([]byte, curve25519.ScalarSize) tH := hkdf.New(sha256.New, nil, i.sshKey.Marshal(), []byte(ed25519Label)) - if _, err := io.ReadFull(tH, tweak[:]); err != nil { + if _, err := io.ReadFull(tH, tweak); err != nil { return nil, err } - curve25519.ScalarMult(&sharedSecret, &i.secretKey, &theirPublicKey) - curve25519.ScalarMult(&sharedSecret, &tweak, &sharedSecret) + sharedSecret, _ = curve25519.X25519(tweak, sharedSecret) - salt := make([]byte, 0, 32*2) - salt = append(salt, theirPublicKey[:]...) - salt = append(salt, i.ourPublicKey[:]...) - h := hkdf.New(sha256.New, sharedSecret[:], salt, []byte(ed25519Label)) + salt := make([]byte, 0, len(publicKey)+len(i.ourPublicKey)) + salt = append(salt, publicKey...) + salt = append(salt, i.ourPublicKey...) + h := hkdf.New(sha256.New, sharedSecret, salt, []byte(ed25519Label)) wrappingKey := make([]byte, chacha20poly1305.KeySize) if _, err := io.ReadFull(h, wrappingKey); err != nil { return nil, err diff --git a/internal/age/x25519.go b/internal/age/x25519.go index 71af804..0a97168 100644 --- a/internal/age/x25519.go +++ b/internal/age/x25519.go @@ -14,9 +14,9 @@ import ( "io" "strings" - "github.com/FiloSottile/age/internal/curve25519" "github.com/FiloSottile/age/internal/format" "golang.org/x/crypto/chacha20poly1305" + "golang.org/x/crypto/curve25519" "golang.org/x/crypto/hkdf" ) @@ -120,7 +120,7 @@ func NewX25519Identity(secretKey []byte) (*X25519Identity, error) { } func GenerateX25519Identity() (*X25519Identity, error) { - secretKey := make([]byte, 32) + secretKey := make([]byte, curve25519.ScalarSize) if _, err := rand.Read(secretKey); err != nil { return nil, fmt.Errorf("internal error: %v", err) } diff --git a/internal/curve25519/x25519.go b/internal/curve25519/x25519.go deleted file mode 100644 index ce44adf..0000000 --- a/internal/curve25519/x25519.go +++ /dev/null @@ -1,86 +0,0 @@ -// Copyright 2019 Google LLC -// -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file or at -// https://developers.google.com/open-source/licenses/bsd - -// Package curve25519 implements the new proposed API for -// golang.org/x/crypto/curve25519 from golang.org/issue/32670. -package curve25519 - -import ( - "crypto/subtle" - "fmt" - - "golang.org/x/crypto/curve25519" -) - -const ( - // ScalarSize is the size of the scalar input to X25519. - ScalarSize = 32 - // PointSize is the size of the point input to X25519. - PointSize = 32 -) - -// Basepoint is the canonical Curve25519 generator. -var Basepoint []byte - -func init() { - Basepoint = []byte{ - 0x09, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - } -} - -func checkBasepoint() { - if subtle.ConstantTimeCompare(Basepoint, []byte{ - 0x09, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, - }) != 1 { - panic("curve25519: global Basepoint value was modified") - } -} - -// X25519 returns the result of the scalar multiplication (scalar * point), -// according to RFC 7748, Section 5. scalar, point and the return value are -// slices of 32 bytes. -// -// scalar can be egenrated at random, for example with crypto/rand. point should -// be either Basepoint or the output of an X25519 call. -// -// If point is Basepoint (but not if it's a different slice with the same -// contents) a precomputed implementation might be used for performance. -func X25519(scalar, point []byte) ([]byte, error) { - // Outline the body of function, to let the allocation be inlined in the - // caller, and possibly avoid escaping to the heap. - var dst [32]byte - return x25519(&dst, scalar, point) -} - -func x25519(dst *[32]byte, scalar, point []byte) ([]byte, error) { - var in [32]byte - if l := len(scalar); l != 32 { - return nil, fmt.Errorf("bad scalar length: %d, expected %d", l, 32) - } - if l := len(point); l != 32 { - return nil, fmt.Errorf("bad point length: %d, expected %d", l, 32) - } - copy(in[:], scalar) - if &point[0] == &Basepoint[0] { - checkBasepoint() - curve25519.ScalarBaseMult(dst, &in) - } else { - var base, zero [32]byte - copy(base[:], point) - curve25519.ScalarMult(dst, &in, &base) - if subtle.ConstantTimeCompare(dst[:], zero[:]) == 1 { - // TODO: test this codepath with all low order points. - return nil, fmt.Errorf("bad input point: low order point") - } - } - return dst[:], nil -}