From 7f8b995ab657a2fe0587c3218f64fcec013cac43 Mon Sep 17 00:00:00 2001 From: Elijah Rodriguez-Beltran Date: Mon, 31 Aug 2026 21:22:51 +0000 Subject: [PATCH] Cache CA certificate pool in CSI plugin TLS verification. Previously, resolveTLSConfig configured VerifyConnection to re-read the CA trust bundle from disk and rebuild the *x509.CertPool on every TLS handshake. This caused unnecessary filesystem I/O and PEM decoding overhead on repeated connections. This change introduces a thread-safe caPoolCache that caches the parsed cert pool in memory and checks file metadata (stat matching) and certificate expiration (earliest NotAfter) to avoid redundant disk reads while preserving automatic CA rotation pickup. TAG=agy CONV=08b81b20-9977-45df-b98e-9e340bc55e2d --- internal/volume/csi/plugin.go | 95 ++++++++++++++++++++++++++++----- internal/volume/csi/tls_test.go | 77 ++++++++++++++++++++++++++ 2 files changed, 160 insertions(+), 12 deletions(-) diff --git a/internal/volume/csi/plugin.go b/internal/volume/csi/plugin.go index 4bb5304138..7282780855 100644 --- a/internal/volume/csi/plugin.go +++ b/internal/volume/csi/plugin.go @@ -18,10 +18,13 @@ import ( "context" "crypto/tls" "crypto/x509" + "encoding/pem" "fmt" "log/slog" "os" "path/filepath" + "sync" + "time" "github.com/agent-substrate/substrate/internal/ateompath" "github.com/agent-substrate/substrate/internal/credbundle" @@ -336,9 +339,10 @@ func resolveTLSConfig(cfg *v1alpha1.CSIDriverConfig, paths tlsPaths) (*tls.Confi return nil, fmt.Errorf("only pod identity TLS is supported in this configuration") } - // Verify CA pool exists and is readable at construction time. - _, err := getCertPool(paths.caCert) - if err != nil { + caCache := newCAPoolCache(paths.caCert) + + // Verify CA pool exists, is readable, and populate the initial cache. + if _, err := caCache.getCertPool(); err != nil { return nil, fmt.Errorf("failed to load CA cert pool from %q: %w", paths.caCert, err) } @@ -358,15 +362,17 @@ func resolveTLSConfig(cfg *v1alpha1.CSIDriverConfig, paths tlsPaths) (*tls.Confi // Standard tls.Config.RootCAs is a static cert pool evaluated at construction time. // To automatically pick up CA trust bundle rotations on disk without restarting the process, // we set InsecureSkipVerify=true and verify the server certificate chain dynamically - // against the latest CA bundle read from disk in VerifyConnection. + // against the CA bundle in VerifyConnection. + // caCache avoids re-reading and re-parsing the CA bundle from disk on every handshake + // unless the file is modified or its certificates have expired. InsecureSkipVerify: true, VerifyConnection: func(state tls.ConnectionState) error { if len(state.PeerCertificates) == 0 { return fmt.Errorf("server did not present certificates") } - // Read CA trust bundle on each TLS connection handshake. - roots, err := getCertPool(paths.caCert) + // Retrieve CA cert pool (cached, reloaded on rotation or expiration). + roots, err := caCache.getCertPool() if err != nil { return fmt.Errorf("failed to load CA cert pool from %q: %w", paths.caCert, err) } @@ -391,14 +397,79 @@ func resolveTLSConfig(cfg *v1alpha1.CSIDriverConfig, paths tlsPaths) (*tls.Confi }, nil } -func getCertPool(path string) (*x509.CertPool, error) { +// caPoolCache holds the parsed *x509.CertPool and file stat / expiration metadata +// so that unchanged and unexpired CA trust bundles are not re-read from disk on every TLS handshake. +type caPoolCache struct { + path string + + mu sync.Mutex + fi os.FileInfo + expiry time.Time + pool *x509.CertPool +} + +func newCAPoolCache(path string) *caPoolCache { + return &caPoolCache{path: path} +} + +// getCertPool returns the parsed CA cert pool, re-reading the file only when it has changed +// on disk (identity, modification time, or size) or when the cached certificates have expired. +func (c *caPoolCache) getCertPool() (*x509.CertPool, error) { + c.mu.Lock() + defer c.mu.Unlock() + + fi, err := os.Stat(c.path) + if err != nil { + return nil, fmt.Errorf("failed to stat CA cert file %q: %w", c.path, err) + } + + if c.pool != nil && os.SameFile(c.fi, fi) && fi.ModTime().Equal(c.fi.ModTime()) && fi.Size() == c.fi.Size() && time.Now().Before(c.expiry) { + return c.pool, nil + } + + pool, expiry, err := parseCertPoolWithExpiry(c.path) + if err != nil { + return nil, err + } + + c.fi, c.pool, c.expiry = fi, pool, expiry + return pool, nil +} + +func parseCertPoolWithExpiry(path string) (*x509.CertPool, time.Time, error) { certBytes, err := os.ReadFile(path) if err != nil { - return nil, fmt.Errorf("failed to read cert file %q: %w", path, err) + return nil, time.Time{}, fmt.Errorf("failed to read cert file %q: %w", path, err) + } + + pool := x509.NewCertPool() + var earliestExpiry time.Time + var count int + + rest := certBytes + for { + var block *pem.Block + block, rest = pem.Decode(rest) + if block == nil { + break + } + if block.Type != "CERTIFICATE" { + continue + } + cert, err := x509.ParseCertificate(block.Bytes) + if err != nil { + return nil, time.Time{}, fmt.Errorf("failed to parse certificate from %q: %w", path, err) + } + pool.AddCert(cert) + count++ + if earliestExpiry.IsZero() || cert.NotAfter.Before(earliestExpiry) { + earliestExpiry = cert.NotAfter + } } - certPool := x509.NewCertPool() - if !certPool.AppendCertsFromPEM(certBytes) { - return nil, fmt.Errorf("failed to parse certs from %q", path) + + if count == 0 { + return nil, time.Time{}, fmt.Errorf("failed to parse certs from %q: no valid CERTIFICATE blocks found", path) } - return certPool, nil + + return pool, earliestExpiry, nil } diff --git a/internal/volume/csi/tls_test.go b/internal/volume/csi/tls_test.go index 68738c6544..7171cc91ed 100644 --- a/internal/volume/csi/tls_test.go +++ b/internal/volume/csi/tls_test.go @@ -364,3 +364,80 @@ func TestMTLSPicksUpCARotation(t *testing.T) { } plugin2.client.Close() } + +func TestCAPoolCache_HitAndFileModification(t *testing.T) { + t.Parallel() + ca1 := newTestCA(t) + ca2 := newTestCA(t) + + dir := t.TempDir() + caPath := filepath.Join(dir, "trust-bundle.pem") + writeFile(t, caPath, ca1.certPEM()) + + cache := newCAPoolCache(caPath) + + pool1, err := cache.getCertPool() + if err != nil { + t.Fatalf("getCertPool (1st call): %v", err) + } + + // 2nd call should return the exact cached instance (pointer equality). + pool2, err := cache.getCertPool() + if err != nil { + t.Fatalf("getCertPool (2nd call): %v", err) + } + if pool1 != pool2 { + t.Errorf("expected cached cert pool pointer equality on unchanged file, got %p != %p", pool1, pool2) + } + + // Modify the file on disk to ca2. + time.Sleep(10 * time.Millisecond) // Ensure mtime advances on fast filesystems + writeFile(t, caPath, ca2.certPEM()) + + // 3rd call should detect file change and return a newly parsed pool. + pool3, err := cache.getCertPool() + if err != nil { + t.Fatalf("getCertPool (3rd call after edit): %v", err) + } + if pool1 == pool3 { + t.Errorf("expected new cert pool after file modification, got same pointer %p", pool3) + } +} + +func TestCAPoolCache_ExpiryReload(t *testing.T) { + t.Parallel() + key := newKey(t) + // CA cert with short lifespan in the past + tmpl := &x509.Certificate{ + SerialNumber: big.NewInt(100), + Subject: pkix.Name{CommonName: "expired-ca"}, + NotBefore: time.Now().Add(-2 * time.Hour), + NotAfter: time.Now().Add(-time.Hour), + IsCA: true, + KeyUsage: x509.KeyUsageCertSign, + BasicConstraintsValid: true, + } + der := createCert(t, tmpl, tmpl, &key.PublicKey, key) + expiredPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) + + dir := t.TempDir() + caPath := filepath.Join(dir, "trust-bundle.pem") + writeFile(t, caPath, expiredPEM) + + cache := newCAPoolCache(caPath) + + pool1, err := cache.getCertPool() + if err != nil { + t.Fatalf("getCertPool (expired CA): %v", err) + } + + // Because NotAfter is in the past, expiry check (time.Now().Before(c.expiry)) fails, + // forcing a reload on the next call even if the file hasn't changed. + pool2, err := cache.getCertPool() + if err != nil { + t.Fatalf("getCertPool (2nd call on expired CA): %v", err) + } + if pool1 == pool2 { + t.Errorf("expected reload for expired CA bundle, got same pointer %p", pool2) + } +}