Feat v0.6.7 native mtls (#42)
All checks were successful
All checks were successful
Co-authored-by: Jeffrey Smith <jasafpro@gmail.com> Co-committed-by: Jeffrey Smith <jasafpro@gmail.com>
This commit was merged in pull request #42.
This commit is contained in:
410
server/auth/mtls_native_test.go
Normal file
410
server/auth/mtls_native_test.go
Normal file
@@ -0,0 +1,410 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"switchboard-core/store"
|
||||
)
|
||||
|
||||
// ── Test CA helpers ──────────────────────────────────────────────────
|
||||
|
||||
type testCA struct {
|
||||
Cert *x509.Certificate
|
||||
Key *ecdsa.PrivateKey
|
||||
CertPEM []byte
|
||||
Pool *x509.CertPool
|
||||
}
|
||||
|
||||
func newTestCA(t *testing.T) *testCA {
|
||||
t.Helper()
|
||||
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatalf("generate CA key: %v", err)
|
||||
}
|
||||
|
||||
tmpl := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1),
|
||||
Subject: pkix.Name{CommonName: "Test CA"},
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
IsCA: true,
|
||||
BasicConstraintsValid: true,
|
||||
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign,
|
||||
}
|
||||
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
|
||||
if err != nil {
|
||||
t.Fatalf("create CA cert: %v", err)
|
||||
}
|
||||
|
||||
cert, err := x509.ParseCertificate(certDER)
|
||||
if err != nil {
|
||||
t.Fatalf("parse CA cert: %v", err)
|
||||
}
|
||||
|
||||
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(cert)
|
||||
|
||||
return &testCA{Cert: cert, Key: key, CertPEM: certPEM, Pool: pool}
|
||||
}
|
||||
|
||||
type testCert struct {
|
||||
Cert *x509.Certificate
|
||||
Key *ecdsa.PrivateKey
|
||||
TLSCert tls.Certificate
|
||||
}
|
||||
|
||||
func (ca *testCA) issueClient(t *testing.T, cn string, email string) *testCert {
|
||||
t.Helper()
|
||||
return ca.issueCert(t, cn, email, []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, time.Now().Add(time.Hour))
|
||||
}
|
||||
|
||||
func (ca *testCA) issueExpiredClient(t *testing.T, cn string) *testCert {
|
||||
t.Helper()
|
||||
return ca.issueCert(t, cn, "", []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, time.Now().Add(-time.Second))
|
||||
}
|
||||
|
||||
func (ca *testCA) issueNode(t *testing.T, cn string) *testCert {
|
||||
t.Helper()
|
||||
return ca.issueCert(t, cn, "", []x509.ExtKeyUsage{
|
||||
x509.ExtKeyUsageServerAuth,
|
||||
x509.ExtKeyUsageClientAuth,
|
||||
}, time.Now().Add(time.Hour))
|
||||
}
|
||||
|
||||
func (ca *testCA) issueCert(t *testing.T, cn, email string, eku []x509.ExtKeyUsage, notAfter time.Time) *testCert {
|
||||
t.Helper()
|
||||
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatalf("generate key: %v", err)
|
||||
}
|
||||
|
||||
tmpl := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(time.Now().UnixNano()),
|
||||
Subject: pkix.Name{CommonName: cn},
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: notAfter,
|
||||
ExtKeyUsage: eku,
|
||||
KeyUsage: x509.KeyUsageDigitalSignature,
|
||||
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)},
|
||||
DNSNames: []string{"localhost"},
|
||||
}
|
||||
if email != "" {
|
||||
tmpl.EmailAddresses = []string{email}
|
||||
}
|
||||
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, tmpl, ca.Cert, &key.PublicKey, ca.Key)
|
||||
if err != nil {
|
||||
t.Fatalf("create cert: %v", err)
|
||||
}
|
||||
|
||||
cert, err := x509.ParseCertificate(certDER)
|
||||
if err != nil {
|
||||
t.Fatalf("parse cert: %v", err)
|
||||
}
|
||||
|
||||
tlsCert := tls.Certificate{
|
||||
Certificate: [][]byte{certDER},
|
||||
PrivateKey: key,
|
||||
}
|
||||
|
||||
return &testCert{Cert: cert, Key: key, TLSCert: tlsCert}
|
||||
}
|
||||
|
||||
// ── FingerprintCert tests ───────────────────────────────────────────
|
||||
|
||||
func TestFingerprintCert(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
c1 := ca.issueClient(t, "alice", "")
|
||||
c2 := ca.issueClient(t, "alice", "") // same CN, different cert
|
||||
|
||||
fp1 := FingerprintCert(c1.Cert)
|
||||
fp2 := FingerprintCert(c2.Cert)
|
||||
|
||||
if fp1 == "" {
|
||||
t.Error("fingerprint should not be empty")
|
||||
}
|
||||
if len(fp1) != 64 { // sha256 hex = 64 chars
|
||||
t.Errorf("fingerprint length = %d, want 64", len(fp1))
|
||||
}
|
||||
if fp1 == fp2 {
|
||||
t.Error("different certs should have different fingerprints")
|
||||
}
|
||||
|
||||
// Same cert → same fingerprint (deterministic)
|
||||
if FingerprintCert(c1.Cert) != fp1 {
|
||||
t.Error("fingerprint should be deterministic")
|
||||
}
|
||||
}
|
||||
|
||||
// ── MTLSNativeProvider unit tests ───────────────────────────────────
|
||||
|
||||
func TestMTLSNativeProvider_Mode(t *testing.T) {
|
||||
p := NewMTLSNativeProvider(MTLSNativeConfig{})
|
||||
if p.Mode() != ModeMTLS {
|
||||
t.Errorf("Mode() = %q, want %q", p.Mode(), ModeMTLS)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMTLSNativeProvider_NoRegistration(t *testing.T) {
|
||||
p := NewMTLSNativeProvider(MTLSNativeConfig{})
|
||||
if p.SupportsRegistration() {
|
||||
t.Error("native mTLS should not support registration")
|
||||
}
|
||||
_, err := p.Register(nil, store.Stores{})
|
||||
if err != ErrNotSupported {
|
||||
t.Errorf("Register() = %v, want ErrNotSupported", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMTLSNativeProvider_NilTLS(t *testing.T) {
|
||||
p := NewMTLSNativeProvider(MTLSNativeConfig{})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||
c.Request.TLS = nil
|
||||
|
||||
_, err := p.Authenticate(c, store.Stores{})
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil TLS")
|
||||
}
|
||||
if !errors.Is(err, ErrNoCert) {
|
||||
t.Errorf("expected ErrNoCert, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMTLSNativeProvider_EmptyPeerCerts(t *testing.T) {
|
||||
p := NewMTLSNativeProvider(MTLSNativeConfig{})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||
c.Request.TLS = &tls.ConnectionState{
|
||||
PeerCertificates: []*x509.Certificate{},
|
||||
}
|
||||
|
||||
_, err := p.Authenticate(c, store.Stores{})
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty peer certs")
|
||||
}
|
||||
if !errors.Is(err, ErrNoCert) {
|
||||
t.Errorf("expected ErrNoCert, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMTLSNativeProvider_NoCN(t *testing.T) {
|
||||
p := NewMTLSNativeProvider(MTLSNativeConfig{AutoActivate: true})
|
||||
ca := newTestCA(t)
|
||||
cert := ca.issueCert(t, "", "", []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, time.Now().Add(time.Hour))
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||
c.Request.TLS = &tls.ConnectionState{
|
||||
PeerCertificates: []*x509.Certificate{cert.Cert},
|
||||
}
|
||||
|
||||
_, err := p.Authenticate(c, store.Stores{})
|
||||
if err == nil {
|
||||
t.Fatal("expected error for cert with no CN")
|
||||
}
|
||||
if !errors.Is(err, ErrInvalidCreds) {
|
||||
t.Errorf("expected ErrInvalidCreds, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMTLSNativeProvider_ExtractsEmail(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
cert := ca.issueClient(t, "alice", "alice@example.com")
|
||||
|
||||
if len(cert.Cert.EmailAddresses) == 0 || cert.Cert.EmailAddresses[0] != "alice@example.com" {
|
||||
t.Fatalf("cert should have email alice@example.com, got %v", cert.Cert.EmailAddresses)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Integration tests (real TLS listener) ───────────────────────────
|
||||
|
||||
func TestTLS_NoClientCert_Rejected(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
serverCert := ca.issueNode(t, "server")
|
||||
|
||||
srv := newMTLSTestServer(t, ca, serverCert)
|
||||
defer srv.Close()
|
||||
|
||||
client := &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: ca.Pool,
|
||||
MinVersion: tls.VersionTLS13,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := client.Get(srv.URL + "/test")
|
||||
if err == nil {
|
||||
t.Fatal("expected TLS handshake error when no client cert is presented")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLS_WrongCA_Rejected(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
wrongCA := newTestCA(t)
|
||||
serverCert := ca.issueNode(t, "server")
|
||||
clientCert := wrongCA.issueClient(t, "intruder", "")
|
||||
|
||||
srv := newMTLSTestServer(t, ca, serverCert)
|
||||
defer srv.Close()
|
||||
|
||||
client := &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: ca.Pool,
|
||||
Certificates: []tls.Certificate{clientCert.TLSCert},
|
||||
MinVersion: tls.VersionTLS13,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := client.Get(srv.URL + "/test")
|
||||
if err == nil {
|
||||
t.Fatal("expected TLS handshake error when client cert is from wrong CA")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLS_ValidClientCert_Accepted(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
serverCert := ca.issueNode(t, "server")
|
||||
clientCert := ca.issueClient(t, "alice", "")
|
||||
|
||||
srv := newMTLSTestServer(t, ca, serverCert)
|
||||
defer srv.Close()
|
||||
|
||||
client := &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: ca.Pool,
|
||||
Certificates: []tls.Certificate{clientCert.TLSCert},
|
||||
MinVersion: tls.VersionTLS13,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
resp, err := client.Get(srv.URL + "/test")
|
||||
if err != nil {
|
||||
t.Fatalf("expected successful connection, got: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != 200 {
|
||||
t.Errorf("status = %d, want 200", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLS_ExpiredCert_Rejected(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
serverCert := ca.issueNode(t, "server")
|
||||
clientCert := ca.issueExpiredClient(t, "expired-user")
|
||||
|
||||
srv := newMTLSTestServer(t, ca, serverCert)
|
||||
defer srv.Close()
|
||||
|
||||
client := &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: ca.Pool,
|
||||
Certificates: []tls.Certificate{clientCert.TLSCert},
|
||||
MinVersion: tls.VersionTLS13,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := client.Get(srv.URL + "/test")
|
||||
if err == nil {
|
||||
t.Fatal("expected TLS handshake error for expired client cert")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLS_PeerCertificateVisible(t *testing.T) {
|
||||
ca := newTestCA(t)
|
||||
serverCert := ca.issueNode(t, "server")
|
||||
clientCert := ca.issueClient(t, "bob", "bob@example.com")
|
||||
|
||||
var seenCN string
|
||||
var seenEmails []string
|
||||
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.TLS != nil && len(r.TLS.PeerCertificates) > 0 {
|
||||
peer := r.TLS.PeerCertificates[0]
|
||||
seenCN = peer.Subject.CommonName
|
||||
seenEmails = peer.EmailAddresses
|
||||
}
|
||||
w.WriteHeader(200)
|
||||
})
|
||||
|
||||
srv := newMTLSTestServerWithHandler(t, ca, serverCert, handler)
|
||||
defer srv.Close()
|
||||
|
||||
client := &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: ca.Pool,
|
||||
Certificates: []tls.Certificate{clientCert.TLSCert},
|
||||
MinVersion: tls.VersionTLS13,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
resp, err := client.Get(srv.URL + "/test")
|
||||
if err != nil {
|
||||
t.Fatalf("request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if seenCN != "bob" {
|
||||
t.Errorf("CN = %q, want %q", seenCN, "bob")
|
||||
}
|
||||
if len(seenEmails) == 0 || seenEmails[0] != "bob@example.com" {
|
||||
t.Errorf("emails = %v, want [bob@example.com]", seenEmails)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Test helpers ────────────────────────────────────────────────────
|
||||
|
||||
func newMTLSTestServer(t *testing.T, ca *testCA, serverCert *testCert) *httptest.Server {
|
||||
t.Helper()
|
||||
return newMTLSTestServerWithHandler(t, ca, serverCert, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(200)
|
||||
fmt.Fprint(w, "ok")
|
||||
}))
|
||||
}
|
||||
|
||||
func newMTLSTestServerWithHandler(t *testing.T, ca *testCA, serverCert *testCert, handler http.Handler) *httptest.Server {
|
||||
t.Helper()
|
||||
srv := httptest.NewUnstartedServer(handler)
|
||||
srv.TLS = &tls.Config{
|
||||
Certificates: []tls.Certificate{serverCert.TLSCert},
|
||||
ClientAuth: tls.RequireAndVerifyClientCert,
|
||||
ClientCAs: ca.Pool,
|
||||
MinVersion: tls.VersionTLS13,
|
||||
}
|
||||
srv.StartTLS()
|
||||
return srv
|
||||
}
|
||||
Reference in New Issue
Block a user