diff --git a/CHANGELOG.md b/CHANGELOG.md index 219a9afc6e..4bae601e0d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,17 @@ project adheres to [Semantic Versioning](http://semver.org/). ## Unreleased +### Custom HTTPAuthPlugin behavior change + +The `HTTPAuthPlugin.NewClient()` method is now called once per `Client` instance and cached rather than being called for every request. Custom plugins that performed per-request operations in `NewClient()` (such as request counters, per-request transport wrapping, or logging/metrics side effects) will now only execute those operations once. All per-request authentication logic must be moved from `NewClient()` to `Prepare()`. All plugins included in OPA have been updated and are unaffected by this change. + +### Runtime, SDK, Tooling + +- plugins/rest: Configurable re-read interval for TLS client certificates via `cert_reread_interval_seconds` field. + Defaults to re-reading on every request for backwards compatibility. + The implementation also uses content hashing to detect changes and avoid re-parsing unchanged TLS certificates and keys. +- plugins/rest: All TLS configurations now inherit the minimum version and TLS ciphersuites as configured for the server. + ## 1.14.1 This is a patch release collecting two bug fixes and various dependency updates for Golang standard library and common package vulnerabilities. diff --git a/internal/config/config.go b/internal/config/config.go index 53dfc6d6cb..561b8083ed 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -29,6 +29,8 @@ type ServiceOptions struct { Keys map[string]*keys.Config Logger logging.Logger DistributedTacingOpts tracing.Options + MinTLSVersion uint16 + CipherSuites *[]uint16 } // ParseServicesConfig returns a set of named service clients. The service @@ -41,10 +43,16 @@ func ParseServicesConfig(opts ServiceOptions) (map[string]rest.Client, error) { var arr []json.RawMessage var obj map[string]json.RawMessage + clientOpts := []func(*rest.Client){ + rest.AuthPluginLookup(opts.AuthPlugin), + rest.Logger(opts.Logger), + rest.DistributedTracingOpts(opts.DistributedTacingOpts), + rest.MinTLSVersion(opts.MinTLSVersion), + rest.CipherSuites(opts.CipherSuites)} if err := util.Unmarshal(opts.Raw, &arr); err == nil { for _, s := range arr { - client, err := rest.New(s, opts.Keys, rest.AuthPluginLookup(opts.AuthPlugin), rest.Logger(opts.Logger), rest.DistributedTracingOpts(opts.DistributedTacingOpts)) + client, err := rest.New(s, opts.Keys, clientOpts...) if err != nil { return nil, err } @@ -52,7 +60,7 @@ func ParseServicesConfig(opts ServiceOptions) (map[string]rest.Client, error) { } } else if util.Unmarshal(opts.Raw, &obj) == nil { for k := range obj { - client, err := rest.New(obj[k], opts.Keys, rest.Name(k), rest.AuthPluginLookup(opts.AuthPlugin), rest.Logger(opts.Logger), rest.DistributedTracingOpts(opts.DistributedTacingOpts)) + client, err := rest.New(obj[k], opts.Keys, append(clientOpts, rest.Name(k))...) if err != nil { return nil, err } diff --git a/v1/config/config.go b/v1/config/config.go index 1912d1f38c..ca1caf917d 100644 --- a/v1/config/config.go +++ b/v1/config/config.go @@ -6,6 +6,7 @@ package config import ( + "crypto/tls" "encoding/json" "errors" "fmt" @@ -22,6 +23,11 @@ import ( "github.com/open-policy-agent/opa/v1/version" ) +const ( + // DefaultMinTLSVersion is the minimum TLS version used by OPA server and REST clients + DefaultMinTLSVersion = tls.VersionTLS12 +) + // ServerConfig represents the different server configuration options. type ServerConfig struct { Metrics json.RawMessage `json:"metrics,omitempty"` diff --git a/v1/plugins/plugins.go b/v1/plugins/plugins.go index 3afd55d909..06e514d8d4 100644 --- a/v1/plugins/plugins.go +++ b/v1/plugins/plugins.go @@ -218,6 +218,8 @@ type Manager struct { bootstrapConfigLabels map[string]string hooks hooks.Hooks enableVersionCheck bool + minTLSVersion uint16 + cipherSuites *[]uint16 versionChecker versioncheck.Checker opaReportNotifyCh chan struct{} stop chan chan struct{} @@ -460,6 +462,20 @@ func WithBundleActivatorPlugin(bundleActivatorPlugin string) func(*Manager) { } } +// WithMinTLSVersion sets the minimum TLS version for REST client connections +func WithMinTLSVersion(v uint16) func(*Manager) { + return func(m *Manager) { + m.minTLSVersion = v + } +} + +// WithCipherSuites sets the cipher suites for REST client connections +func WithCipherSuites(cs *[]uint16) func(*Manager) { + return func(m *Manager) { + m.cipherSuites = cs + } +} + // New creates a new Manager using config. func New(raw []byte, id string, store storage.Store, opts ...func(*Manager)) (*Manager, error) { parsedConfig, err := config.ParseConfig(raw, id) @@ -842,6 +858,8 @@ func (m *Manager) DefaultServiceOpts(config *config.Config) cfg.ServiceOptions { Logger: m.logger, Keys: m.keys, DistributedTacingOpts: m.distributedTacingOpts, + MinTLSVersion: m.minTLSVersion, + CipherSuites: m.cipherSuites, } } diff --git a/v1/plugins/rest/auth.go b/v1/plugins/rest/auth.go index 14c2f266ac..b35c35d733 100644 --- a/v1/plugins/rest/auth.go +++ b/v1/plugins/rest/auth.go @@ -5,9 +5,9 @@ package rest import ( + "cmp" "context" "crypto/rand" - "crypto/rsa" "crypto/sha256" "crypto/sha512" "crypto/tls" @@ -33,6 +33,7 @@ import ( "github.com/lestrrat-go/jwx/v3/jws" "github.com/open-policy-agent/opa/internal/providers/aws" "github.com/open-policy-agent/opa/internal/uuid" + "github.com/open-policy-agent/opa/v1/config" "github.com/open-policy-agent/opa/v1/keys" "github.com/open-policy-agent/opa/v1/logging" ) @@ -44,59 +45,6 @@ const ( defaultClientAssertionType = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" ) -// DefaultTLSConfig defines standard TLS configurations based on the Config -func DefaultTLSConfig(c Config) (*tls.Config, error) { - t := &tls.Config{} - url, err := url.Parse(c.URL) - if err != nil { - return nil, err - } - if url.Scheme == "https" { - t.InsecureSkipVerify = c.AllowInsecureTLS - } - - if c.TLS != nil && c.TLS.CACert != "" { - caCert, err := os.ReadFile(c.TLS.CACert) - if err != nil { - return nil, err - } - - var rootCAs *x509.CertPool - if c.TLS.SystemCARequired { - rootCAs, err = x509.SystemCertPool() - if err != nil { - return nil, err - } - } else { - rootCAs = x509.NewCertPool() - } - - ok := rootCAs.AppendCertsFromPEM(caCert) - if !ok { - return nil, errors.New("unable to parse and append CA certificate to certificate pool") - } - t.RootCAs = rootCAs - } - - return t, nil -} - -// DefaultRoundTripperClient is a reasonable set of defaults for HTTP auth plugins -func DefaultRoundTripperClient(t *tls.Config, timeout int64) *http.Client { - // Ensure we use a http.Transport with proper settings: the zero values are not - // a good choice, as they cause leaking connections: - // https://github.com/golang/go/issues/19620 - - // copy, we don't want to alter the default client's Transport - tr := http.DefaultTransport.(*http.Transport).Clone() - tr.ResponseHeaderTimeout = time.Duration(timeout) * time.Second - tr.TLSClientConfig = t - - c := *http.DefaultClient - c.Transport = tr - return &c -} - // defaultAuthPlugin represents baseline 'no auth' behavior if no alternative plugin is specified for a service type defaultAuthPlugin struct{} @@ -112,11 +60,6 @@ func (*defaultAuthPlugin) Prepare(*http.Request) error { return nil } -type serverTLSConfig struct { - CACert string `json:"ca_cert,omitempty"` - SystemCARequired bool `json:"system_ca_required,omitempty"` -} - // bearerAuthPlugin represents authentication via a bearer token in the HTTP Authorization header type bearerAuthPlugin struct { Token string `json:"token"` @@ -316,6 +259,8 @@ type oauth2ClientCredentialsAuthPlugin struct { signingKeyParsed any tokenCache *oauth2Token tlsSkipVerify bool + minTLSVersion uint16 + cipherSuites *[]uint16 logger logging.Logger } @@ -557,8 +502,10 @@ func (ap *oauth2ClientCredentialsAuthPlugin) NewClient(c Config) (*http.Client, } } - // Inherit skip verify from the "parent" settings. Should this be configurable on the credentials too? + // Inherit TLS settings from the "parent" config ap.tlsSkipVerify = c.AllowInsecureTLS + ap.minTLSVersion = c.minTLSVersion + ap.cipherSuites = c.cipherSuites ap.logger = c.logger @@ -714,7 +661,15 @@ func (ap *oauth2ClientCredentialsAuthPlugin) requestToken(ctx context.Context) ( r.Header.Add(k, v) } - client := DefaultRoundTripperClient(&tls.Config{InsecureSkipVerify: ap.tlsSkipVerify}, 10) + tlsConfig := &tls.Config{ + MinVersion: cmp.Or(ap.minTLSVersion, uint16(config.DefaultMinTLSVersion)), + InsecureSkipVerify: ap.tlsSkipVerify, + } + if ap.cipherSuites != nil { + tlsConfig.CipherSuites = *ap.cipherSuites + } + + client := DefaultRoundTripperClient(tlsConfig, 10) response, err := client.Do(r) if err != nil { return nil, err @@ -761,120 +716,6 @@ func (ap *oauth2ClientCredentialsAuthPlugin) Prepare(req *http.Request) error { return nil } -// clientTLSAuthPlugin represents authentication via client certificate on a TLS connection -type clientTLSAuthPlugin struct { - Cert string `json:"cert"` - PrivateKey string `json:"private_key"` - PrivateKeyPassphrase string `json:"private_key_passphrase,omitempty"` - CACert string `json:"ca_cert,omitempty"` // Deprecated: Use `services[_].tls.ca_cert` instead - SystemCARequired bool `json:"system_ca_required,omitempty"` // Deprecated: Use `services[_].tls.system_ca_required` instead -} - -func (ap *clientTLSAuthPlugin) NewClient(c Config) (*http.Client, error) { - tlsConfig, err := DefaultTLSConfig(c) - if err != nil { - return nil, err - } - - if ap.Cert == "" { - return nil, errors.New("client certificate is needed when client TLS is enabled") - } - if ap.PrivateKey == "" { - return nil, errors.New("private key is needed when client TLS is enabled") - } - - var keyPEMBlock []byte - data, err := os.ReadFile(ap.PrivateKey) - if err != nil { - return nil, err - } - - block, _ := pem.Decode(data) - if block == nil { - return nil, errors.New("PEM data could not be found") - } - - // nolint: staticcheck // We don't want to forbid users from using this encryption. - if x509.IsEncryptedPEMBlock(block) { - if ap.PrivateKeyPassphrase == "" { - return nil, errors.New("client certificate passphrase is needed, because the certificate is password encrypted") - } - // nolint: staticcheck // We don't want to forbid users from using this encryption. - block, err := x509.DecryptPEMBlock(block, []byte(ap.PrivateKeyPassphrase)) - if err != nil { - return nil, err - } - key, err := x509.ParsePKCS8PrivateKey(block) - if err != nil { - key, err = x509.ParsePKCS1PrivateKey(block) - if err != nil { - return nil, fmt.Errorf("private key should be a PEM or plain PKCS1 or PKCS8; parse error: %v", err) - } - } - rsa, ok := key.(*rsa.PrivateKey) - if !ok { - return nil, errors.New("private key is invalid") - } - keyPEMBlock = pem.EncodeToMemory( - &pem.Block{ - Type: "RSA PRIVATE KEY", - Bytes: x509.MarshalPKCS1PrivateKey(rsa), - }, - ) - } else { - keyPEMBlock = data - } - - certPEMBlock, err := os.ReadFile(ap.Cert) - if err != nil { - return nil, err - } - - cert, err := tls.X509KeyPair(certPEMBlock, keyPEMBlock) - if err != nil { - return nil, err - } - tlsConfig.Certificates = []tls.Certificate{cert} - - var client *http.Client - - if c.TLS != nil && c.TLS.CACert != "" { - client = DefaultRoundTripperClient(tlsConfig, *c.ResponseHeaderTimeoutSeconds) - } else { - if ap.CACert != "" { - c.logger.Warn("Deprecated 'services[_].credentials.client_tls.ca_cert' configuration specified. Use 'services[_].tls.ca_cert' instead. See https://www.openpolicyagent.org/docs/latest/configuration/#services") - caCert, err := os.ReadFile(ap.CACert) - if err != nil { - return nil, err - } - - var caCertPool *x509.CertPool - if ap.SystemCARequired { - caCertPool, err = x509.SystemCertPool() - if err != nil { - return nil, err - } - } else { - caCertPool = x509.NewCertPool() - } - - ok := caCertPool.AppendCertsFromPEM(caCert) - if !ok { - return nil, errors.New("unable to parse and append CA certificate to certificate pool") - } - tlsConfig.RootCAs = caCertPool - } - - client = DefaultRoundTripperClient(tlsConfig, *c.ResponseHeaderTimeoutSeconds) - } - - return client, nil -} - -func (*clientTLSAuthPlugin) Prepare(_ *http.Request) error { - return nil -} - // awsSigningAuthPlugin represents authentication using AWS V4 HMAC signing in the Authorization header type awsSigningAuthPlugin struct { AWSEnvironmentCredentials *awsEnvironmentCredentialService `json:"environment_credentials,omitempty"` diff --git a/v1/plugins/rest/auth_tls.go b/v1/plugins/rest/auth_tls.go new file mode 100644 index 0000000000..8f0bb37a3b --- /dev/null +++ b/v1/plugins/rest/auth_tls.go @@ -0,0 +1,250 @@ +// Copyright 2026 The OPA Authors. All rights reserved. +// Use of this source code is governed by an Apache2 +// license that can be found in the LICENSE file. + +package rest + +import ( + "cmp" + "crypto/rsa" + "crypto/sha256" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "errors" + "fmt" + "net/http" + "net/url" + "os" + "sync" + "time" + + "github.com/open-policy-agent/opa/v1/config" +) + +// DefaultTLSConfig defines standard TLS configurations based on the Config +func DefaultTLSConfig(c Config) (*tls.Config, error) { + t := &tls.Config{ + MinVersion: cmp.Or(c.minTLSVersion, uint16(config.DefaultMinTLSVersion)), + } + + if c.cipherSuites != nil { + t.CipherSuites = *c.cipherSuites + } + + url, err := url.Parse(c.URL) + if err != nil { + return nil, err + } + if url.Scheme == "https" { + t.InsecureSkipVerify = c.AllowInsecureTLS + } + + if c.TLS != nil && c.TLS.CACert != "" { + caCert, err := os.ReadFile(c.TLS.CACert) + if err != nil { + return nil, err + } + + var rootCAs *x509.CertPool + if c.TLS.SystemCARequired { + rootCAs, err = x509.SystemCertPool() + if err != nil { + return nil, err + } + } else { + rootCAs = x509.NewCertPool() + } + + ok := rootCAs.AppendCertsFromPEM(caCert) + if !ok { + return nil, errors.New("unable to parse and append CA certificate to certificate pool") + } + t.RootCAs = rootCAs + } + + return t, nil +} + +// DefaultRoundTripperClient is a reasonable set of defaults for HTTP auth plugins +func DefaultRoundTripperClient(t *tls.Config, timeout int64) *http.Client { + // Ensure we use a http.Transport with proper settings: the zero values are not + // a good choice, as they cause leaking connections: + // https://github.com/golang/go/issues/19620 + + // copy, we don't want to alter the default client's Transport + tr := http.DefaultTransport.(*http.Transport).Clone() + tr.ResponseHeaderTimeout = time.Duration(timeout) * time.Second + tr.TLSClientConfig = t + + c := *http.DefaultClient + c.Transport = tr + return &c +} + +type serverTLSConfig struct { + CACert string `json:"ca_cert,omitempty"` + SystemCARequired bool `json:"system_ca_required,omitempty"` +} + +// clientTLSAuthPlugin represents authentication via client certificate on a TLS connection +type clientTLSAuthPlugin struct { + Cert string `json:"cert"` + PrivateKey string `json:"private_key"` + PrivateKeyPassphrase string `json:"private_key_passphrase,omitempty"` + CACert string `json:"ca_cert,omitempty"` // Deprecated: Use `services[_].tls.ca_cert` instead + SystemCARequired bool `json:"system_ca_required,omitempty"` // Deprecated: Use `services[_].tls.system_ca_required` instead + CertRereadIntervalSeconds *int64 `json:"cert_reread_interval_seconds,omitempty"` + + mu sync.RWMutex + cachedCert *tls.Certificate + certFileHash [32]byte + keyFileHash [32]byte + lastLoadTime time.Time +} + +func (ap *clientTLSAuthPlugin) loadCertificate() (*tls.Certificate, error) { + rereadIntervalSeconds := int64(0) + if ap.CertRereadIntervalSeconds != nil { + rereadIntervalSeconds = *ap.CertRereadIntervalSeconds + } + + ap.mu.RLock() + if ap.cachedCert != nil && rereadIntervalSeconds > 0 { + timeSinceLastLoad := time.Since(ap.lastLoadTime).Seconds() + if timeSinceLastLoad < float64(rereadIntervalSeconds) { + cert := ap.cachedCert + ap.mu.RUnlock() + return cert, nil + } + } + ap.mu.RUnlock() + + certPEMBlock, err := os.ReadFile(ap.Cert) + if err != nil { + return nil, fmt.Errorf("failed to read client certificate file: %w", err) + } + + keyData, err := os.ReadFile(ap.PrivateKey) + if err != nil { + return nil, fmt.Errorf("failed to read client key file: %w", err) + } + + certHash := sha256.Sum256(certPEMBlock) + keyHash := sha256.Sum256(keyData) + + ap.mu.RLock() + if ap.cachedCert != nil && ap.certFileHash == certHash && ap.keyFileHash == keyHash { + cert := ap.cachedCert + ap.mu.RUnlock() + return cert, nil + } + ap.mu.RUnlock() + + var keyPEMBlock []byte + block, _ := pem.Decode(keyData) + if block == nil { + return nil, errors.New("PEM data could not be found") + } + + // nolint: staticcheck // We don't want to forbid users from using this encryption. + if x509.IsEncryptedPEMBlock(block) { + if ap.PrivateKeyPassphrase == "" { + return nil, errors.New("client private key passphrase is needed, because the certificate is password encrypted") + } + // nolint: staticcheck // We don't want to forbid users from using this encryption. + decryptedBlock, err := x509.DecryptPEMBlock(block, []byte(ap.PrivateKeyPassphrase)) + if err != nil { + return nil, err + } + key, err := x509.ParsePKCS8PrivateKey(decryptedBlock) + if err != nil { + key, err = x509.ParsePKCS1PrivateKey(decryptedBlock) + if err != nil { + return nil, fmt.Errorf("private key should be a PEM or plain PKCS1 or PKCS8; parse error: %v", err) + } + } + rsaKey, ok := key.(*rsa.PrivateKey) + if !ok { + return nil, errors.New("private key is invalid") + } + keyPEMBlock = pem.EncodeToMemory( + &pem.Block{ + Type: "RSA PRIVATE KEY", + Bytes: x509.MarshalPKCS1PrivateKey(rsaKey), + }, + ) + } else { + keyPEMBlock = keyData + } + + cert, err := tls.X509KeyPair(certPEMBlock, keyPEMBlock) + if err != nil { + return nil, fmt.Errorf("failed to parse public/private key pair: %v", err) + } + + ap.mu.Lock() + ap.cachedCert = &cert + ap.certFileHash = certHash + ap.keyFileHash = keyHash + ap.lastLoadTime = time.Now() + ap.mu.Unlock() + + return &cert, nil +} + +func (ap *clientTLSAuthPlugin) NewClient(c Config) (*http.Client, error) { + tlsConfig, err := DefaultTLSConfig(c) + if err != nil { + return nil, err + } + + if ap.Cert == "" { + return nil, errors.New("client certificate is needed when client TLS is enabled") + } + if ap.PrivateKey == "" { + return nil, errors.New("private key is needed when client TLS is enabled") + } + + tlsConfig.GetClientCertificate = func(*tls.CertificateRequestInfo) (*tls.Certificate, error) { + return ap.loadCertificate() + } + + var client *http.Client + + if c.TLS != nil && c.TLS.CACert != "" { + client = DefaultRoundTripperClient(tlsConfig, *c.ResponseHeaderTimeoutSeconds) + } else { + if ap.CACert != "" { + c.logger.Warn("Deprecated 'services[_].credentials.client_tls.ca_cert' configuration specified. Use 'services[_].tls.ca_cert' instead. See https://www.openpolicyagent.org/docs/latest/configuration/#services") + caCert, err := os.ReadFile(ap.CACert) + if err != nil { + return nil, err + } + + var caCertPool *x509.CertPool + if ap.SystemCARequired { + caCertPool, err = x509.SystemCertPool() + if err != nil { + return nil, err + } + } else { + caCertPool = x509.NewCertPool() + } + + ok := caCertPool.AppendCertsFromPEM(caCert) + if !ok { + return nil, errors.New("unable to parse and append CA certificate to certificate pool") + } + tlsConfig.RootCAs = caCertPool + } + + client = DefaultRoundTripperClient(tlsConfig, *c.ResponseHeaderTimeoutSeconds) + } + + return client, nil +} + +func (*clientTLSAuthPlugin) Prepare(*http.Request) error { + return nil +} diff --git a/v1/plugins/rest/auth_tls_test.go b/v1/plugins/rest/auth_tls_test.go new file mode 100644 index 0000000000..21babff3b8 --- /dev/null +++ b/v1/plugins/rest/auth_tls_test.go @@ -0,0 +1,432 @@ +package rest + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "fmt" + "math/big" + "net/http" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/open-policy-agent/opa/v1/keys" + "github.com/open-policy-agent/opa/v1/logging" + "github.com/open-policy-agent/opa/v1/util/test" +) + +func generateTestCertificate(t *testing.T, certPath, keyPath string, serialNumber int64) { + t.Helper() + + privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("failed to generate private key: %v", err) + } + + template := x509.Certificate{ + SerialNumber: big.NewInt(serialNumber), + Subject: pkix.Name{ + Organization: []string{"Test Org"}, + CommonName: "Test Cert", + }, + NotBefore: time.Now(), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + BasicConstraintsValid: true, + } + + derBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey) + if err != nil { + t.Fatalf("failed to create certificate: %v", err) + } + + certOut, err := os.Create(certPath) + if err != nil { + t.Fatalf("failed to open cert file for writing: %v", err) + } + if err := pem.Encode(certOut, &pem.Block{Type: "CERTIFICATE", Bytes: derBytes}); err != nil { + t.Fatalf("failed to write certificate: %v", err) + } + if err := certOut.Close(); err != nil { + t.Fatalf("error closing cert file: %v", err) + } + + keyOut, err := os.Create(keyPath) + if err != nil { + t.Fatalf("failed to open key file for writing: %v", err) + } + privBytes, err := x509.MarshalECPrivateKey(privateKey) + if err != nil { + t.Fatalf("failed to marshal private key: %v", err) + } + if err := pem.Encode(keyOut, &pem.Block{Type: "EC PRIVATE KEY", Bytes: privBytes}); err != nil { + t.Fatalf("failed to write private key: %v", err) + } + if err := keyOut.Close(); err != nil { + t.Fatalf("error closing key file: %v", err) + } +} + +func TestClientTLSAuthPlugin_CertificateRotation(t *testing.T) { + tmpDir := t.TempDir() + certPath := filepath.Join(tmpDir, "cert.pem") + keyPath := filepath.Join(tmpDir, "key.pem") + + generateTestCertificate(t, certPath, keyPath, 1) + + plugin := &clientTLSAuthPlugin{ + Cert: certPath, + PrivateKey: keyPath, + } + + config := Config{ + URL: "https://example.com", + ResponseHeaderTimeoutSeconds: &[]int64{10}[0], + logger: logging.New(), + } + + client, err := plugin.NewClient(config) + if err != nil { + t.Fatalf("NewClient() failed: %v", err) + } + + transport, ok := client.Transport.(*http.Transport) + if !ok { + t.Fatal("client transport is not *http.Transport") + } + + if transport.TLSClientConfig.GetClientCertificate == nil { + t.Fatal("client transport has no GetClientCertificate") + } + cert1, err := transport.TLSClientConfig.GetClientCertificate(&tls.CertificateRequestInfo{}) + if err != nil { + t.Fatalf("first GetClientCertificate failed: %v", err) + } + + if len(cert1.Certificate) == 0 { + t.Fatal("first certificate is empty") + } + + parsedCert1, err := x509.ParseCertificate(cert1.Certificate[0]) + if err != nil { + t.Fatalf("failed to parse first certificate: %v", err) + } + + if parsedCert1.SerialNumber.Int64() != 1 { + t.Errorf("first certificate serial number = %d, want 1", parsedCert1.SerialNumber.Int64()) + } + + cert2, err := transport.TLSClientConfig.GetClientCertificate(&tls.CertificateRequestInfo{}) + if err != nil { + t.Fatalf("second GetClientCertificate failed: %v", err) + } + + parsedCert2, err := x509.ParseCertificate(cert2.Certificate[0]) + if err != nil { + t.Fatalf("failed to parse second certificate: %v", err) + } + + if parsedCert2.SerialNumber.Int64() != 1 { + t.Errorf("second certificate serial number = %d, want 1 (should be cached)", parsedCert2.SerialNumber.Int64()) + } + + generateTestCertificate(t, certPath, keyPath, 2) + + cert3, err := transport.TLSClientConfig.GetClientCertificate(&tls.CertificateRequestInfo{}) + if err != nil { + t.Fatalf("third GetClientCertificate failed: %v", err) + } + + parsedCert3, err := x509.ParseCertificate(cert3.Certificate[0]) + if err != nil { + t.Fatalf("failed to parse third certificate: %v", err) + } + + if parsedCert3.SerialNumber.Int64() != 2 { + t.Errorf("third certificate serial number = %d, want 2 (should be reloaded)", parsedCert3.SerialNumber.Int64()) + } + + if parsedCert1.SerialNumber.Cmp(parsedCert3.SerialNumber) == 0 { + t.Error("certificate was not rotated after file change") + } +} + +func TestClientTLSAuthPlugin_ConfigParsing(t *testing.T) { + tmpDir := t.TempDir() + certPath := filepath.Join(tmpDir, "cert.pem") + keyPath := filepath.Join(tmpDir, "key.pem") + + generateTestCertificate(t, certPath, keyPath, 1) + + tests := []struct { + name string + buildConfig func(cert, key, ca string) string + expectSystemCARequired bool + expectError bool + }{ + { + name: "system_ca_required true", + buildConfig: func(cert, key, ca string) string { + return fmt.Sprintf(`{ + "name": "test", + "url": "https://example.com", + "credentials": { + "client_tls": { + "cert": %q, + "private_key": %q, + "system_ca_required": true + } + } + }`, cert, key) + }, + expectSystemCARequired: true, + }, + { + name: "system_ca_required false", + buildConfig: func(cert, key, ca string) string { + return fmt.Sprintf(`{ + "name": "test", + "url": "https://example.com", + "credentials": { + "client_tls": { + "cert": %q, + "private_key": %q, + "system_ca_required": false + } + } + }`, cert, key) + }, + expectSystemCARequired: false, + }, + { + name: "deprecated ca_cert field with system_ca_required", + buildConfig: func(cert, key, ca string) string { + return fmt.Sprintf(`{ + "name": "test", + "url": "https://example.com", + "credentials": { + "client_tls": { + "cert": %q, + "private_key": %q, + "ca_cert": %q, + "system_ca_required": true + } + } + }`, cert, key, ca) + }, + expectSystemCARequired: true, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + config := tc.buildConfig(certPath, keyPath, certPath) + + client, err := New([]byte(config), map[string]*keys.Config{}) + if tc.expectError { + if err == nil { + t.Fatal("expected error but got none") + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if client.config.Credentials.ClientTLS == nil { + t.Fatal("ClientTLS credentials not parsed") + } + + if client.config.Credentials.ClientTLS.SystemCARequired != tc.expectSystemCARequired { + t.Errorf("SystemCARequired = %v, want %v", + client.config.Credentials.ClientTLS.SystemCARequired, + tc.expectSystemCARequired) + } + }) + } +} + +func TestClientCert(t *testing.T) { + t.Parallel() + + ts := testServer{ + t: t, + tls: true, + expectClientCert: true, + } + ts.start() + defer ts.stop() + + files := map[string]string{ + "client.pem": string(ts.clientCertPem), + "client.key": string(ts.clientCertKey), + } + + test.WithTempFS(files, func(path string) { + certPath := filepath.Join(path, "client.pem") + keyPath := filepath.Join(path, "client.key") + + client := newTestClient(t, &ts, certPath, keyPath) + + ctx := t.Context() + if _, err := client.Do(ctx, "GET", "test"); err != nil { + t.Fatalf("Unexpected error: %v", err) + } + + // Scramble the keys in the server + ts.stop() + ts.start() + + // Ensure the keys don't work anymore, make a new client as the url will have changed + client = newTestClient(t, &ts, certPath, keyPath) + _, err := client.Do(ctx, "GET", "test") + expectedErrMsg := func(s string) bool { + switch { + case strings.Contains(s, "tls: unknown certificate authority"): + case strings.Contains(s, "tls: bad certificate"): + default: + return false + } + return true + } + if err == nil || !expectedErrMsg(err.Error()) { + t.Fatalf("Unexpected error %v", err) + } + + // Update the key files and try again.. + if err := os.WriteFile(filepath.Join(path, "client.pem"), ts.clientCertPem, 0600); err != nil { + t.Fatalf("Unexpected error: %s", err) + } + if err := os.WriteFile(filepath.Join(path, "client.key"), ts.clientCertKey, 0600); err != nil { + t.Fatalf("Unexpected error: %s", err) + } + if _, err := client.Do(ctx, "GET", "test"); err != nil { + t.Fatalf("Unexpected error: %v", err) + } + }) +} + +func TestClientCertPassword(t *testing.T) { + t.Parallel() + + ts := testServer{ + t: t, + tls: true, + expectClientCert: true, + clientCertPassword: "password", + } + ts.start() + defer ts.stop() + + files := map[string]string{ + "client.pem": string(ts.clientCertPem), + "client.key": string(ts.clientCertKey), + } + + test.WithTempFS(files, func(path string) { + certPath := filepath.Join(path, "client.pem") + keyPath := filepath.Join(path, "client.key") + + client := newTestClient(t, &ts, certPath, keyPath) + + ctx := t.Context() + if _, err := client.Do(ctx, "GET", "test"); err != nil { + t.Fatalf("Unexpected error: %v", err) + } + }) +} + +func TestClientTLSWithCustomCACert(t *testing.T) { + t.Parallel() + + ts := testServer{ + t: t, + tls: true, + expectClientCert: true, + } + ts.start() + defer ts.stop() + + files := map[string]string{ + "client.pem": string(ts.clientCertPem), + "client.key": string(ts.clientCertKey), + "ca.pem": string(ts.rootCertPEM), + } + + test.WithTempFS(files, func(path string) { + certPath := filepath.Join(path, "client.pem") + keyPath := filepath.Join(path, "client.key") + ts.caCert = filepath.Join(path, "ca.pem") + + client := newTestClient(t, &ts, certPath, keyPath) + + ctx := t.Context() + if _, err := client.Do(ctx, "GET", "test"); err != nil { + t.Fatalf("Unexpected error: %v", err) + } + }) +} + +func TestClientTLSWithCustomCACertAndSystemCA(t *testing.T) { + t.Parallel() + + ts := testServer{ + t: t, + tls: true, + expectClientCert: true, + expectSystemCA: true, + } + ts.start() + defer ts.stop() + + files := map[string]string{ + "client.pem": string(ts.clientCertPem), + "client.key": string(ts.clientCertKey), + "ca.pem": string(ts.rootCertPEM), + } + + test.WithTempFS(files, func(path string) { + certPath := filepath.Join(path, "client.pem") + keyPath := filepath.Join(path, "client.key") + ts.caCert = filepath.Join(path, "ca.pem") + + client := newTestClient(t, &ts, certPath, keyPath) + + ctx := t.Context() + if _, err := client.Do(ctx, "GET", "test"); err != nil { + t.Fatalf("Unexpected error: %v", err) + } + }) +} + +func newTestClient(t *testing.T, ts *testServer, certPath string, keypath string) *Client { + config := fmt.Sprintf(`{ + "name": "foo", + "url": %q, + "allow_insecure_tls": true, + "tls": {"ca_cert": %q, system_ca_required: %v}, + "credentials": { + "client_tls": { + "cert": %q, + "private_key": %q + } + } + }`, ts.server.URL, ts.caCert, ts.expectSystemCA, certPath, keypath) + client, err := New([]byte(config), map[string]*keys.Config{}) + if err != nil { + t.Fatalf("Unexpected error: %v", err) + } + if ts.clientCertPassword != "" { + client.Config().Credentials.ClientTLS.PrivateKeyPassphrase = ts.clientCertPassword + } + + return &client +} diff --git a/v1/plugins/rest/rest.go b/v1/plugins/rest/rest.go index 37cde035e2..5c741f4eec 100644 --- a/v1/plugins/rest/rest.go +++ b/v1/plugins/rest/rest.go @@ -16,6 +16,7 @@ import ( "net/http/httputil" "reflect" "strings" + "sync" "github.com/open-policy-agent/opa/internal/version" "github.com/open-policy-agent/opa/v1/keys" @@ -39,8 +40,13 @@ var maskedHeaderKeys = map[string]struct{}{ // An HTTPAuthPlugin represents a mechanism to construct and configure HTTP authentication for a REST service type HTTPAuthPlugin interface { - // implementations can assume NewClient will be called before Prepare + // NewClient is called once per Client instance and the result is cached. + // Implementations MUST NOT perform per-request operations here. + // All per-request authentication logic MUST be in Prepare(). NewClient(Config) (*http.Client, error) + + // Prepare is called before every HTTP request. + // Implementations should perform per-request authentication here. Prepare(*http.Request) error } @@ -61,9 +67,11 @@ type Config struct { AzureManagedIdentity *azureManagedIdentitiesAuthPlugin `json:"azure_managed_identity,omitempty"` Plugin *string `json:"plugin,omitempty"` } `json:"credentials"` - Type string `json:"type,omitempty"` - keys map[string]*keys.Config - logger logging.Logger + Type string `json:"type,omitempty"` + keys map[string]*keys.Config + logger logging.Logger + minTLSVersion uint16 + cipherSuites *[]uint16 } // Equal returns true if this client config is equal to the other. @@ -112,6 +120,14 @@ func (c *Config) AuthPlugin(lookup AuthPluginLookupFunc) (HTTPAuthPlugin, error) return candidate, nil } +// clientCache holds the cached HTTP client and related state with thread-safe initialization +type clientCache struct { + mu sync.Mutex + httpClient *http.Client + authPlugin HTTPAuthPlugin + initErr error +} + // Client implements an HTTP/REST client for communicating with remote // services. type Client struct { @@ -123,6 +139,7 @@ type Client struct { logger logging.Logger loggerFields map[string]any distributedTacingOpts tracing.Options + cache *clientCache } // Name returns an option that overrides the service name on the client. @@ -156,6 +173,20 @@ func DistributedTracingOpts(tr tracing.Options) func(*Client) { } } +// MinTLSVersion sets the minimum TLS version for the client +func MinTLSVersion(v uint16) func(*Client) { + return func(c *Client) { + c.config.minTLSVersion = v + } +} + +// CipherSuites sets the cipher suites for the client +func CipherSuites(cs *[]uint16) func(*Client) { + return func(c *Client) { + c.config.cipherSuites = cs + } +} + // New returns a new Client for config. func New(config []byte, keys map[string]*keys.Config, opts ...func(*Client)) (Client, error) { var parsedConfig Config @@ -174,6 +205,7 @@ func New(config []byte, keys map[string]*keys.Config, opts ...func(*Client)) (Cl client := Client{ config: parsedConfig, + cache: &clientCache{}, } for _, f := range opts { @@ -189,6 +221,36 @@ func New(config []byte, keys map[string]*keys.Config, opts ...func(*Client)) (Cl return client, nil } +func (c *Client) ensureHTTPClient() error { + c.cache.mu.Lock() + defer c.cache.mu.Unlock() + + if c.cache.httpClient != nil && c.cache.authPlugin != nil { + return c.cache.initErr + } + + plugin, err := c.config.AuthPlugin(c.authPluginLookup) + if err != nil { + c.cache.initErr = err + return err + } + + hc, err := plugin.NewClient(c.config) + if err != nil { + c.cache.initErr = err + return err + } + + if len(c.distributedTacingOpts) > 0 { + hc.Transport = tracing.NewTransport(hc.Transport, c.distributedTacingOpts) + } + + c.cache.httpClient = hc + c.cache.authPlugin = plugin + c.cache.initErr = nil + return nil +} + // AuthPluginLookup returns the lookup function to find a custom registered // auth plugin by its name. func (c Client) AuthPluginLookup() AuthPluginLookupFunc { @@ -208,6 +270,7 @@ func (c Client) Config() *Config { // SetResponseHeaderTimeout sets the "ResponseHeaderTimeout" in the http client's Transport func (c Client) SetResponseHeaderTimeout(timeout *int64) Client { c.config.ResponseHeaderTimeoutSeconds = timeout + c.cache = &clientCache{} return c } @@ -253,20 +316,10 @@ func (c Client) WithBytes(body []byte) Client { // Do executes a request using the client. func (c Client) Do(ctx context.Context, method, path string) (*http.Response, error) { - plugin, err := c.config.AuthPlugin(c.authPluginLookup) - if err != nil { + if err := c.ensureHTTPClient(); err != nil { return nil, err } - hc, err := plugin.NewClient(c.config) - if err != nil { - return nil, err - } - - if len(c.distributedTacingOpts) > 0 { - hc.Transport = tracing.NewTransport(hc.Transport, c.distributedTacingOpts) - } - path = strings.Trim(path, "/") var body io.Reader @@ -296,7 +349,12 @@ func (c Client) Do(ctx context.Context, method, path string) (*http.Response, er req.Header.Add(key, value) } - if err := plugin.Prepare(req); err != nil { + c.cache.mu.Lock() + authPlugin := c.cache.authPlugin + httpClient := c.cache.httpClient + c.cache.mu.Unlock() + + if err := authPlugin.Prepare(req); err != nil { return nil, err } @@ -310,7 +368,7 @@ func (c Client) Do(ctx context.Context, method, path string) (*http.Response, er c.logger.WithFields(c.loggerFields).Debug("Sending request.") } - resp, err := hc.Do(req) + resp, err := httpClient.Do(req) if resp != nil && c.logger.GetLevel() >= logging.Debug { // Only log for debug purposes. If an error occurred, the caller should handle diff --git a/v1/plugins/rest/rest_test.go b/v1/plugins/rest/rest_test.go index e2e1000d30..ea6aed47bc 100644 --- a/v1/plugins/rest/rest_test.go +++ b/v1/plugins/rest/rest_test.go @@ -1409,159 +1409,6 @@ func newTestBearerClient(t *testing.T, ts *testServer, tokenPath string) *Client return &client } -func TestClientCert(t *testing.T) { - t.Parallel() - - ts := testServer{ - t: t, - tls: true, - expectClientCert: true, - } - ts.start() - defer ts.stop() - - files := map[string]string{ - "client.pem": string(ts.clientCertPem), - "client.key": string(ts.clientCertKey), - } - - test.WithTempFS(files, func(path string) { - certPath := filepath.Join(path, "client.pem") - keyPath := filepath.Join(path, "client.key") - - client := newTestClient(t, &ts, certPath, keyPath) - - ctx := t.Context() - if _, err := client.Do(ctx, "GET", "test"); err != nil { - t.Fatalf("Unexpected error: %v", err) - } - - // Scramble the keys in the server - ts.stop() - ts.start() - - // Ensure the keys don't work anymore, make a new client as the url will have changed - client = newTestClient(t, &ts, certPath, keyPath) - _, err := client.Do(ctx, "GET", "test") - expectedErrMsg := func(s string) bool { - switch { - case strings.Contains(s, "tls: unknown certificate authority"): - case strings.Contains(s, "tls: bad certificate"): - default: - return false - } - return true - } - if err == nil || !expectedErrMsg(err.Error()) { - t.Fatalf("Unexpected error %v", err) - } - - // Update the key files and try again.. - if err := os.WriteFile(filepath.Join(path, "client.pem"), ts.clientCertPem, 0600); err != nil { - t.Fatalf("Unexpected error: %s", err) - } - if err := os.WriteFile(filepath.Join(path, "client.key"), ts.clientCertKey, 0600); err != nil { - t.Fatalf("Unexpected error: %s", err) - } - if _, err := client.Do(ctx, "GET", "test"); err != nil { - t.Fatalf("Unexpected error: %v", err) - } - }) -} - -func TestClientCertPassword(t *testing.T) { - t.Parallel() - - ts := testServer{ - t: t, - tls: true, - expectClientCert: true, - clientCertPassword: "password", - } - ts.start() - defer ts.stop() - - files := map[string]string{ - "client.pem": string(ts.clientCertPem), - "client.key": string(ts.clientCertKey), - } - - test.WithTempFS(files, func(path string) { - certPath := filepath.Join(path, "client.pem") - keyPath := filepath.Join(path, "client.key") - - client := newTestClient(t, &ts, certPath, keyPath) - - ctx := t.Context() - if _, err := client.Do(ctx, "GET", "test"); err != nil { - t.Fatalf("Unexpected error: %v", err) - } - }) -} - -func TestClientTLSWithCustomCACert(t *testing.T) { - t.Parallel() - - ts := testServer{ - t: t, - tls: true, - expectClientCert: true, - } - ts.start() - defer ts.stop() - - files := map[string]string{ - "client.pem": string(ts.clientCertPem), - "client.key": string(ts.clientCertKey), - "ca.pem": string(ts.rootCertPEM), - } - - test.WithTempFS(files, func(path string) { - certPath := filepath.Join(path, "client.pem") - keyPath := filepath.Join(path, "client.key") - ts.caCert = filepath.Join(path, "ca.pem") - - client := newTestClient(t, &ts, certPath, keyPath) - - ctx := t.Context() - if _, err := client.Do(ctx, "GET", "test"); err != nil { - t.Fatalf("Unexpected error: %v", err) - } - }) -} - -func TestClientTLSWithCustomCACertAndSystemCA(t *testing.T) { - t.Parallel() - - ts := testServer{ - t: t, - tls: true, - expectClientCert: true, - expectSystemCA: true, - } - ts.start() - defer ts.stop() - - files := map[string]string{ - "client.pem": string(ts.clientCertPem), - "client.key": string(ts.clientCertKey), - "ca.pem": string(ts.rootCertPEM), - } - - test.WithTempFS(files, func(path string) { - certPath := filepath.Join(path, "client.pem") - keyPath := filepath.Join(path, "client.key") - ts.caCert = filepath.Join(path, "ca.pem") - - client := newTestClient(t, &ts, certPath, keyPath) - - ctx := t.Context() - if _, err := client.Do(ctx, "GET", "test"); err != nil { - t.Fatalf("Unexpected error: %v", err) - } - }) -} - func TestOauth2ClientCredentials(t *testing.T) { t.Parallel() @@ -2125,30 +1972,6 @@ func TestDebugLoggingRequestMaskAuthorizationHeader(t *testing.T) { } } -func newTestClient(t *testing.T, ts *testServer, certPath string, keypath string) *Client { - config := fmt.Sprintf(`{ - "name": "foo", - "url": %q, - "allow_insecure_tls": true, - "tls": {"ca_cert": %q, system_ca_required: %v}, - "credentials": { - "client_tls": { - "cert": %q, - "private_key": %q - } - } - }`, ts.server.URL, ts.caCert, ts.expectSystemCA, certPath, keypath) - client, err := New([]byte(config), map[string]*keys.Config{}) - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } - if ts.clientCertPassword != "" { - client.Config().Credentials.ClientTLS.PrivateKeyPassphrase = ts.clientCertPassword - } - - return &client -} - type testPluginCustomizer func(c *Config) type testServer struct { diff --git a/v1/runtime/runtime.go b/v1/runtime/runtime.go index 72b0275847..f2325e6734 100644 --- a/v1/runtime/runtime.go +++ b/v1/runtime/runtime.go @@ -516,6 +516,8 @@ func NewRuntime(ctx context.Context, params Params) (*Runtime, error) { plugins.WithDistributedTracingOpts(params.DistributedTracingOpts), plugins.WithBundleActivatorPlugin(params.BundleActivatorPlugin), plugins.WithHooks(params.Hooks), + plugins.WithMinTLSVersion(params.MinTLSVersion), + plugins.WithCipherSuites(params.CipherSuites), ) if err != nil { return nil, fmt.Errorf("config error: %w", err) diff --git a/v1/server/server.go b/v1/server/server.go index 6de543abac..6aaf529eef 100644 --- a/v1/server/server.go +++ b/v1/server/server.go @@ -34,6 +34,7 @@ import ( "github.com/open-policy-agent/opa/internal/json/patch" "github.com/open-policy-agent/opa/v1/ast" "github.com/open-policy-agent/opa/v1/bundle" + "github.com/open-policy-agent/opa/v1/config" "github.com/open-policy-agent/opa/v1/hooks" "github.com/open-policy-agent/opa/v1/logging" "github.com/open-policy-agent/opa/v1/metrics" @@ -80,8 +81,6 @@ const ( ) const ( - defaultMinTLSVersion = tls.VersionTLS12 - // Set of handlers for use in the "handler" dimension of the duration metric. PromHandlerV0Data = "v0/data" PromHandlerV1Data = "v1/data" @@ -416,7 +415,7 @@ func (s *Server) WithMinTLSVersion(minTLSVersion uint16) *Server { if slices.Contains(supportedTLSVersions, minTLSVersion) { s.minTLSVersion = minTLSVersion } else { - s.minTLSVersion = defaultMinTLSVersion + s.minTLSVersion = config.DefaultMinTLSVersion } return s } @@ -691,7 +690,7 @@ func (s *Server) getListenerForHTTPSServer(u *url.URL, h http.Handler, t httpLis if s.minTLSVersion != 0 { cfg.MinVersion = s.minTLSVersion } else { - cfg.MinVersion = defaultMinTLSVersion + cfg.MinVersion = config.DefaultMinTLSVersion } if s.cipherSuites != nil {