In future peer connections could be coming from CLI clients and not just machines.
513 lines
14 KiB
Go
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}
|
|
}
|