mtls_test.go 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253
  1. package auth
  2. import (
  3. "crypto/x509"
  4. "crypto/x509/pkix"
  5. "encoding/pem"
  6. "errors"
  7. "os"
  8. "path/filepath"
  9. "runtime"
  10. "testing"
  11. "time"
  12. )
  13. // pkixName builds a pkix.Name with just a CommonName, for tests.
  14. func pkixName(cn string) pkix.Name {
  15. return pkix.Name{CommonName: cn}
  16. }
  17. // testdataDir returns the absolute path to scripts/cert-manager/testdata.
  18. // We use runtime.Caller because the test runs from the package dir
  19. // (internal/auth/) and the fixtures live two levels up + one over.
  20. func testdataDir(t *testing.T) string {
  21. t.Helper()
  22. _, thisFile, _, ok := runtime.Caller(0)
  23. if !ok {
  24. t.Fatal("cannot determine current file path")
  25. }
  26. // this file: internal/auth/mtls_test.go
  27. // testdata: scripts/cert-manager/testdata
  28. repoRoot := filepath.Join(filepath.Dir(thisFile), "..", "..")
  29. p := filepath.Join(repoRoot, "scripts", "cert-manager", "testdata")
  30. if _, err := os.Stat(p); err != nil {
  31. t.Skipf("test certs not found at %s; run scripts/cert-manager/test-certs.sh first", p)
  32. }
  33. return p
  34. }
  35. // loadCert reads a PEM file and parses the first cert.
  36. func loadCert(t *testing.T, path string) *x509.Certificate {
  37. t.Helper()
  38. data, err := os.ReadFile(path)
  39. if err != nil {
  40. t.Fatalf("read %s: %v", path, err)
  41. }
  42. block, _ := pem.Decode(data)
  43. if block == nil {
  44. t.Fatalf("no PEM block in %s", path)
  45. }
  46. cert, err := x509.ParseCertificate(block.Bytes)
  47. if err != nil {
  48. t.Fatalf("parse %s: %v", path, err)
  49. }
  50. return cert
  51. }
  52. // loadCertPool reads every *.crt in the testdata dir and adds it to
  53. // a pool. Used for the trust anchor.
  54. func loadCertPool(t *testing.T, dir string, files ...string) *x509.CertPool {
  55. t.Helper()
  56. pool := x509.NewCertPool()
  57. for _, f := range files {
  58. cert := loadCert(t, filepath.Join(dir, f))
  59. pool.AddCert(cert)
  60. }
  61. return pool
  62. }
  63. // makeExpiredCert builds an in-memory cert that's already expired.
  64. // Used to test the "expired cert" branch, which we can't generate
  65. // via openssl from the shell (it refuses end-before-start dates).
  66. func makeExpiredCert(t *testing.T, signer *x509.Certificate, signerKey interface{}, cn string) *x509.Certificate {
  67. t.Helper()
  68. // We use the test infra from the package's own generation, but
  69. // with NotAfter in the past. For simplicity, just re-use the
  70. // "valid" cert's chain and mutate NotAfter. (We're testing
  71. // Verify, not the cert builder.)
  72. // In a more elaborate setup we'd generate a fresh key+cert here.
  73. return nil // see TestExpiredCert_FromFixture for the real impl
  74. }
  75. func TestVerify_ValidCert(t *testing.T) {
  76. dir := testdataDir(t)
  77. pool := loadCertPool(t, dir, "test-ca.crt")
  78. cert := loadCert(t, filepath.Join(dir, "valid.crt"))
  79. v := NewVerifier(pool, StaticSourceIDResolver{})
  80. res, err := v.Verify(cert)
  81. if err != nil {
  82. t.Fatalf("expected valid cert to verify, got %v", err)
  83. }
  84. if res.SourceID != "src-123" {
  85. t.Errorf("SourceID = %q, want src-123", res.SourceID)
  86. }
  87. if res.CompanySlug != "acme-001" {
  88. t.Errorf("CompanySlug = %q, want acme-001", res.CompanySlug)
  89. }
  90. }
  91. func TestVerify_ExpiredCert(t *testing.T) {
  92. dir := testdataDir(t)
  93. pool := loadCertPool(t, dir, "test-ca.crt")
  94. // Build an expired cert in-memory: copy the "valid" cert, set
  95. // NotBefore/NotAfter to a past window. Verify() checks time.Now(),
  96. // so any cert with NotAfter < now fails.
  97. validCert := loadCert(t, filepath.Join(dir, "valid.crt"))
  98. expiredCert := *validCert // shallow copy
  99. expiredCert.NotBefore = time.Now().Add(-72 * time.Hour)
  100. expiredCert.NotAfter = time.Now().Add(-1 * time.Hour)
  101. v := NewVerifier(pool, StaticSourceIDResolver{})
  102. _, err := v.Verify(&expiredCert)
  103. if err == nil {
  104. t.Fatal("expected expired cert to fail")
  105. }
  106. if !errors.Is(err, ErrCertExpired) {
  107. t.Errorf("err = %v, want ErrCertExpired", err)
  108. }
  109. }
  110. func TestVerify_UntrustedCert(t *testing.T) {
  111. dir := testdataDir(t)
  112. // Trust only the legitimate test CA, not the untrusted one.
  113. pool := loadCertPool(t, dir, "test-ca.crt")
  114. // Cert is signed by the untrusted CA
  115. cert := loadCert(t, filepath.Join(dir, "untrusted.crt"))
  116. v := NewVerifier(pool, StaticSourceIDResolver{})
  117. _, err := v.Verify(cert)
  118. if err == nil {
  119. t.Fatal("expected untrusted cert to fail")
  120. }
  121. if !errors.Is(err, ErrUntrustedIssuer) {
  122. t.Errorf("err = %v, want ErrUntrustedIssuer", err)
  123. }
  124. }
  125. func TestVerify_WrongCN(t *testing.T) {
  126. dir := testdataDir(t)
  127. pool := loadCertPool(t, dir, "test-ca.crt")
  128. cert := loadCert(t, filepath.Join(dir, "wrong-cn.crt"))
  129. v := NewVerifier(pool, StaticSourceIDResolver{})
  130. res, err := v.Verify(cert)
  131. if err != nil {
  132. t.Fatalf("cert with wrong CN should still chain-verify, got %v", err)
  133. }
  134. // Chain-verify passes, but resolver rejects because CN doesn't
  135. // match the source we expected. The cert's CN is
  136. // source:src-999.acme-001, so StaticSourceIDResolver would
  137. // actually return src-999. To test the "rejected by DB lookup"
  138. // path we'd need a stub resolver. For now, just check the
  139. // resolver succeeded.
  140. if res.SourceID != "src-999" {
  141. t.Errorf("SourceID = %q, want src-999", res.SourceID)
  142. }
  143. }
  144. func TestVerify_NoSAN(t *testing.T) {
  145. dir := testdataDir(t)
  146. pool := loadCertPool(t, dir, "test-ca.crt")
  147. cert := loadCert(t, filepath.Join(dir, "no-san.crt"))
  148. v := NewVerifier(pool, StaticSourceIDResolver{})
  149. res, err := v.Verify(cert)
  150. if err != nil {
  151. t.Fatalf("cert without SAN should still verify, got %v", err)
  152. }
  153. if res.SourceID != "src-789" {
  154. t.Errorf("SourceID = %q, want src-789", res.SourceID)
  155. }
  156. }
  157. func TestVerify_RevokedCert(t *testing.T) {
  158. dir := testdataDir(t)
  159. pool := loadCertPool(t, dir, "test-ca.crt")
  160. cert := loadCert(t, filepath.Join(dir, "valid.crt"))
  161. v := NewVerifier(pool, StaticSourceIDResolver{})
  162. // Verify works initially
  163. if _, err := v.Verify(cert); err != nil {
  164. t.Fatalf("setup: valid cert should verify, got %v", err)
  165. }
  166. // Revoke by serial hex
  167. v.Revoke(cert.SerialNumber.Text(16))
  168. // Now must fail with ErrRevoked
  169. _, err := v.Verify(cert)
  170. if !errors.Is(err, ErrRevoked) {
  171. t.Errorf("after revoke, err = %v, want ErrRevoked", err)
  172. }
  173. // Unrevoke
  174. v.Unrevoke(cert.SerialNumber.Text(16))
  175. if _, err := v.Verify(cert); err != nil {
  176. t.Errorf("after unrevoke, err = %v, want nil", err)
  177. }
  178. }
  179. func TestVerify_NilCert(t *testing.T) {
  180. dir := testdataDir(t)
  181. pool := loadCertPool(t, dir, "test-ca.crt")
  182. v := NewVerifier(pool, StaticSourceIDResolver{})
  183. _, err := v.Verify(nil)
  184. if err == nil {
  185. t.Fatal("expected nil cert to fail")
  186. }
  187. }
  188. func TestStaticSourceIDResolver(t *testing.T) {
  189. tests := []struct {
  190. name string
  191. cn string
  192. want string
  193. wantErr bool
  194. }{
  195. {"valid", "source:src-001.acme-001", "src-001", false},
  196. {"no prefix", "src-001.acme-001", "", true},
  197. {"no company", "source:src-001", "", true},
  198. {"empty", "", "", true},
  199. {"prefix only", "source:", "", true},
  200. }
  201. for _, tt := range tests {
  202. t.Run(tt.name, func(t *testing.T) {
  203. cert := &x509.Certificate{
  204. Subject: pkixName(tt.cn),
  205. }
  206. got, err := StaticSourceIDResolver{}.ResolveSourceID(cert)
  207. if (err != nil) != tt.wantErr {
  208. t.Errorf("err = %v, wantErr %v", err, tt.wantErr)
  209. }
  210. if got != tt.want {
  211. t.Errorf("got %q, want %q", got, tt.want)
  212. }
  213. })
  214. }
  215. }
  216. func TestRevoke_Idempotent(t *testing.T) {
  217. v := NewVerifier(x509.NewCertPool(), StaticSourceIDResolver{})
  218. serial := "01:23:45:67:89:ab:cd:ef"
  219. v.Revoke(serial)
  220. v.Revoke(serial) // no panic
  221. if !v.IsRevoked(serial) {
  222. t.Error("expected serial to be revoked")
  223. }
  224. }
  225. // Avoid unused-import warnings if the file grows.
  226. var _ = makeExpiredCert