Files
brianmcgee 353d720ccc pkg/crypto: change machine name to just name
In future peer connections could be coming from CLI clients and not just machines.
2026-03-30 11:48:32 +01:00

513 lines
14 KiB
Go

package tls_test
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/x509"
"crypto/x509/pkix"
"encoding/asn1"
"fmt"
"math/big"
"net"
"testing"
"time"
dmcrypto "git.clan.lol/clan/data-mesher/pkg/crypto"
"git.clan.lol/clan/data-mesher/pkg/libp2p/security/tls"
"git.clan.lol/clan/data-mesher/test"
ic "github.com/libp2p/go-libp2p/core/crypto"
"github.com/libp2p/go-libp2p/core/peerstore"
"github.com/libp2p/go-libp2p/p2p/host/peerstore/pstoremem"
tptu "github.com/libp2p/go-libp2p/p2p/net/upgrader"
"github.com/stretchr/testify/require"
"golang.org/x/sync/errgroup"
)
// createTransport builds a tls.Transport from a private key and cert bundle.
func createTransport(
t *testing.T,
key *dmcrypto.PrivateKey,
bundle test.CertBundle,
) (*tls.Transport, peerstore.Peerstore) {
t.Helper()
p2pKey, err := key.ToLibP2P()
require.NoError(t, err)
ps, err := pstoremem.NewPeerstore()
require.NoError(t, err)
constructor := tls.New(bundle.CAKeys, bundle.Cert, ps)
transport, err := constructor(tls.ID, p2pKey, []tptu.StreamMuxer{})
require.NoError(t, err)
return transport, ps
}
// runHandshake runs a TLS handshake between two transports over an in-memory pipe, returning both secure connections.
func runHandshake(
t *testing.T,
clientTransport, serverTransport *tls.Transport,
clientKey, serverKey *dmcrypto.PrivateKey,
errCheck func(err error),
) {
t.Helper()
// default to no error check
if errCheck == nil {
errCheck = func(err error) {
require.NoError(t, err)
}
}
// create an in-memory pipe
clientConn, serverConn := net.Pipe()
// get peer IDs for both peers
serverPeerID, err := serverKey.PeerID()
require.NoError(t, err)
clientPeerID, err := clientKey.PeerID()
require.NoError(t, err)
// ensure a 10 second timeout
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
// run the handshake in parallel
eg, ctx := errgroup.WithContext(ctx)
eg.Go(func() error {
conn, connErr := serverTransport.SecureInbound(ctx, serverConn, clientPeerID)
if connErr == nil {
_ = conn.Close()
}
if connErr != nil {
return fmt.Errorf("secure inbound: %w", connErr)
}
return nil
})
eg.Go(func() error {
conn, connErr := clientTransport.SecureOutbound(ctx, clientConn, serverPeerID)
if connErr == nil {
_ = conn.Close()
}
if connErr != nil {
return fmt.Errorf("secure outbound: %w", connErr)
}
return nil
})
errCheck(eg.Wait())
}
func TestHandshake_ValidCerts(t *testing.T) {
t.Parallel()
as := require.New(t)
// generate some keys
clientKey, _ := test.GenerateIdentityKey(t)
serverKey, _ := test.GenerateIdentityKey(t)
// use the same CA for both peers
caKey, err := dmcrypto.GenerateKey(rand.Reader)
as.NoError(err)
// sign a cert for both peers
now := time.Now()
clientCert, err := dmcrypto.SignCertificate(
"client-machine", clientKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
serverCert, err := dmcrypto.SignCertificate(
"server-machine", serverKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
// create a transport for each peer
caKeys := []*dmcrypto.PublicKey{caKey.Public}
clientBundle := test.CertBundle{Cert: clientCert, CAKeys: caKeys}
serverBundle := test.CertBundle{Cert: serverCert, CAKeys: caKeys}
clientTransport, clientPS := createTransport(t, clientKey, clientBundle)
serverTransport, serverPS := createTransport(t, serverKey, serverBundle)
// run the handshake
runHandshake(t, clientTransport, serverTransport, clientKey, serverKey, nil)
// verify peerstore recorded the correct network (CA key) for each remote peer
serverPeerID, err := serverKey.PeerID()
as.NoError(err)
clientPeerID, err := clientKey.PeerID()
as.NoError(err)
// client's peerstore records the server's network (CA key)
val, err := clientPS.Get(serverPeerID, tls.PeerMetaNetwork)
as.NoError(err)
as.Equal(caKey.Public, val)
// server's peerstore records the client's network (CA key)
val, err = serverPS.Get(clientPeerID, tls.PeerMetaNetwork)
as.NoError(err)
as.Equal(caKey.Public, val)
// verify peerstore recorded the correct machine names
clientMachineName, err := tls.GetPeerName(clientPS, serverPeerID)
as.NoError(err)
as.Equal("server-machine", clientMachineName)
serverMachineName, err := tls.GetPeerName(serverPS, clientPeerID)
as.NoError(err)
as.Equal("client-machine", serverMachineName)
}
func TestHandshake_NameStorage(t *testing.T) {
t.Parallel()
as := require.New(t)
// generate keys for two distinct machines
aliceKey, _ := test.GenerateIdentityKey(t)
bobKey, _ := test.GenerateIdentityKey(t)
// use the same CA for both peers
caKey, err := dmcrypto.GenerateKey(rand.Reader)
as.NoError(err)
// sign certs with different peer names
now := time.Now()
aliceCert, err := dmcrypto.SignCertificate(
"alice-workstation", aliceKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
bobCert, err := dmcrypto.SignCertificate(
"bob-laptop", bobKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
// create transports for each peer
caKeys := []*dmcrypto.PublicKey{caKey.Public}
aliceBundle := test.CertBundle{Cert: aliceCert, CAKeys: caKeys}
bobBundle := test.CertBundle{Cert: bobCert, CAKeys: caKeys}
aliceTransport, alicePS := createTransport(t, aliceKey, aliceBundle)
bobTransport, bobPS := createTransport(t, bobKey, bobBundle)
// run the handshake
runHandshake(t, aliceTransport, bobTransport, aliceKey, bobKey, nil)
// verify each peer's peerstore contains the other's peer name
alicePeerID, err := aliceKey.PeerID()
as.NoError(err)
bobPeerID, err := bobKey.PeerID()
as.NoError(err)
// alice's peerstore should have bob's peer name
bobMachineName, err := tls.GetPeerName(alicePS, bobPeerID)
as.NoError(err)
as.Equal("bob-laptop", bobMachineName)
// bob's peerstore should have alice's peer name
aliceMachineName, err := tls.GetPeerName(bobPS, alicePeerID)
as.NoError(err)
as.Equal("alice-workstation", aliceMachineName)
}
func TestHandshake_ExpiredCert(t *testing.T) {
t.Parallel()
as := require.New(t)
// generate some keys
clientKey, _ := test.GenerateIdentityKey(t)
serverKey, _ := test.GenerateIdentityKey(t)
// use the same CA for both peers
caKey, err := dmcrypto.GenerateKey(rand.Reader)
as.NoError(err)
now := time.Now()
// client has a valid cert
clientCert, err := dmcrypto.SignCertificate(
"client-machine", clientKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
// server has an expired cert
serverCert, err := dmcrypto.SignCertificate(
"server-machine", serverKey.Public, caKey, now.Add(-2*time.Hour), now.Add(-time.Hour),
)
as.NoError(err)
// create a transport for each peer
caKeys := []*dmcrypto.PublicKey{caKey.Public}
clientBundle := test.CertBundle{Cert: clientCert, CAKeys: caKeys}
serverBundle := test.CertBundle{Cert: serverCert, CAKeys: caKeys}
clientTransport, _ := createTransport(t, clientKey, clientBundle)
serverTransport, _ := createTransport(t, serverKey, serverBundle)
// the handshake should fail because the server's cert is expired
errCheck := func(err error) {
as.Error(err)
}
runHandshake(t, clientTransport, serverTransport, clientKey, serverKey, errCheck)
}
func TestHandshake_UntrustedCA(t *testing.T) {
t.Parallel()
as := require.New(t)
// generate some keys
clientKey, _ := test.GenerateIdentityKey(t)
serverKey, _ := test.GenerateIdentityKey(t)
caKeyOne, err := dmcrypto.GenerateKey(rand.Reader)
as.NoError(err)
caKeyTwo, err := dmcrypto.GenerateKey(rand.Reader)
as.NoError(err)
now := time.Now()
// client trusts CA One, server's cert is signed by CA Two
clientCert, err := dmcrypto.SignCertificate(
"client-machine", clientKey.Public, caKeyOne, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
serverCert, err := dmcrypto.SignCertificate(
"server-machine", serverKey.Public, caKeyTwo, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
// client only trusts CA One
clientBundle := test.CertBundle{Cert: clientCert, CAKeys: []*dmcrypto.PublicKey{caKeyOne.Public}}
// server only trusts CA Two
serverBundle := test.CertBundle{Cert: serverCert, CAKeys: []*dmcrypto.PublicKey{caKeyTwo.Public}}
// create transports for each peer
clientTransport, _ := createTransport(t, clientKey, clientBundle)
serverTransport, _ := createTransport(t, serverKey, serverBundle)
errCheck := func(err error) {
as.Error(err)
}
runHandshake(t, clientTransport, serverTransport, clientKey, serverKey, errCheck)
}
func TestHandshake_PeerIDMismatch(t *testing.T) {
t.Parallel()
as := require.New(t)
// generate some keys
clientKey, _ := test.GenerateIdentityKey(t)
serverKey, _ := test.GenerateIdentityKey(t)
wrongKey, _ := test.GenerateIdentityKey(t)
// use the same CA for both peers
caKey, err := dmcrypto.GenerateKey(rand.Reader)
as.NoError(err)
now := time.Now()
// sign a cert for both peers
clientCert, err := dmcrypto.SignCertificate(
"client-machine", clientKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
serverCert, err := dmcrypto.SignCertificate(
"server-machine", serverKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
as.NoError(err)
// create a transport for each peer
caKeys := []*dmcrypto.PublicKey{caKey.Public}
clientBundle := test.CertBundle{Cert: clientCert, CAKeys: caKeys}
serverBundle := test.CertBundle{Cert: serverCert, CAKeys: caKeys}
clientTransport, _ := createTransport(t, clientKey, clientBundle)
serverTransport, _ := createTransport(t, serverKey, serverBundle)
// create an in-memory pipe
clientConn, serverConn := net.Pipe()
// client expects a different peer ID than the server actually has
wrongPeerID, err := wrongKey.PeerID()
require.NoError(t, err)
clientPeerID, err := clientKey.PeerID()
require.NoError(t, err)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
eg, ctx := errgroup.WithContext(ctx)
eg.Go(func() error {
conn, connErr := serverTransport.SecureInbound(ctx, serverConn, clientPeerID)
if connErr == nil {
_ = conn.Close()
}
if connErr != nil {
return fmt.Errorf("secure inbound: %w", connErr)
}
return nil
})
eg.Go(func() error {
// client dials expecting wrongPeerID
conn, connErr := clientTransport.SecureOutbound(ctx, clientConn, wrongPeerID)
if connErr == nil {
_ = conn.Close()
}
if connErr != nil {
return fmt.Errorf("secure outbound: %w", connErr)
}
return nil
})
err = eg.Wait()
as.Error(err)
}
func TestHandshake_NoDMCertExtension(t *testing.T) {
t.Parallel()
// test VerifyCACert directly: create an x509 cert with only the libp2p key extension but NO data-mesher cert
// extension, then verify it is rejected.
key, _ := test.GenerateIdentityKey(t)
caKey, err := dmcrypto.GenerateKey(rand.Reader)
require.NoError(t, err)
p2pKey, err := key.ToLibP2P()
require.NoError(t, err)
chain := createBareX509Cert(t, p2pKey)
_, _, err = tls.VerifyCACert(chain, []*dmcrypto.PublicKey{caKey.Public}, p2pKey.GetPublic())
require.Error(t, err)
require.Contains(t, err.Error(), "data-mesher certificate extension")
}
func TestVerifyCACert_CertKeyMismatch(t *testing.T) {
t.Parallel()
// data-mesher cert has a different public key than the libp2p identity
nodeKey, _ := test.GenerateIdentityKey(t)
otherKey, _ := test.GenerateIdentityKey(t)
caKey, err := dmcrypto.GenerateKey(rand.Reader)
require.NoError(t, err)
now := time.Now()
// sign a cert for otherKey, but present it with nodeKey's libp2p identity
dmCert, err := dmcrypto.SignCertificate(
"other-machine", otherKey.Public, caKey, now.Add(-time.Minute), now.Add(time.Hour),
)
require.NoError(t, err)
p2pKey, err := nodeKey.ToLibP2P()
require.NoError(t, err)
// create x509 cert that embeds nodeKey's libp2p identity but otherKey's DM cert
x509Cert := createX509CertWithDMCert(t, p2pKey, dmCert)
_, _, err = tls.VerifyCACert(x509Cert, []*dmcrypto.PublicKey{caKey.Public}, p2pKey.GetPublic())
require.Error(t, err)
require.Contains(t, err.Error(), "does not match")
}
// createBareX509Cert creates a self-signed x509 certificate with only the libp2p key extension (no data-mesher cert
// extension).
func createBareX509Cert(t *testing.T, privKey ic.PrivKey) []*x509.Certificate {
t.Helper()
certKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)
ext, err := tls.GenerateSignedExtension(privKey, certKey.Public())
require.NoError(t, err)
tmpl := &x509.Certificate{
SerialNumber: big.NewInt(1),
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
ExtraExtensions: []pkix.Extension{ext},
}
certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, certKey.Public(), certKey)
require.NoError(t, err)
parsed, err := x509.ParseCertificate(certDER)
require.NoError(t, err)
return []*x509.Certificate{parsed}
}
// createX509CertWithDMCert creates a self-signed x509 certificate with both the libp2p key extension and a data-mesher
// cert extension containing the provided data-mesher certificate.
func createX509CertWithDMCert(t *testing.T, privKey ic.PrivKey, dmCert *dmcrypto.Certificate) []*x509.Certificate {
t.Helper()
certKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)
ext, err := tls.GenerateSignedExtension(privKey, certKey.Public())
require.NoError(t, err)
dmExt := pkix.Extension{
Id: dmCertExtensionOID(),
Value: dmCert.Bytes(),
}
tmpl := &x509.Certificate{
SerialNumber: big.NewInt(1),
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
ExtraExtensions: []pkix.Extension{ext, dmExt},
}
certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, certKey.Public(), certKey)
require.NoError(t, err)
parsed, err := x509.ParseCertificate(certDER)
require.NoError(t, err)
return []*x509.Certificate{parsed}
}
// dmCertExtensionOID returns the OID for the data-mesher cert extension.
func dmCertExtensionOID() asn1.ObjectIdentifier {
return asn1.ObjectIdentifier{1, 3, 6, 1, 4, 1, 53594, 1, 2}
}