Refactoring crypto code for future reuse. (#25148)
Refactoring crypto code for future reuse for #24869. No functional changes.
This commit is contained in:
@@ -25,8 +25,8 @@ import (
|
||||
"github.com/fleetdm/fleet/v4/server/fleet"
|
||||
apple_mdm "github.com/fleetdm/fleet/v4/server/mdm/apple"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/nanomdm/mdm"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/cryptoutil/x509util"
|
||||
scepserver "github.com/fleetdm/fleet/v4/server/mdm/scep/server"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/x509util"
|
||||
httptransport "github.com/go-kit/kit/transport/http"
|
||||
"github.com/go-kit/log"
|
||||
kitlog "github.com/go-kit/log"
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
package cryptoutil
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/ed25519"
|
||||
"crypto/elliptic"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"crypto/x509"
|
||||
"encoding/asn1"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// GenerateSubjectKeyID generates Subject Key Identifier (SKI) using SHA-256
|
||||
// hash of the public key bytes according to RFC 7093 section 2.
|
||||
func GenerateSubjectKeyID(pub crypto.PublicKey) ([]byte, error) {
|
||||
var pubBytes []byte
|
||||
var err error
|
||||
switch pub := pub.(type) {
|
||||
case *rsa.PublicKey:
|
||||
pubBytes, err = asn1.Marshal(*pub)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case *ecdsa.PublicKey:
|
||||
pubBytes = elliptic.Marshal(pub.Curve, pub.X, pub.Y)
|
||||
default:
|
||||
return nil, errors.New("only ECDSA and RSA public keys are supported")
|
||||
}
|
||||
|
||||
hash := sha256.Sum256(pubBytes)
|
||||
|
||||
// According to RFC 7093, The keyIdentifier is composed of the leftmost
|
||||
// 160-bits of the SHA-256 hash of the value of the BIT STRING
|
||||
// subjectPublicKey (excluding the tag, length, and number of unused bits).
|
||||
return hash[:20], nil
|
||||
}
|
||||
|
||||
// ParsePrivateKey parses a PEM encoded private key and returns a crypto.PrivateKey.
|
||||
// It can be used for private keys passed in from environment variables or command line or files.
|
||||
func ParsePrivateKey(privKeyPEM []byte, keyName string) (crypto.PrivateKey, error) {
|
||||
block, _ := pem.Decode(privKeyPEM)
|
||||
if block == nil {
|
||||
return nil, fmt.Errorf("failed to decode %s", keyName)
|
||||
}
|
||||
|
||||
// The code below is based on tls.parsePrivateKey
|
||||
// https://cs.opensource.google/go/go/+/release-branch.go1.23:src/crypto/tls/tls.go;l=355-372
|
||||
if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
|
||||
switch key := key.(type) {
|
||||
case *rsa.PrivateKey, *ecdsa.PrivateKey, ed25519.PrivateKey:
|
||||
return key, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unmarshaled PKCS8 %s is not an RSA, ECDSA, or Ed25519 private key", keyName)
|
||||
}
|
||||
}
|
||||
if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("failed to parse %s of type %s", keyName, block.Type)
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package cryptoutil
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"math/big"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGenerateSubjectKeyID(t *testing.T) {
|
||||
ecKey, err := ecdsa.GenerateKey(elliptic.P224(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, test := range []struct {
|
||||
testName string
|
||||
pub crypto.PublicKey
|
||||
}{
|
||||
{"RSA", &rsa.PublicKey{N: big.NewInt(123), E: 65537}},
|
||||
{"ECDSA", ecKey.Public()},
|
||||
} {
|
||||
test := test
|
||||
t.Run(test.testName, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ski, err := GenerateSubjectKeyID(test.pub)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(ski) != 20 {
|
||||
t.Fatalf("unexpected subject public key identifier length: %d", len(ski))
|
||||
}
|
||||
ski2, err := GenerateSubjectKeyID(test.pub)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !testSKIEq(ski, ski2) {
|
||||
t.Fatal("subject key identifier generation is not deterministic")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testSKIEq(a, b []byte) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
|
||||
for i := range a {
|
||||
if a[i] != b[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func TestParsePrivateKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
// nil block not allowed
|
||||
_, err := ParsePrivateKey(nil, "APNS private key")
|
||||
assert.ErrorContains(t, err, "failed to decode")
|
||||
|
||||
// encrypted pkcs8 not supported
|
||||
pkcs8Encrypted, err := os.ReadFile("testdata/pkcs8-encrypted.key")
|
||||
require.NoError(t, err)
|
||||
_, err = ParsePrivateKey(pkcs8Encrypted, "APNS private key")
|
||||
assert.ErrorContains(t, err, "failed to parse APNS private key of type ENCRYPTED PRIVATE KEY")
|
||||
|
||||
// X25519 pkcs8 not supported
|
||||
pkcs8Encrypted, err = os.ReadFile("testdata/pkcs8-x25519.key")
|
||||
require.NoError(t, err)
|
||||
_, err = ParsePrivateKey(pkcs8Encrypted, "APNS private key")
|
||||
assert.ErrorContains(t, err, "unmarshaled PKCS8 APNS private key is not")
|
||||
|
||||
// In this test, the pkcs1 key and pkcs8 keys are the same key, just different formats
|
||||
pkcs1, err := os.ReadFile("testdata/pkcs1.key")
|
||||
require.NoError(t, err)
|
||||
pkcs1Key, err := ParsePrivateKey(pkcs1, "APNS private key")
|
||||
require.NoError(t, err)
|
||||
|
||||
pkcs8, err := os.ReadFile("testdata/pkcs8-rsa.key")
|
||||
require.NoError(t, err)
|
||||
pkcs8Key, err := ParsePrivateKey(pkcs8, "APNS private key")
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, pkcs1Key, pkcs8Key)
|
||||
}
|
||||
Vendored
@@ -10,7 +10,7 @@ import (
|
||||
"io/ioutil"
|
||||
"os"
|
||||
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/cryptoutil/x509util"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/x509util"
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -1,36 +0,0 @@
|
||||
package cryptoutil
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"encoding/asn1"
|
||||
"errors"
|
||||
)
|
||||
|
||||
// GenerateSubjectKeyID generates Subject Key Identifier (SKI) using SHA-256
|
||||
// hash of the public key bytes according to RFC 7093 section 2.
|
||||
func GenerateSubjectKeyID(pub crypto.PublicKey) ([]byte, error) {
|
||||
var pubBytes []byte
|
||||
var err error
|
||||
switch pub := pub.(type) {
|
||||
case *rsa.PublicKey:
|
||||
pubBytes, err = asn1.Marshal(*pub)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case *ecdsa.PublicKey:
|
||||
pubBytes = elliptic.Marshal(pub.Curve, pub.X, pub.Y)
|
||||
default:
|
||||
return nil, errors.New("only ECDSA and RSA public keys are supported")
|
||||
}
|
||||
|
||||
hash := sha256.Sum256(pubBytes)
|
||||
|
||||
// According to RFC 7093, The keyIdentifier is composed of the leftmost
|
||||
// 160-bits of the SHA-256 hash of the value of the BIT STRING
|
||||
// subjectPublicKey (excluding the tag, length, and number of unused bits).
|
||||
return hash[:20], nil
|
||||
}
|
||||
@@ -1,58 +0,0 @@
|
||||
package cryptoutil
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"math/big"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGenerateSubjectKeyID(t *testing.T) {
|
||||
ecKey, err := ecdsa.GenerateKey(elliptic.P224(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, test := range []struct {
|
||||
testName string
|
||||
pub crypto.PublicKey
|
||||
}{
|
||||
{"RSA", &rsa.PublicKey{N: big.NewInt(123), E: 65537}},
|
||||
{"ECDSA", ecKey.Public()},
|
||||
} {
|
||||
test := test
|
||||
t.Run(test.testName, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ski, err := GenerateSubjectKeyID(test.pub)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(ski) != 20 {
|
||||
t.Fatalf("unexpected subject public key identifier length: %d", len(ski))
|
||||
}
|
||||
ski2, err := GenerateSubjectKeyID(test.pub)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !testSKIEq(ski, ski2) {
|
||||
t.Fatal("subject key identifier generation is not deterministic")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testSKIEq(a, b []byte) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
|
||||
for i := range a {
|
||||
if a[i] != b[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"math/big"
|
||||
"time"
|
||||
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/cryptoutil"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/cryptoutil"
|
||||
)
|
||||
|
||||
// CACert represents a new self-signed CA certificate
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"crypto/x509"
|
||||
"time"
|
||||
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/cryptoutil"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/cryptoutil"
|
||||
"github.com/smallstep/scep"
|
||||
)
|
||||
|
||||
|
||||
+3
-32
@@ -4,11 +4,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/ed25519"
|
||||
"crypto/rsa"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
@@ -34,6 +30,7 @@ import (
|
||||
"github.com/fleetdm/fleet/v4/server/mdm"
|
||||
apple_mdm "github.com/fleetdm/fleet/v4/server/mdm/apple"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/assets"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/cryptoutil"
|
||||
nanomdm "github.com/fleetdm/fleet/v4/server/mdm/nanomdm/mdm"
|
||||
"github.com/fleetdm/fleet/v4/server/ptr"
|
||||
"github.com/go-kit/log/level"
|
||||
@@ -2496,10 +2493,9 @@ func (svc *Service) GetMDMAppleCSR(ctx context.Context) ([]byte, error) {
|
||||
}
|
||||
} else {
|
||||
rawApnsKey := savedAssets[fleet.MDMAssetAPNSKey]
|
||||
block, _ := pem.Decode(rawApnsKey.Value)
|
||||
apnsKey, err = parseAPNSPrivateKey(ctx, block)
|
||||
apnsKey, err = cryptoutil.ParsePrivateKey(rawApnsKey.Value, "APNS private key")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, ctxerr.Wrap(ctx, err, "parse APNS private key")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2546,31 +2542,6 @@ func (svc *Service) GetMDMAppleCSR(ctx context.Context) ([]byte, error) {
|
||||
return signedCSRB64, nil
|
||||
}
|
||||
|
||||
func parseAPNSPrivateKey(ctx context.Context, block *pem.Block) (crypto.PrivateKey, error) {
|
||||
if block == nil {
|
||||
return nil, ctxerr.New(ctx, "failed to decode saved APNS key")
|
||||
}
|
||||
|
||||
// The code below is based on tls.parsePrivateKey
|
||||
// https://cs.opensource.google/go/go/+/release-branch.go1.23:src/crypto/tls/tls.go;l=355-372
|
||||
if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
|
||||
switch key := key.(type) {
|
||||
case *rsa.PrivateKey, *ecdsa.PrivateKey, ed25519.PrivateKey:
|
||||
return key, nil
|
||||
default:
|
||||
return nil, errors.New("unmarshaled PKCS8 APNS key is not an RSA, ECDSA, or Ed25519 private key")
|
||||
}
|
||||
}
|
||||
if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
|
||||
return nil, ctxerr.New(ctx, fmt.Sprintf("failed to parse APNS private key of type %s", block.Type))
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
// POST /mdm/apple/apns_certificate
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"database/sql"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"math/big"
|
||||
"net/http"
|
||||
@@ -29,7 +28,7 @@ import (
|
||||
"github.com/fleetdm/fleet/v4/server/contexts/license"
|
||||
"github.com/fleetdm/fleet/v4/server/contexts/viewer"
|
||||
"github.com/fleetdm/fleet/v4/server/fleet"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/cryptoutil/x509util"
|
||||
"github.com/fleetdm/fleet/v4/server/mdm/scep/x509util"
|
||||
"github.com/fleetdm/fleet/v4/server/mock"
|
||||
"github.com/fleetdm/fleet/v4/server/ptr"
|
||||
"github.com/fleetdm/fleet/v4/server/test"
|
||||
@@ -2185,44 +2184,3 @@ func TestBatchSetMDMProfilesLabels(t *testing.T) {
|
||||
assert.Equal(t, ProfileLabels{IncludeAny: true}, *profileLabels["DIncAny"])
|
||||
assert.Equal(t, ProfileLabels{ExcludeAny: true}, *profileLabels["DExclAny"])
|
||||
}
|
||||
|
||||
func TestParseAPNSPrivateKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
// nil block not allowed
|
||||
ctx := context.Background()
|
||||
_, err := parseAPNSPrivateKey(ctx, nil)
|
||||
assert.ErrorContains(t, err, "failed to decode")
|
||||
|
||||
// encrypted pkcs8 not supported
|
||||
pkcs8Encrypted, err := os.ReadFile("testdata/pkcs8-encrypted.key")
|
||||
require.NoError(t, err)
|
||||
block, _ := pem.Decode(pkcs8Encrypted)
|
||||
assert.NotNil(t, block)
|
||||
_, err = parseAPNSPrivateKey(ctx, block)
|
||||
assert.ErrorContains(t, err, "failed to parse APNS private key of type ENCRYPTED PRIVATE KEY")
|
||||
|
||||
// X25519 pkcs8 not supported
|
||||
pkcs8Encrypted, err = os.ReadFile("testdata/pkcs8-x25519.key")
|
||||
require.NoError(t, err)
|
||||
block, _ = pem.Decode(pkcs8Encrypted)
|
||||
assert.NotNil(t, block)
|
||||
_, err = parseAPNSPrivateKey(ctx, block)
|
||||
assert.ErrorContains(t, err, "unmarshaled PKCS8 APNS key is not")
|
||||
|
||||
// In this test, the pkcs1 key and pkcs8 keys are the same key, just different formats
|
||||
pkcs1, err := os.ReadFile("testdata/pkcs1.key")
|
||||
require.NoError(t, err)
|
||||
block, _ = pem.Decode(pkcs1)
|
||||
assert.NotNil(t, block)
|
||||
pkcs1Key, err := parseAPNSPrivateKey(ctx, block)
|
||||
require.NoError(t, err)
|
||||
|
||||
pkcs8, err := os.ReadFile("testdata/pkcs8-rsa.key")
|
||||
require.NoError(t, err)
|
||||
block, _ = pem.Decode(pkcs8)
|
||||
assert.NotNil(t, block)
|
||||
pkcs8Key, err := parseAPNSPrivateKey(ctx, block)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, pkcs1Key, pkcs8Key)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user