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" "armature/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 }