551 lines
14 KiB
Go
551 lines
14 KiB
Go
package crypto_test
|
|
|
|
import (
|
|
"crypto/rand"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.clan.lol/clan/data-mesher/pkg/crypto"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
func TestCertificate_SignAndVerify(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA keypair
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// generate machine keypair
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// sign the machine's public key with the CA (valid for 1 hour)
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
as.NotNil(cert)
|
|
|
|
// verify the certificate
|
|
err = cert.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
|
|
// certificate's public key should match the machine's public key
|
|
as.True(cert.IdentityKey.Equal(machineKey.Public))
|
|
|
|
// verify timestamps
|
|
as.Equal(notBefore, cert.NotBefore)
|
|
as.Equal(notAfter, cert.NotAfter)
|
|
}
|
|
|
|
func TestCertificate_VerifyWithWrongCA(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate two Networks
|
|
networkOne, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
networkTwo, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// generate machine keypair
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// sign with CA One
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, networkOne, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// verify with CA Two should fail
|
|
err = cert.Verify(networkTwo.Public)
|
|
as.ErrorIs(err, crypto.ErrInvalidSignature)
|
|
as.Contains(err.Error(), "signature verification failed")
|
|
}
|
|
|
|
func TestCertificate_ParseAndString(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(24 * time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// serialize to string
|
|
certStr := cert.String()
|
|
as.NotEmpty(certStr)
|
|
|
|
// parse back
|
|
parsed, err := crypto.ParseCertificate(certStr)
|
|
as.NoError(err)
|
|
|
|
// should be identical
|
|
as.True(parsed.IdentityKey.Equal(cert.IdentityKey))
|
|
as.Equal(cert.Name, parsed.Name)
|
|
|
|
// note: timestamps are unix seconds when serialised
|
|
as.Equal(cert.NotBefore.Truncate(time.Second), parsed.NotBefore)
|
|
as.Equal(cert.NotAfter.Truncate(time.Second), parsed.NotAfter)
|
|
|
|
as.Equal(cert.Signature, parsed.Signature)
|
|
|
|
// should still verify
|
|
err = parsed.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
}
|
|
|
|
func TestCertificate_Bytes(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// get bytes
|
|
certBytes := cert.Bytes()
|
|
// certificate size is now fixed at 176 bytes
|
|
as.Len(certBytes, crypto.CertificateSize) // 32 + 8 + 8 + 64 + 64 = 176 bytes
|
|
as.Equal(176, crypto.CertificateSize)
|
|
}
|
|
|
|
func TestCertificate_NilInputs(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
networkKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
identityKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
// nil network key
|
|
_, err = crypto.SignCertificate("test-machine", identityKey.Public, nil, notBefore, notAfter)
|
|
as.Error(err)
|
|
as.ErrorContains(err, "network key is required")
|
|
|
|
// nil identity key
|
|
_, err = crypto.SignCertificate("test-machine", nil, networkKey, notBefore, notAfter)
|
|
as.Error(err)
|
|
as.ErrorContains(err, "identity key is required")
|
|
|
|
// empty name
|
|
_, err = crypto.SignCertificate("", identityKey.Public, networkKey, notBefore, notAfter)
|
|
as.Error(err)
|
|
as.ErrorContains(err, "name is required")
|
|
|
|
// name too long (> 64 bytes)
|
|
longName := string(make([]byte, 65)) // 65 bytes
|
|
_, err = crypto.SignCertificate(longName, identityKey.Public, networkKey, notBefore, notAfter)
|
|
as.Error(err)
|
|
as.ErrorContains(err, "exceeds maximum length of 64 bytes")
|
|
}
|
|
|
|
func TestCertificate_Expiry(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create an expired certificate (expired 1 hour ago)
|
|
notBefore := time.Now().Add(-2 * time.Hour)
|
|
notAfter := time.Now().Add(-1 * time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// full verification should fail due to expiry
|
|
err = cert.Verify(caKey.Public)
|
|
as.Error(err)
|
|
as.ErrorIs(err, crypto.ErrInvalidCertificate)
|
|
as.Contains(err.Error(), "certificate has expired")
|
|
|
|
// signature-only verification should still pass
|
|
err = cert.VerifySignatureOnly(caKey.Public)
|
|
as.NoError(err)
|
|
|
|
// IsExpired should return true
|
|
as.True(cert.IsExpired())
|
|
|
|
// IsValid should return false
|
|
as.False(cert.IsValid())
|
|
}
|
|
|
|
func TestCertificate_NotYetValid(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate that's not valid yet (valid in 1 hour)
|
|
notBefore := time.Now().Add(1 * time.Hour)
|
|
notAfter := time.Now().Add(2 * time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// full verification should fail
|
|
err = cert.Verify(caKey.Public)
|
|
as.Error(err)
|
|
as.ErrorIs(err, crypto.ErrInvalidCertificate)
|
|
as.Contains(err.Error(), "not yet valid")
|
|
|
|
// signature-only verification should still pass
|
|
err = cert.VerifySignatureOnly(caKey.Public)
|
|
as.NoError(err)
|
|
|
|
// IsExpired should return false (not expired, just not valid yet)
|
|
as.False(cert.IsExpired())
|
|
|
|
// IsValid should return false
|
|
as.False(cert.IsValid())
|
|
}
|
|
|
|
func TestCertificate_InvalidTimeRange(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate with notAfter before notBefore
|
|
notBefore := time.Now().Add(2 * time.Hour)
|
|
notAfter := time.Now().Add(1 * time.Hour)
|
|
|
|
_, err = crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.Error(err)
|
|
as.Contains(err.Error(), "notAfter must be after notBefore")
|
|
}
|
|
|
|
func TestCertificate_TimestampIntegrity(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a valid certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(24 * time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// serialize and parse back
|
|
certStr := cert.String()
|
|
parsed, err := crypto.ParseCertificate(certStr)
|
|
as.NoError(err)
|
|
|
|
// tamper with the notAfter timestamp (extend validity by 1 year)
|
|
parsed.NotAfter = parsed.NotAfter.Add(365 * 24 * time.Hour)
|
|
|
|
// signature verification should fail because timestamps are part of signed data
|
|
err = parsed.VerifySignatureOnly(caKey.Public)
|
|
as.Error(err)
|
|
as.ErrorIs(err, crypto.ErrInvalidSignature)
|
|
as.Contains(err.Error(), "signature verification failed")
|
|
}
|
|
|
|
func TestCertificate_VerifySignatureOnly(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a valid certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// signature-only verification should pass
|
|
err = cert.VerifySignatureOnly(caKey.Public)
|
|
as.NoError(err)
|
|
|
|
// test nil cases
|
|
err = cert.VerifySignatureOnly(nil)
|
|
as.Error(err)
|
|
as.Contains(err.Error(), "ca public key is required")
|
|
}
|
|
|
|
func TestCertificate_PEM(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(24 * time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// get PEM representation
|
|
pemData := cert.PEM()
|
|
as.NotEmpty(pemData)
|
|
|
|
// verify PEM format
|
|
pemStr := string(pemData)
|
|
as.Contains(pemStr, "-----BEGIN DATA-MESHER CERTIFICATE-----")
|
|
as.Contains(pemStr, "-----END DATA-MESHER CERTIFICATE-----")
|
|
|
|
// parse back from PEM
|
|
parsed, err := crypto.ParseCertificate(pemStr)
|
|
as.NoError(err)
|
|
|
|
// should be identical
|
|
as.True(parsed.IdentityKey.Equal(cert.IdentityKey))
|
|
as.Equal(cert.Name, parsed.Name)
|
|
|
|
// note: timestamps are unix seconds when serialised
|
|
as.Equal(cert.NotBefore.Truncate(time.Second), parsed.NotBefore)
|
|
as.Equal(cert.NotAfter.Truncate(time.Second), parsed.NotAfter)
|
|
|
|
as.Equal(cert.Signature, parsed.Signature)
|
|
|
|
// should still verify
|
|
err = parsed.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
}
|
|
|
|
func TestCertificate_ParsePEMWithoutHeaders(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keypairs
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(24 * time.Hour)
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// get base64 string (without PEM headers)
|
|
base64Str := cert.String()
|
|
|
|
// add some whitespace/newlines to simulate copy-paste
|
|
dataWithWhitespace := "\n " + base64Str + " \n"
|
|
|
|
// should still parse correctly
|
|
parsed, err := crypto.ParseCertificate(dataWithWhitespace)
|
|
as.NoError(err)
|
|
as.True(parsed.IdentityKey.Equal(cert.IdentityKey))
|
|
|
|
// should still verify
|
|
err = parsed.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
}
|
|
|
|
func TestCertificate_PeerID(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
// create a certificate
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("test-machine", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// derive peer ID from certificate
|
|
certPeerID, err := cert.PeerID()
|
|
as.NoError(err)
|
|
|
|
// derive peer ID from the machine's private key directly
|
|
machinePeerID, err := machineKey.PeerID()
|
|
as.NoError(err)
|
|
|
|
// they should match
|
|
as.Equal(machinePeerID, certPeerID)
|
|
}
|
|
|
|
func TestCertificate_MarshalBinaryRoundTrip(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("binary-test", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
data, err := cert.MarshalBinary()
|
|
as.NoError(err)
|
|
as.Len(data, crypto.CertificateSize)
|
|
|
|
var parsed crypto.Certificate
|
|
|
|
err = parsed.UnmarshalBinary(data)
|
|
as.NoError(err)
|
|
|
|
as.True(parsed.IdentityKey.Equal(cert.IdentityKey))
|
|
as.Equal(cert.Name, parsed.Name)
|
|
as.Equal(cert.NotBefore.Truncate(time.Second), parsed.NotBefore)
|
|
as.Equal(cert.NotAfter.Truncate(time.Second), parsed.NotAfter)
|
|
as.Equal(cert.Signature, parsed.Signature)
|
|
|
|
err = parsed.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
}
|
|
|
|
func TestCertificate_MarshalJSONRoundTrip(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
cert, err := crypto.SignCertificate("json-test", machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
jsonData, err := cert.MarshalJSON()
|
|
as.NoError(err)
|
|
// should be a quoted base64 string
|
|
as.Equal(byte('"'), jsonData[0])
|
|
as.Equal(byte('"'), jsonData[len(jsonData)-1])
|
|
|
|
var parsed crypto.Certificate
|
|
|
|
err = parsed.UnmarshalJSON(jsonData)
|
|
as.NoError(err)
|
|
|
|
as.True(parsed.IdentityKey.Equal(cert.IdentityKey))
|
|
as.Equal(cert.Name, parsed.Name)
|
|
as.Equal(cert.NotBefore.Truncate(time.Second), parsed.NotBefore)
|
|
as.Equal(cert.NotAfter.Truncate(time.Second), parsed.NotAfter)
|
|
as.Equal(cert.Signature, parsed.Signature)
|
|
|
|
err = parsed.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
}
|
|
|
|
func TestCertificate_MachineNameLengths(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// generate CA and machine keys
|
|
caKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
machineKey, err := crypto.GenerateKey(rand.Reader)
|
|
as.NoError(err)
|
|
|
|
notBefore := time.Now()
|
|
notAfter := notBefore.Add(time.Hour)
|
|
|
|
testCases := []struct {
|
|
name string
|
|
machineName string
|
|
}{
|
|
{"short name", "a"},
|
|
{"medium name", "my-test-machine"},
|
|
{"max length name", "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ01"}, // exactly 64 bytes
|
|
}
|
|
|
|
for _, tc := range testCases {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
as := require.New(t)
|
|
|
|
// create certificate
|
|
cert, err := crypto.SignCertificate(tc.machineName, machineKey.Public, caKey, notBefore, notAfter)
|
|
as.NoError(err)
|
|
|
|
// serialize to bytes and back
|
|
certBytes := cert.Bytes()
|
|
parsed, err := crypto.ParseCertificateBytes(certBytes)
|
|
as.NoError(err)
|
|
|
|
// machine name should be preserved
|
|
as.Equal(tc.machineName, parsed.Name)
|
|
|
|
// should verify correctly
|
|
err = parsed.Verify(caKey.Public)
|
|
as.NoError(err)
|
|
})
|
|
}
|
|
}
|