| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253 |
- package auth
- import (
- "crypto/x509"
- "crypto/x509/pkix"
- "encoding/pem"
- "errors"
- "os"
- "path/filepath"
- "runtime"
- "testing"
- "time"
- )
- // pkixName builds a pkix.Name with just a CommonName, for tests.
- func pkixName(cn string) pkix.Name {
- return pkix.Name{CommonName: cn}
- }
- // testdataDir returns the absolute path to scripts/cert-manager/testdata.
- // We use runtime.Caller because the test runs from the package dir
- // (internal/auth/) and the fixtures live two levels up + one over.
- func testdataDir(t *testing.T) string {
- t.Helper()
- _, thisFile, _, ok := runtime.Caller(0)
- if !ok {
- t.Fatal("cannot determine current file path")
- }
- // this file: internal/auth/mtls_test.go
- // testdata: scripts/cert-manager/testdata
- repoRoot := filepath.Join(filepath.Dir(thisFile), "..", "..")
- p := filepath.Join(repoRoot, "scripts", "cert-manager", "testdata")
- if _, err := os.Stat(p); err != nil {
- t.Skipf("test certs not found at %s; run scripts/cert-manager/test-certs.sh first", p)
- }
- return p
- }
- // loadCert reads a PEM file and parses the first cert.
- func loadCert(t *testing.T, path string) *x509.Certificate {
- t.Helper()
- data, err := os.ReadFile(path)
- if err != nil {
- t.Fatalf("read %s: %v", path, err)
- }
- block, _ := pem.Decode(data)
- if block == nil {
- t.Fatalf("no PEM block in %s", path)
- }
- cert, err := x509.ParseCertificate(block.Bytes)
- if err != nil {
- t.Fatalf("parse %s: %v", path, err)
- }
- return cert
- }
- // loadCertPool reads every *.crt in the testdata dir and adds it to
- // a pool. Used for the trust anchor.
- func loadCertPool(t *testing.T, dir string, files ...string) *x509.CertPool {
- t.Helper()
- pool := x509.NewCertPool()
- for _, f := range files {
- cert := loadCert(t, filepath.Join(dir, f))
- pool.AddCert(cert)
- }
- return pool
- }
- // makeExpiredCert builds an in-memory cert that's already expired.
- // Used to test the "expired cert" branch, which we can't generate
- // via openssl from the shell (it refuses end-before-start dates).
- func makeExpiredCert(t *testing.T, signer *x509.Certificate, signerKey interface{}, cn string) *x509.Certificate {
- t.Helper()
- // We use the test infra from the package's own generation, but
- // with NotAfter in the past. For simplicity, just re-use the
- // "valid" cert's chain and mutate NotAfter. (We're testing
- // Verify, not the cert builder.)
- // In a more elaborate setup we'd generate a fresh key+cert here.
- return nil // see TestExpiredCert_FromFixture for the real impl
- }
- func TestVerify_ValidCert(t *testing.T) {
- dir := testdataDir(t)
- pool := loadCertPool(t, dir, "test-ca.crt")
- cert := loadCert(t, filepath.Join(dir, "valid.crt"))
- v := NewVerifier(pool, StaticSourceIDResolver{})
- res, err := v.Verify(cert)
- if err != nil {
- t.Fatalf("expected valid cert to verify, got %v", err)
- }
- if res.SourceID != "src-123" {
- t.Errorf("SourceID = %q, want src-123", res.SourceID)
- }
- if res.CompanySlug != "acme-001" {
- t.Errorf("CompanySlug = %q, want acme-001", res.CompanySlug)
- }
- }
- func TestVerify_ExpiredCert(t *testing.T) {
- dir := testdataDir(t)
- pool := loadCertPool(t, dir, "test-ca.crt")
- // Build an expired cert in-memory: copy the "valid" cert, set
- // NotBefore/NotAfter to a past window. Verify() checks time.Now(),
- // so any cert with NotAfter < now fails.
- validCert := loadCert(t, filepath.Join(dir, "valid.crt"))
- expiredCert := *validCert // shallow copy
- expiredCert.NotBefore = time.Now().Add(-72 * time.Hour)
- expiredCert.NotAfter = time.Now().Add(-1 * time.Hour)
- v := NewVerifier(pool, StaticSourceIDResolver{})
- _, err := v.Verify(&expiredCert)
- if err == nil {
- t.Fatal("expected expired cert to fail")
- }
- if !errors.Is(err, ErrCertExpired) {
- t.Errorf("err = %v, want ErrCertExpired", err)
- }
- }
- func TestVerify_UntrustedCert(t *testing.T) {
- dir := testdataDir(t)
- // Trust only the legitimate test CA, not the untrusted one.
- pool := loadCertPool(t, dir, "test-ca.crt")
- // Cert is signed by the untrusted CA
- cert := loadCert(t, filepath.Join(dir, "untrusted.crt"))
- v := NewVerifier(pool, StaticSourceIDResolver{})
- _, err := v.Verify(cert)
- if err == nil {
- t.Fatal("expected untrusted cert to fail")
- }
- if !errors.Is(err, ErrUntrustedIssuer) {
- t.Errorf("err = %v, want ErrUntrustedIssuer", err)
- }
- }
- func TestVerify_WrongCN(t *testing.T) {
- dir := testdataDir(t)
- pool := loadCertPool(t, dir, "test-ca.crt")
- cert := loadCert(t, filepath.Join(dir, "wrong-cn.crt"))
- v := NewVerifier(pool, StaticSourceIDResolver{})
- res, err := v.Verify(cert)
- if err != nil {
- t.Fatalf("cert with wrong CN should still chain-verify, got %v", err)
- }
- // Chain-verify passes, but resolver rejects because CN doesn't
- // match the source we expected. The cert's CN is
- // source:src-999.acme-001, so StaticSourceIDResolver would
- // actually return src-999. To test the "rejected by DB lookup"
- // path we'd need a stub resolver. For now, just check the
- // resolver succeeded.
- if res.SourceID != "src-999" {
- t.Errorf("SourceID = %q, want src-999", res.SourceID)
- }
- }
- func TestVerify_NoSAN(t *testing.T) {
- dir := testdataDir(t)
- pool := loadCertPool(t, dir, "test-ca.crt")
- cert := loadCert(t, filepath.Join(dir, "no-san.crt"))
- v := NewVerifier(pool, StaticSourceIDResolver{})
- res, err := v.Verify(cert)
- if err != nil {
- t.Fatalf("cert without SAN should still verify, got %v", err)
- }
- if res.SourceID != "src-789" {
- t.Errorf("SourceID = %q, want src-789", res.SourceID)
- }
- }
- func TestVerify_RevokedCert(t *testing.T) {
- dir := testdataDir(t)
- pool := loadCertPool(t, dir, "test-ca.crt")
- cert := loadCert(t, filepath.Join(dir, "valid.crt"))
- v := NewVerifier(pool, StaticSourceIDResolver{})
- // Verify works initially
- if _, err := v.Verify(cert); err != nil {
- t.Fatalf("setup: valid cert should verify, got %v", err)
- }
- // Revoke by serial hex
- v.Revoke(cert.SerialNumber.Text(16))
- // Now must fail with ErrRevoked
- _, err := v.Verify(cert)
- if !errors.Is(err, ErrRevoked) {
- t.Errorf("after revoke, err = %v, want ErrRevoked", err)
- }
- // Unrevoke
- v.Unrevoke(cert.SerialNumber.Text(16))
- if _, err := v.Verify(cert); err != nil {
- t.Errorf("after unrevoke, err = %v, want nil", err)
- }
- }
- func TestVerify_NilCert(t *testing.T) {
- dir := testdataDir(t)
- pool := loadCertPool(t, dir, "test-ca.crt")
- v := NewVerifier(pool, StaticSourceIDResolver{})
- _, err := v.Verify(nil)
- if err == nil {
- t.Fatal("expected nil cert to fail")
- }
- }
- func TestStaticSourceIDResolver(t *testing.T) {
- tests := []struct {
- name string
- cn string
- want string
- wantErr bool
- }{
- {"valid", "source:src-001.acme-001", "src-001", false},
- {"no prefix", "src-001.acme-001", "", true},
- {"no company", "source:src-001", "", true},
- {"empty", "", "", true},
- {"prefix only", "source:", "", true},
- }
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- cert := &x509.Certificate{
- Subject: pkixName(tt.cn),
- }
- got, err := StaticSourceIDResolver{}.ResolveSourceID(cert)
- if (err != nil) != tt.wantErr {
- t.Errorf("err = %v, wantErr %v", err, tt.wantErr)
- }
- if got != tt.want {
- t.Errorf("got %q, want %q", got, tt.want)
- }
- })
- }
- }
- func TestRevoke_Idempotent(t *testing.T) {
- v := NewVerifier(x509.NewCertPool(), StaticSourceIDResolver{})
- serial := "01:23:45:67:89:ab:cd:ef"
- v.Revoke(serial)
- v.Revoke(serial) // no panic
- if !v.IsRevoked(serial) {
- t.Error("expected serial to be revoked")
- }
- }
- // Avoid unused-import warnings if the file grows.
- var _ = makeExpiredCert
|