diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd.go b/enterprise/aibridgeproxyd/aibridgeproxyd.go index 0f6c86637a..adf6848aa6 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd.go @@ -6,6 +6,7 @@ import ( "crypto/x509" "errors" "net/http" + "sync" "time" "github.com/elazarl/goproxy" @@ -14,6 +15,13 @@ import ( "cdr.dev/slog" ) +// loadMitmOnce ensures the MITM certificate is loaded exactly once. +// goproxy.GoproxyCa is a package-level global variable shared across all +// goproxy.ProxyHttpServer instances in the process. In tests, multiple proxy +// servers run in parallel, and without this guard they would race on writing +// to GoproxyCa. In production, only one server runs, so this has no impact. +var loadMitmOnce sync.Once + // Server is the AI MITM (Man-in-the-Middle) proxy server. // It is responsible for: // - intercepting HTTPS requests to AI providers @@ -91,23 +99,28 @@ func (s *Server) Close() error { return s.httpServer.Shutdown(ctx) } -// loadMitmCertificate loads the CA certificate and key for MITM into goproxy. +// loadMitmCertificate loads the CA certificate and private key for MITM proxying. +// This function is safe to call concurrently - the certificate is only loaded once +// into the global goproxy.GoproxyCa variable. func loadMitmCertificate(certFile, keyFile string) error { tlsCert, err := tls.LoadX509KeyPair(certFile, keyFile) if err != nil { - return xerrors.Errorf("load x509 keypair: %w", err) + return xerrors.Errorf("load CA certificate: %w", err) } x509Cert, err := x509.ParseCertificate(tlsCert.Certificate[0]) if err != nil { - return xerrors.Errorf("parse certificate: %w", err) + return xerrors.Errorf("parse CA certificate: %w", err) } - goproxy.GoproxyCa = tls.Certificate{ - Certificate: tlsCert.Certificate, - PrivateKey: tlsCert.PrivateKey, - Leaf: x509Cert, - } + // Only protect the global assignment with sync.Once + loadMitmOnce.Do(func() { + goproxy.GoproxyCa = tls.Certificate{ + Certificate: tlsCert.Certificate, + PrivateKey: tlsCert.PrivateKey, + Leaf: x509Cert, + } + }) return nil } diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go index 5efc31e3df..8b5b05773a 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go @@ -15,10 +15,12 @@ import ( "net/url" "os" "path/filepath" + "sync" "testing" "time" "github.com/stretchr/testify/require" + "golang.org/x/xerrors" "cdr.dev/slog/sloggers/slogtest" @@ -26,18 +28,43 @@ import ( "github.com/coder/coder/v2/testutil" ) -// generateTestCA creates a temporary CA certificate and key for testing. -func generateTestCA(t *testing.T) (certFile, keyFile string) { +var ( + // testCAOnce ensures the shared CA is generated exactly once. + // sync.Once guarantees single execution even with parallel tests. + // Note: no retry on failure. + testCAOnce sync.Once + // Shared CA certificate and key paths, and any error from generation. + // These are set once by testCAOnce and read by all tests. + testCACert string + testCAKey string + errTestSharedCA error +) + +// getSharedTestCA returns a shared CA certificate for all tests. +// This avoids race conditions with goproxy.GoproxyCa which is a global variable. +// Using sync.Once ensures the CA is generated exactly once, even when tests run +// in parallel. All tests share the same CA, so goproxy.GoproxyCa is only set once. +func getSharedTestCA(t *testing.T) (certFile, keyFile string) { t.Helper() - privateKey, err := rsa.GenerateKey(rand.Reader, 2048) - require.NoError(t, err) + testCAOnce.Do(func() { + testCACert, testCAKey, errTestSharedCA = generateSharedTestCA() + }) - template := &x509.Certificate{ - SerialNumber: big.NewInt(1), - Subject: pkix.Name{ - CommonName: "Test CA", - }, + require.NoError(t, errTestSharedCA, "failed to generate shared test CA") + return testCACert, testCAKey +} + +// generateSharedTestCA creates a shared CA certificate and key for testing. +func generateSharedTestCA() (certFile, keyFile string, err error) { + caKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + return "", "", xerrors.Errorf("generate CA key: %w", err) + } + + caTemplate := x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "Shared Test CA"}, NotBefore: time.Now(), NotAfter: time.Now().Add(time.Hour), KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign, @@ -45,22 +72,26 @@ func generateTestCA(t *testing.T) (certFile, keyFile string) { IsCA: true, } - certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) - require.NoError(t, err) + caCertDER, err := x509.CreateCertificate(rand.Reader, &caTemplate, &caTemplate, &caKey.PublicKey, caKey) + if err != nil { + return "", "", xerrors.Errorf("create CA certificate: %w", err) + } - tempDir := t.TempDir() - certFile = filepath.Join(tempDir, "ca.crt") - keyFile = filepath.Join(tempDir, "ca.key") + tmpDir := os.TempDir() + certPath := filepath.Join(tmpDir, "aibridgeproxyd_test_ca.crt") + keyPath := filepath.Join(tmpDir, "aibridgeproxyd_test_ca.key") - certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) - err = os.WriteFile(certFile, certPEM, 0o600) - require.NoError(t, err) + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: caCertDER}) + if err := os.WriteFile(certPath, certPEM, 0o600); err != nil { + return "", "", xerrors.Errorf("write cert file: %w", err) + } - keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)}) - err = os.WriteFile(keyFile, keyPEM, 0o600) - require.NoError(t, err) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(caKey)}) + if err := os.WriteFile(keyPath, keyPEM, 0o600); err != nil { + return "", "", xerrors.Errorf("write key file: %w", err) + } - return certFile, keyFile + return certPath, keyPath, nil } func TestNew(t *testing.T) { @@ -69,7 +100,7 @@ func TestNew(t *testing.T) { t.Run("MissingListenAddr", func(t *testing.T) { t.Parallel() - certFile, keyFile := generateTestCA(t) + certFile, keyFile := getSharedTestCA(t) logger := slogtest.Make(t, nil) _, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{ @@ -124,7 +155,7 @@ func TestNew(t *testing.T) { t.Run("Success", func(t *testing.T) { t.Parallel() - certFile, keyFile := generateTestCA(t) + certFile, keyFile := getSharedTestCA(t) logger := slogtest.Make(t, nil) srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{ @@ -143,7 +174,7 @@ func TestNew(t *testing.T) { func TestClose(t *testing.T) { t.Parallel() - certFile, keyFile := generateTestCA(t) + certFile, keyFile := getSharedTestCA(t) logger := slogtest.Make(t, nil) srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{ @@ -171,7 +202,7 @@ func TestProxy_MITM(t *testing.T) { })) defer targetServer.Close() - certFile, keyFile := generateTestCA(t) + certFile, keyFile := getSharedTestCA(t) logger := slogtest.Make(t, nil) // Start the proxy server.