Files

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)
})
}
}