diff --git a/http.go b/http.go index 32fd62362..2cb38d4b8 100644 --- a/http.go +++ b/http.go @@ -11,8 +11,6 @@ import ( "k8s.io/apiserver/pkg/server/dynamiccertificates" - oscrypto "github.com/openshift/library-go/pkg/crypto" - "github.com/openshift/oauth-proxy/util" ) @@ -72,14 +70,23 @@ func (s *Server) ServeHTTP() { log.Printf("HTTP: closing %s", listener.Addr()) } -func (s *Server) ServeHTTPS(ctx context.Context) { - addr := s.Opts.HttpsAddress +func (s *Server) buildTLSConfig() *tls.Config { + config := &tls.Config{ + MinVersion: s.Opts.tlsMinVersionValue, + NextProtos: []string{"http/1.1"}, + } - config := oscrypto.SecureTLSConfig(&tls.Config{}) - if config.NextProtos == nil { - config.NextProtos = []string{"http/1.1"} + if len(s.Opts.tlsCipherSuiteIDs) > 0 { + config.CipherSuites = s.Opts.tlsCipherSuiteIDs } + return config +} + +func (s *Server) ServeHTTPS(ctx context.Context) { + addr := s.Opts.HttpsAddress + config := s.buildTLSConfig() + var err error servingCertProvider, err := dynamiccertificates.NewDynamicServingContentFromFiles("serving", s.Opts.TLSCertFile, s.Opts.TLSKeyFile) if err != nil { diff --git a/http_test.go b/http_test.go new file mode 100644 index 000000000..0d0c7ea49 --- /dev/null +++ b/http_test.go @@ -0,0 +1,75 @@ +package main + +import ( + "crypto/tls" + "testing" +) + +func TestBuildTLSConfig_Defaults(t *testing.T) { + s := &Server{Opts: &Options{tlsMinVersionValue: tls.VersionTLS13}} + config := s.buildTLSConfig() + + if config.MinVersion != tls.VersionTLS13 { + t.Errorf("Expected default MinVersion TLS 1.3 (%d), got %d", tls.VersionTLS13, config.MinVersion) + } +} + +func TestBuildTLSConfig_MinVersion(t *testing.T) { + tests := []struct { + name string + version uint16 + expected uint16 + }{ + {"TLS 1.2", tls.VersionTLS12, tls.VersionTLS12}, + {"TLS 1.3", tls.VersionTLS13, tls.VersionTLS13}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := &Server{Opts: &Options{tlsMinVersionValue: tt.version}} + config := s.buildTLSConfig() + + if config.MinVersion != tt.expected { + t.Errorf("Expected MinVersion %d, got %d", tt.expected, config.MinVersion) + } + }) + } +} + +func TestBuildTLSConfig_CipherSuites(t *testing.T) { + tests := []struct { + name string + suiteIDs []uint16 + expectCount int + }{ + { + name: "single cipher", + suiteIDs: []uint16{tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256}, + expectCount: 1, + }, + { + name: "multiple ciphers", + suiteIDs: []uint16{tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384}, + expectCount: 2, + }, + { + name: "no ciphers uses default", + suiteIDs: nil, + expectCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := &Server{Opts: &Options{ + tlsMinVersionValue: tls.VersionTLS13, + tlsCipherSuiteIDs: tt.suiteIDs, + }} + config := s.buildTLSConfig() + + if len(config.CipherSuites) != tt.expectCount { + t.Errorf("Expected %d cipher suites, got %d", tt.expectCount, len(config.CipherSuites)) + } + }) + } +} diff --git a/main.go b/main.go index 637874d70..1cd6d0aff 100644 --- a/main.go +++ b/main.go @@ -32,6 +32,7 @@ func main() { openshiftCAs := NewStringArray() clientCA := "" upstreamCAs := NewStringArray() + tlsCipherSuitesStr := "" config := flagSet.String("config", "", "path to config file") showVersion := flagSet.Bool("version", false, "print version string") @@ -42,6 +43,8 @@ func main() { flagSet.String("tls-cert", "", "path to certificate file") flagSet.String("tls-key", "", "path to private key file") flagSet.StringVar(&clientCA, "tls-client-ca", clientCA, "path to a CA file for admitting client certificates.") + flagSet.String("tls-min-version", "", "minimum TLS version (e.g., VersionTLS12, VersionTLS13). Defaults to TLS 1.3") + flagSet.StringVar(&tlsCipherSuitesStr, "tls-cipher-suites", "", "comma-separated list of TLS cipher suites") flagSet.String("redirect-url", "", "the OAuth Redirect URL. ie: \"https://internalapp.yourcompany.com/oauth/callback\"") flagSet.Bool("set-xauthrequest", false, "set X-Auth-Request-User and X-Auth-Request-Email response headers (useful in Nginx auth_request mode)") flagSet.Var(upstreams, "upstream", "the http url(s) of the upstream endpoint or file:// paths for static files. Routing is based on the path") @@ -131,6 +134,18 @@ func main() { cfg.LoadEnvForStruct(opts) options.Resolve(opts, flagSet, cfg) + // Parse comma-separated TLS cipher suites from CLI (overrides config file) + if tlsCipherSuitesStr != "" { + cipherSuites := strings.Split(tlsCipherSuitesStr, ",") + opts.TLSCipherSuites = make([]string, 0, len(cipherSuites)) + for _, cipher := range cipherSuites { + trimmed := strings.TrimSpace(cipher) + if trimmed != "" { + opts.TLSCipherSuites = append(opts.TLSCipherSuites, trimmed) + } + } + } + var p providers.Provider switch opts.Provider { case "openshift": diff --git a/options.go b/options.go index 33f3faf01..3ce51c3ed 100644 --- a/options.go +++ b/options.go @@ -9,6 +9,7 @@ import ( "net/http" "net/url" "regexp" + "sort" "strings" "time" @@ -40,6 +41,8 @@ type Options struct { TLSCertFile string `flag:"tls-cert" cfg:"tls_cert_file"` TLSKeyFile string `flag:"tls-key" cfg:"tls_key_file"` TLSClientCAFile string `flag:"tls-client-ca" cfg:"tls_client_ca"` + TLSMinVersion string `flag:"tls-min-version" cfg:"tls_min_version"` + TLSCipherSuites []string `cfg:"tls_cipher_suites"` // No flag tag - manually parsed after Resolve() AuthenticatedEmailsFile string `flag:"authenticated-emails-file" cfg:"authenticated_emails_file"` EmailDomains []string `flag:"email-domain" cfg:"email_domains"` @@ -111,12 +114,14 @@ type Options struct { Timeout time.Duration `flag:"upstream-timeout" cfg:"upstream_timeout"` // internal values that are set after config validation - redirectURL *url.URL - proxyURLs []*url.URL - CompiledAuthRegex []*regexp.Regexp - CompiledSkipRegex []*regexp.Regexp - provider providers.Provider - signatureData *SignatureData + redirectURL *url.URL + proxyURLs []*url.URL + CompiledAuthRegex []*regexp.Regexp + CompiledSkipRegex []*regexp.Regexp + provider providers.Provider + signatureData *SignatureData + tlsMinVersionValue uint16 + tlsCipherSuiteIDs []uint16 } type SignatureData struct { @@ -310,6 +315,30 @@ func (o *Options) Validate(p providers.Provider) error { msgs = append(msgs, "tls-client-ca requires tls-key-file or tls-cert-file to be set to listen on tls") } + o.tlsMinVersionValue = tls.VersionTLS13 + if o.TLSMinVersion != "" { + if v, ok := tlsVersionMap[o.TLSMinVersion]; ok { + o.tlsMinVersionValue = v + } else { + msgs = append(msgs, fmt.Sprintf("unrecognized tls-min-version %q; valid values: %s", o.TLSMinVersion, validTLSVersions())) + } + } + + if len(o.TLSCipherSuites) > 0 { + var unrecognized []string + o.tlsCipherSuiteIDs = make([]uint16, 0, len(o.TLSCipherSuites)) + for _, name := range o.TLSCipherSuites { + if id, ok := tlsCipherSuiteMap[name]; ok { + o.tlsCipherSuiteIDs = append(o.tlsCipherSuiteIDs, id) + } else { + unrecognized = append(unrecognized, name) + } + } + if len(unrecognized) > 0 { + msgs = append(msgs, fmt.Sprintf("unrecognized tls-cipher-suites: %s", strings.Join(unrecognized, ", "))) + } + } + switch o.CookieSameSite { case "", "none", "lax", "strict": default: @@ -427,3 +456,25 @@ func secretBytes(secret string) []byte { } return []byte(secret) } + +var tlsVersionMap = map[string]uint16{ + "VersionTLS12": tls.VersionTLS12, + "VersionTLS13": tls.VersionTLS13, +} + +var tlsCipherSuiteMap = func() map[string]uint16 { + m := make(map[string]uint16) + for _, cs := range tls.CipherSuites() { + m[cs.Name] = cs.ID + } + return m +}() + +func validTLSVersions() string { + versions := make([]string, 0, len(tlsVersionMap)) + for k := range tlsVersionMap { + versions = append(versions, k) + } + sort.Strings(versions) + return strings.Join(versions, ", ") +} diff --git a/options_test.go b/options_test.go index 7beb5a61b..f72cdcd5a 100644 --- a/options_test.go +++ b/options_test.go @@ -2,6 +2,7 @@ package main import ( "crypto" + "crypto/tls" "fmt" "net/url" "strings" @@ -229,3 +230,40 @@ func TestValidateCookieSameSite(t *testing.T) { }) } } + +func TestValidateTLSMinVersionDefault(t *testing.T) { + o := testOptions() + assert.Equal(t, nil, o.Validate(&testProvider{})) + assert.Equal(t, uint16(tls.VersionTLS13), o.tlsMinVersionValue) +} + +func TestValidateTLSMinVersion(t *testing.T) { + o := testOptions() + o.TLSMinVersion = "VersionTLS12" + assert.Equal(t, nil, o.Validate(&testProvider{})) + assert.Equal(t, uint16(tls.VersionTLS12), o.tlsMinVersionValue) +} + +func TestValidateTLSMinVersionInvalid(t *testing.T) { + o := testOptions() + o.TLSMinVersion = "VersionTLS99" + err := o.Validate(&testProvider{}) + assert.NotEqual(t, nil, err) + assert.Equal(t, true, strings.Contains(err.Error(), "unrecognized tls-min-version")) +} + +func TestValidateTLSCipherSuites(t *testing.T) { + o := testOptions() + o.TLSCipherSuites = []string{"TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256"} + assert.Equal(t, nil, o.Validate(&testProvider{})) + assert.Equal(t, 1, len(o.tlsCipherSuiteIDs)) + assert.Equal(t, uint16(tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256), o.tlsCipherSuiteIDs[0]) +} + +func TestValidateTLSCipherSuitesInvalid(t *testing.T) { + o := testOptions() + o.TLSCipherSuites = []string{"TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256", "BOGUS_CIPHER"} + err := o.Validate(&testProvider{}) + assert.NotEqual(t, nil, err) + assert.Equal(t, true, strings.Contains(err.Error(), "BOGUS_CIPHER")) +}