diff --git a/cmd/component/server/cmd/default_config.go b/cmd/component/server/cmd/default_config.go index afad3d4..62c5e36 100644 --- a/cmd/component/server/cmd/default_config.go +++ b/cmd/component/server/cmd/default_config.go @@ -45,10 +45,12 @@ func getDefaultServerConfig() *config.Config { Port: 8080, }, HttpsSSL: config.HttpsSSLConfig{ - Enabled: true, - BindAddress: "0.0.0.0", - Port: 8443, - CertDir: "/mc_home/certs/https_ssl", + Enabled: true, + BindAddress: "0.0.0.0", + Port: 8443, + CertDir: "/mc_home/certs/https_ssl", + ValidityDays: 365, // optional; self-signed cert lifetime in days + RenewBeforeDays: 30, // optional; regenerate when remaining validity is below this many days }, HttpsACME: config.HttpsACMEConfig{ Enabled: false, diff --git a/pkg/service/http_listener/https/ssl.go b/pkg/service/http_listener/https/ssl.go index ecc2dac..3ffd94d 100644 --- a/pkg/service/http_listener/https/ssl.go +++ b/pkg/service/http_listener/https/ssl.go @@ -11,9 +11,13 @@ import ( "errors" "fmt" "math/big" + "os" + "path/filepath" + "sync" "time" "github.com/mycontroller-org/server/v2/pkg/types/config" + schedulerTY "github.com/mycontroller-org/server/v2/pkg/types/scheduler" "github.com/mycontroller-org/server/v2/pkg/utils" "go.uber.org/zap" ) @@ -24,52 +28,299 @@ const ( RSABits = 2048 OrganizationName = "MyController.org" - Validity = 365 // days + ValidityDays = 365 // days - default generated certificate lifetime + RenewBeforeDays = 30 // days - default regenerate when remaining validity is less than this GeneratedCertFileName = "mc_generated.crt" GeneratedKeyFileName = "mc_generated.key" + + // daily renewal check for MyController-managed self-signed certificates + sslRenewalJobName = "https_ssl_cert_renewal" + sslRenewalCron = "@every 24h" ) -// GetSSLTLSConfig returns ssl certificate -func GetSSLTLSConfig(logger *zap.Logger, cfg config.HttpsSSLConfig) (*tls.Config, error) { +// SSLManager loads HTTPS certificates and, when MyController manages them, +// periodically renews the cert (reusing the private key) and hot-reloads it. +type SSLManager struct { + logger *zap.Logger + cfg config.HttpsSSLConfig + managed bool + mu sync.RWMutex + cert *tls.Certificate + + scheduler schedulerTY.CoreScheduler +} + +// NewSSLManager prepares certificates for HTTPS/SSL. +// Managed (auto-generated) certs are renewed on a daily schedule when a scheduler is provided. +// Custom certs (custom.crt + custom.key) are loaded as-is and never auto-renewed. +func NewSSLManager(logger *zap.Logger, cfg config.HttpsSSLConfig, scheduler schedulerTY.CoreScheduler) (*SSLManager, error) { if cfg.CertDir == "" { return nil, errors.New("cert_dir is missing") } - certFile := fmt.Sprintf("%s/%s", cfg.CertDir, GeneratedCertFileName) - keyFile := fmt.Sprintf("%s/%s", cfg.CertDir, GeneratedKeyFileName) + m := &SSLManager{ + logger: logger, + cfg: cfg, + managed: isManagedByMyController(cfg), + scheduler: scheduler, + } + + // only relevant for MyController-managed self-signed certificates + if m.managed { + validityDays := resolveValidityDays(cfg.ValidityDays) + renewBeforeDays := resolveRenewBeforeDays(cfg.RenewBeforeDays) + if validityDays <= renewBeforeDays { + logger.Warn("invalid SSL self-signed certificate configuration: validity_days must be greater than renew_before_days", + zap.Int("validityDays", validityDays), zap.Int("renewBeforeDays", renewBeforeDays), + zap.String("hint", "increase validity_days or decrease renew_before_days; otherwise a newly generated certificate remains below the renewal threshold and will be regenerated on every check"), + ) + } + } - // check the certificate on disk, if available use it and skip the following steps - customCertFile := fmt.Sprintf("%s/%s", cfg.CertDir, CustomCertFileName) - customKeyFile := fmt.Sprintf("%s/%s", cfg.CertDir, CustomKeyFileName) + if err := m.loadOrCreate(); err != nil { + return nil, err + } - if utils.IsFileExists(customCertFile) && utils.IsFileExists(customKeyFile) { - certFile = customCertFile - keyFile = customKeyFile - } else { // generate certificate - err := generateSSLCert(logger, cfg.CertDir) - if err != nil { - return nil, err + return m, nil +} + +// TLSConfig returns a tls.Config that serves the current certificate and +// picks up renewals without restarting the listener. +func (m *SSLManager) TLSConfig() *tls.Config { + return &tls.Config{ + GetCertificate: m.getCertificate, + } +} + +// Managed reports whether certificates are generated/renewed by MyController. +func (m *SSLManager) Managed() bool { + return m.managed +} + +// StartDailyRenewalCheck schedules a once-per-day renewal check when SSL is +// managed by MyController. No-op for custom certificates or when scheduler is nil. +func (m *SSLManager) StartDailyRenewalCheck() error { + if !m.managed { + m.logger.Debug("SSL certificate is custom; daily renewal check is disabled") + return nil + } + if m.scheduler == nil { + m.logger.Warn("core scheduler not available; SSL daily renewal check is disabled") + return nil + } + + err := m.scheduler.AddFunc(sslRenewalJobName, sslRenewalCron, m.dailyRenewalCheck) + if err != nil { + return fmt.Errorf("schedule SSL daily renewal check: %w", err) + } + m.logger.Info("scheduled daily SSL certificate renewal check", zap.String("job", sslRenewalJobName), zap.String("cron", sslRenewalCron), + zap.Int("renewBeforeDays", resolveRenewBeforeDays(m.cfg.RenewBeforeDays)), zap.Int("validityDays", resolveValidityDays(m.cfg.ValidityDays))) + return nil +} + +// Close stops the daily renewal check if it was scheduled. +func (m *SSLManager) Close() error { + if m.scheduler != nil && m.managed { + m.scheduler.RemoveFunc(sslRenewalJobName) + } + return nil +} + +// CheckAndRenew regenerates the managed certificate when remaining validity is +// below the configured threshold. Safe to call from the daily job or tests. +func (m *SSLManager) CheckAndRenew() error { + if !m.managed { + return nil + } + + certFile := filepath.Join(m.cfg.CertDir, GeneratedCertFileName) + keyFile := filepath.Join(m.cfg.CertDir, GeneratedKeyFileName) + validityDays := resolveValidityDays(m.cfg.ValidityDays) + renewBeforeDays := resolveRenewBeforeDays(m.cfg.RenewBeforeDays) + + if !shouldRegenerateCert(m.logger, certFile, keyFile, renewBeforeDays) { + m.logger.Debug("SSL certificate still valid; no renewal needed", zap.String("certFile", certFile), zap.Int("renewBeforeDays", renewBeforeDays)) + return nil + } + + if err := generateSSLCert(m.logger, m.cfg.CertDir, validityDays); err != nil { + return err + } + + certificate, err := tls.LoadX509KeyPair(certFile, keyFile) + if err != nil { + return err + } + + m.mu.Lock() + m.cert = &certificate + m.mu.Unlock() + + m.logger.Info("SSL certificate renewed and hot-reloaded") + return nil +} + +func (m *SSLManager) getCertificate(_ *tls.ClientHelloInfo) (*tls.Certificate, error) { + m.mu.RLock() + defer m.mu.RUnlock() + if m.cert == nil { + return nil, errors.New("SSL certificate not loaded") + } + return m.cert, nil +} + +func (m *SSLManager) dailyRenewalCheck() { + if err := m.CheckAndRenew(); err != nil { + m.logger.Error("error on daily SSL certificate renewal check", zap.Error(err)) + } +} + +func (m *SSLManager) loadOrCreate() error { + var certFile, keyFile string + + if !m.managed { + certFile = filepath.Join(m.cfg.CertDir, CustomCertFileName) + keyFile = filepath.Join(m.cfg.CertDir, CustomKeyFileName) + m.logger.Info("using custom SSL certificate", zap.String("certFile", certFile), zap.String("keyFile", keyFile)) + } else { + certFile = filepath.Join(m.cfg.CertDir, GeneratedCertFileName) + keyFile = filepath.Join(m.cfg.CertDir, GeneratedKeyFileName) + validityDays := resolveValidityDays(m.cfg.ValidityDays) + renewBeforeDays := resolveRenewBeforeDays(m.cfg.RenewBeforeDays) + + if shouldRegenerateCert(m.logger, certFile, keyFile, renewBeforeDays) { + if err := generateSSLCert(m.logger, m.cfg.CertDir, validityDays); err != nil { + return err + } + } else { + m.logger.Info("using existing generated SSL certificate", zap.String("certFile", certFile), zap.String("keyFile", keyFile)) } } - tlsConfig := &tls.Config{Certificates: make([]tls.Certificate, 1)} certificate, err := tls.LoadX509KeyPair(certFile, keyFile) + if err != nil { + return err + } + + m.mu.Lock() + m.cert = &certificate + m.mu.Unlock() + return nil +} + +// isManagedByMyController is true when no operator-supplied custom cert/key pair is present. +func isManagedByMyController(cfg config.HttpsSSLConfig) bool { + customCertFile := filepath.Join(cfg.CertDir, CustomCertFileName) + customKeyFile := filepath.Join(cfg.CertDir, CustomKeyFileName) + return !utils.IsFileExists(customCertFile) || !utils.IsFileExists(customKeyFile) +} + +// GetSSLTLSConfig returns ssl certificate (startup-only path without daily renewal). +// Prefer NewSSLManager when the HTTPS listener needs hot-reload renewal. +func GetSSLTLSConfig(logger *zap.Logger, cfg config.HttpsSSLConfig) (*tls.Config, error) { + m, err := NewSSLManager(logger, cfg, nil) if err != nil { return nil, err } + return m.TLSConfig(), nil +} + +// resolveValidityDays returns cfg value when positive, otherwise the default lifetime. +func resolveValidityDays(days int) int { + if days > 0 { + return days + } + return ValidityDays +} + +// resolveRenewBeforeDays returns cfg value when positive, otherwise the default threshold. +func resolveRenewBeforeDays(days int) int { + if days > 0 { + return days + } + return RenewBeforeDays +} + +// shouldRegenerateCert returns true when the generated cert/key are missing, +// cannot be parsed, or remaining validity is less than renewBeforeDays. +func shouldRegenerateCert(logger *zap.Logger, certFile, keyFile string, renewBeforeDays int) bool { + if !utils.IsFileExists(certFile) || !utils.IsFileExists(keyFile) { + logger.Info("generated SSL certificate not found, will create a new one", zap.String("certFile", certFile), zap.String("keyFile", keyFile)) + return true + } + + remaining, err := certificateRemainingValidity(certFile) + if err != nil { + logger.Warn("unable to read existing SSL certificate, will create a new one", zap.String("certFile", certFile), zap.Error(err)) + return true + } - tlsConfig.Certificates[0] = certificate + renewBefore := time.Duration(renewBeforeDays) * 24 * time.Hour + if remaining < renewBefore { + logger.Info("SSL certificate remaining validity is below threshold, will regenerate", zap.String("certFile", certFile), + zap.Duration("remaining", remaining), zap.Int("renewBeforeDays", renewBeforeDays)) + return true + } - return tlsConfig, nil + return false } -// generateSSLCert generates ssl certificate and key -func generateSSLCert(logger *zap.Logger, certDir string) error { - // check the certificate on disk, if available use it and skip the following steps +// certificateRemainingValidity returns the time left until the certificate expires. +func certificateRemainingValidity(certFile string) (time.Duration, error) { + certPEM, err := os.ReadFile(certFile) + if err != nil { + return 0, err + } + + block, _ := pem.Decode(certPEM) + if block == nil { + return 0, errors.New("failed to decode PEM certificate") + } + + cert, err := x509.ParseCertificate(block.Bytes) + if err != nil { + return 0, err + } + + return time.Until(cert.NotAfter), nil +} + +// loadOrCreatePrivateKey reuses an existing RSA key when present and valid; otherwise creates one. +func loadOrCreatePrivateKey(logger *zap.Logger, certDir string) (*rsa.PrivateKey, bool, error) { + keyPath := filepath.Join(certDir, GeneratedKeyFileName) + if utils.IsFileExists(keyPath) { + keyPEM, err := os.ReadFile(keyPath) + if err == nil { + block, _ := pem.Decode(keyPEM) + if block != nil { + if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil { + if rsaKey, ok := key.(*rsa.PrivateKey); ok { + logger.Info("reusing existing SSL private key", zap.String("keyFile", GeneratedKeyFileName)) + return rsaKey, true, nil + } + } + // older OpenSSL-style PKCS#1 keys + if rsaKey, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil { + logger.Info("reusing existing SSL private key", zap.String("keyFile", GeneratedKeyFileName)) + return rsaKey, true, nil + } + } + } + logger.Warn("unable to reuse existing SSL private key, will generate a new one", zap.String("keyFile", keyPath)) + } - // reference: https://golang.org/src/crypto/tls/generate_cert.go privateKey, err := rsa.GenerateKey(rand.Reader, RSABits) + if err != nil { + return nil, false, err + } + return privateKey, false, nil +} + +// generateSSLCert generates (or renews) the SSL certificate, reusing the private key when possible. +func generateSSLCert(logger *zap.Logger, certDir string, validityDays int) error { + // reference: https://golang.org/src/crypto/tls/generate_cert.go + privateKey, keyReused, err := loadOrCreatePrivateKey(logger, certDir) if err != nil { return err } @@ -81,11 +332,14 @@ func generateSSLCert(logger *zap.Logger, certDir string) error { return err } + notBefore := time.Now() + notAfter := notBefore.Add(time.Hour * 24 * time.Duration(validityDays)) + template := x509.Certificate{ SerialNumber: serialNumber, Subject: pkix.Name{Organization: []string{OrganizationName}}, - NotBefore: time.Now(), - NotAfter: time.Now().Add(time.Hour * 24 * Validity), + NotBefore: notBefore, + NotAfter: notAfter, KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, BasicConstraintsValid: true, @@ -106,19 +360,24 @@ func generateSSLCert(logger *zap.Logger, certDir string) error { return err } - var keyBuf bytes.Buffer - privBytes, err := x509.MarshalPKCS8PrivateKey(privateKey) - if err != nil { - return err - } - if err := pem.Encode(&keyBuf, &pem.Block{Type: "PRIVATE KEY", Bytes: privBytes}); err != nil { - return err - } + if !keyReused { + var keyBuf bytes.Buffer + privBytes, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + return err + } + if err := pem.Encode(&keyBuf, &pem.Block{Type: "PRIVATE KEY", Bytes: privBytes}); err != nil { + return err + } - err = utils.WriteFile(certDir, GeneratedKeyFileName, keyBuf.Bytes()) - if err != nil { - return err + err = utils.WriteFile(certDir, GeneratedKeyFileName, keyBuf.Bytes()) + if err != nil { + return err + } } + logger.Info("generated new SSL certificate", zap.String("certDir", certDir), zap.String("certFile", GeneratedCertFileName), zap.String("keyFile", GeneratedKeyFileName), + zap.Bool("keyReused", keyReused), zap.Time("notBefore", notBefore), zap.Time("notAfter", notAfter), zap.Int("validityDays", validityDays)) + return nil } diff --git a/pkg/service/http_listener/https/ssl_test.go b/pkg/service/http_listener/https/ssl_test.go new file mode 100644 index 0000000..fb0b288 --- /dev/null +++ b/pkg/service/http_listener/https/ssl_test.go @@ -0,0 +1,640 @@ +package https + +import ( + "bytes" + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/mycontroller-org/server/v2/pkg/types/config" + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func TestShouldRegenerateCert_MissingFiles(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + certFile := filepath.Join(dir, GeneratedCertFileName) + keyFile := filepath.Join(dir, GeneratedKeyFileName) + + if !shouldRegenerateCert(logger, certFile, keyFile, RenewBeforeDays) { + t.Fatal("expected regeneration when cert and key files are missing") + } +} + +func TestShouldRegenerateCert_ValidCert(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + certFile, keyFile := writeTestCert(t, dir, time.Now(), time.Now().Add(365*24*time.Hour)) + + if shouldRegenerateCert(logger, certFile, keyFile, RenewBeforeDays) { + t.Fatal("expected existing long-lived cert to be reused") + } +} + +func TestShouldRegenerateCert_ExpiringSoon(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + // remaining validity: 10 days (< default RenewBeforeDays) + certFile, keyFile := writeTestCert(t, dir, time.Now().Add(-355*24*time.Hour), time.Now().Add(10*24*time.Hour)) + + if !shouldRegenerateCert(logger, certFile, keyFile, RenewBeforeDays) { + t.Fatal("expected regeneration when remaining validity is less than renew_before_days") + } +} + +func TestShouldRegenerateCert_AboveThreshold(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + // remaining validity slightly above RenewBeforeDays so clock skew does not flip the result + certFile, keyFile := writeTestCert(t, dir, time.Now().Add(-345*24*time.Hour), time.Now().Add(time.Duration(RenewBeforeDays+1)*24*time.Hour)) + + if shouldRegenerateCert(logger, certFile, keyFile, RenewBeforeDays) { + t.Fatal("expected cert to be reused when remaining validity is above the threshold") + } +} + +func TestShouldRegenerateCert_CustomRenewBeforeDays(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + // remaining validity: 15 days — regenerate only when threshold is higher than 15 + certFile, keyFile := writeTestCert(t, dir, time.Now().Add(-350*24*time.Hour), time.Now().Add(15*24*time.Hour)) + + if shouldRegenerateCert(logger, certFile, keyFile, 10) { + t.Fatal("expected cert to be reused when remaining validity is above custom threshold of 10 days") + } + if !shouldRegenerateCert(logger, certFile, keyFile, 30) { + t.Fatal("expected regeneration when remaining validity is below custom threshold of 30 days") + } +} + +func TestShouldRegenerateCert_InvalidCert(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + certFile := filepath.Join(dir, GeneratedCertFileName) + keyFile := filepath.Join(dir, GeneratedKeyFileName) + + if err := os.WriteFile(certFile, []byte("not-a-cert"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(keyFile, []byte("not-a-key"), 0o600); err != nil { + t.Fatal(err) + } + + if !shouldRegenerateCert(logger, certFile, keyFile, RenewBeforeDays) { + t.Fatal("expected regeneration when certificate cannot be parsed") + } +} + +func TestGetSSLTLSConfig_ReusesGeneratedCert(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + + cfg := config.HttpsSSLConfig{CertDir: dir} + tls1, err := GetSSLTLSConfig(logger, cfg) + if err != nil { + t.Fatalf("first GetSSLTLSConfig failed: %v", err) + } + if cert, err := tls1.GetCertificate(nil); err != nil || cert == nil { + t.Fatalf("expected GetCertificate to return a cert: %v", err) + } + + certPath := filepath.Join(dir, GeneratedCertFileName) + keyPath := filepath.Join(dir, GeneratedKeyFileName) + certBefore, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + keyBefore, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + + // second call must reuse the same files + tls2, err := GetSSLTLSConfig(logger, cfg) + if err != nil { + t.Fatalf("second GetSSLTLSConfig failed: %v", err) + } + if cert, err := tls2.GetCertificate(nil); err != nil || cert == nil { + t.Fatalf("expected GetCertificate on reuse: %v", err) + } + + certAfter, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + keyAfter, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + + if !bytes.Equal(certBefore, certAfter) { + t.Fatal("certificate was regenerated on restart; expected reuse") + } + if !bytes.Equal(keyBefore, keyAfter) { + t.Fatal("private key was regenerated on restart; expected reuse") + } +} + +func TestGetSSLTLSConfig_RegeneratesWhenExpiring(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + + writeTestCert(t, dir, time.Now().Add(-360*24*time.Hour), time.Now().Add(5*24*time.Hour)) + certPath := filepath.Join(dir, GeneratedCertFileName) + keyPath := filepath.Join(dir, GeneratedKeyFileName) + certBefore, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + keyBefore, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + + cfg := config.HttpsSSLConfig{CertDir: dir} + _, err = GetSSLTLSConfig(logger, cfg) + if err != nil { + t.Fatalf("GetSSLTLSConfig failed: %v", err) + } + + certAfter, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + keyAfter, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + if bytes.Equal(certBefore, certAfter) { + t.Fatal("expected certificate to be regenerated when remaining validity is below threshold") + } + if !bytes.Equal(keyBefore, keyAfter) { + t.Fatal("expected private key to be reused on renewal") + } + + remaining, err := certificateRemainingValidity(certPath) + if err != nil { + t.Fatal(err) + } + minExpected := time.Duration(ValidityDays-1) * 24 * time.Hour + if remaining < minExpected { + t.Fatalf("regenerated cert remaining validity too short: %v", remaining) + } +} + +func TestGetSSLTLSConfig_CustomValidityAndRenewBefore(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + + writeTestCert(t, dir, time.Now().Add(-340*24*time.Hour), time.Now().Add(25*24*time.Hour)) + certPath := filepath.Join(dir, GeneratedCertFileName) + certBefore, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + + cfg := config.HttpsSSLConfig{ + CertDir: dir, + ValidityDays: 90, + RenewBeforeDays: 30, + } + _, err = GetSSLTLSConfig(logger, cfg) + if err != nil { + t.Fatalf("GetSSLTLSConfig failed: %v", err) + } + + certAfter, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + if bytes.Equal(certBefore, certAfter) { + t.Fatal("expected regeneration with custom renew_before_days=30") + } + + remaining, err := certificateRemainingValidity(certPath) + if err != nil { + t.Fatal(err) + } + minExpected := 89 * 24 * time.Hour + maxExpected := 90*24*time.Hour + time.Hour + if remaining < minExpected || remaining > maxExpected { + t.Fatalf("expected ~90 days validity from config, got %v", remaining) + } +} + +func TestGetSSLTLSConfig_DefaultsWhenUnset(t *testing.T) { + if got := resolveValidityDays(0); got != ValidityDays { + t.Fatalf("expected default validity %d, got %d", ValidityDays, got) + } + if got := resolveValidityDays(-1); got != ValidityDays { + t.Fatalf("expected default validity for negative, got %d", got) + } + if got := resolveValidityDays(180); got != 180 { + t.Fatalf("expected custom validity 180, got %d", got) + } + if got := resolveRenewBeforeDays(0); got != RenewBeforeDays { + t.Fatalf("expected default renew_before %d, got %d", RenewBeforeDays, got) + } + if got := resolveRenewBeforeDays(7); got != 7 { + t.Fatalf("expected custom renew_before 7, got %d", got) + } +} + +func TestNewSSLManager_WarnsOnInvalidManagedConfig(t *testing.T) { + core, logs := observer.New(zap.WarnLevel) + logger := zap.New(core) + dir := t.TempDir() + + // equal + _, err := NewSSLManager(logger, config.HttpsSSLConfig{ + CertDir: dir, + ValidityDays: 30, + RenewBeforeDays: 30, + }, nil) + if err != nil { + t.Fatal(err) + } + + // renew_before > validity + _, err = NewSSLManager(logger, config.HttpsSSLConfig{ + CertDir: t.TempDir(), + ValidityDays: 15, + RenewBeforeDays: 30, + }, nil) + if err != nil { + t.Fatal(err) + } + + // valid config should not warn + _, err = NewSSLManager(logger, config.HttpsSSLConfig{ + CertDir: t.TempDir(), + ValidityDays: 365, + RenewBeforeDays: 30, + }, nil) + if err != nil { + t.Fatal(err) + } + + warnCount := 0 + for _, e := range logs.All() { + if e.Level == zap.WarnLevel { + warnCount++ + } + } + if warnCount < 2 { + t.Fatalf("expected at least 2 warnings for invalid configs, got %d", warnCount) + } +} + +func TestGenerateSSLCert_ValidityOneYear(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + + if err := generateSSLCert(logger, dir, ValidityDays); err != nil { + t.Fatalf("generateSSLCert failed: %v", err) + } + + certPath := filepath.Join(dir, GeneratedCertFileName) + remaining, err := certificateRemainingValidity(certPath) + if err != nil { + t.Fatal(err) + } + + minExpected := time.Duration(ValidityDays-1) * 24 * time.Hour + maxExpected := time.Duration(ValidityDays)*24*time.Hour + time.Hour + if remaining < minExpected || remaining > maxExpected { + t.Fatalf("expected ~%d days validity, got %v", ValidityDays, remaining) + } +} + +func TestGenerateSSLCert_CustomValidity(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + const days = 30 + + if err := generateSSLCert(logger, dir, days); err != nil { + t.Fatalf("generateSSLCert failed: %v", err) + } + + certPath := filepath.Join(dir, GeneratedCertFileName) + remaining, err := certificateRemainingValidity(certPath) + if err != nil { + t.Fatal(err) + } + + minExpected := time.Duration(days-1) * 24 * time.Hour + maxExpected := time.Duration(days)*24*time.Hour + time.Hour + if remaining < minExpected || remaining > maxExpected { + t.Fatalf("expected ~%d days validity, got %v", days, remaining) + } +} + +func TestGenerateSSLCert_ReusesPrivateKey(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + + if err := generateSSLCert(logger, dir, ValidityDays); err != nil { + t.Fatal(err) + } + keyPath := filepath.Join(dir, GeneratedKeyFileName) + keyBefore, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + + if err := generateSSLCert(logger, dir, ValidityDays); err != nil { + t.Fatal(err) + } + keyAfter, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(keyBefore, keyAfter) { + t.Fatal("expected private key file to be unchanged on renew") + } +} + +func TestSSLManager_ManagedVsCustom(t *testing.T) { + logger := zap.NewNop() + + // managed: no custom files + managedDir := t.TempDir() + m, err := NewSSLManager(logger, config.HttpsSSLConfig{CertDir: managedDir}, nil) + if err != nil { + t.Fatal(err) + } + if !m.Managed() { + t.Fatal("expected managed certificate when custom files are absent") + } + + // custom: both custom.crt and custom.key present + customDir := t.TempDir() + writeCustomCert(t, customDir) + c, err := NewSSLManager(logger, config.HttpsSSLConfig{CertDir: customDir}, nil) + if err != nil { + t.Fatal(err) + } + if c.Managed() { + t.Fatal("expected custom certificate when custom.crt and custom.key exist") + } +} + +func TestSSLManager_CheckAndRenew_HotReload(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + + // seed an expiring managed cert + writeTestCert(t, dir, time.Now().Add(-360*24*time.Hour), time.Now().Add(5*24*time.Hour)) + certPath := filepath.Join(dir, GeneratedCertFileName) + keyPath := filepath.Join(dir, GeneratedKeyFileName) + keyBefore, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + + m, err := NewSSLManager(logger, config.HttpsSSLConfig{CertDir: dir}, nil) + if err != nil { + t.Fatal(err) + } + if !m.Managed() { + t.Fatal("expected managed SSL") + } + + // force another near-expiry and call CheckAndRenew again + writeTestCert(t, dir, time.Now().Add(-360*24*time.Hour), time.Now().Add(5*24*time.Hour)) + // restore original key so renew reuses it (writeTestCert overwrites key) + if err := os.WriteFile(keyPath, keyBefore, 0o600); err != nil { + t.Fatal(err) + } + + beforeReload, err := m.TLSConfig().GetCertificate(nil) + if err != nil { + t.Fatal(err) + } + + if err := m.CheckAndRenew(); err != nil { + t.Fatalf("CheckAndRenew failed: %v", err) + } + + afterReload, err := m.TLSConfig().GetCertificate(nil) + if err != nil { + t.Fatal(err) + } + if bytes.Equal(beforeReload.Certificate[0], afterReload.Certificate[0]) { + t.Fatal("expected in-memory certificate to be hot-reloaded after CheckAndRenew") + } + + remaining, err := certificateRemainingValidity(certPath) + if err != nil { + t.Fatal(err) + } + if remaining < time.Duration(ValidityDays-1)*24*time.Hour { + t.Fatalf("expected renewed long-lived cert, remaining=%v", remaining) + } + + keyAfter, err := os.ReadFile(keyPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(keyBefore, keyAfter) { + t.Fatal("expected key reuse on CheckAndRenew") + } +} + +func TestSSLManager_CheckAndRenew_SkipsWhenValid(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + cfg := config.HttpsSSLConfig{CertDir: dir} + + m, err := NewSSLManager(logger, cfg, nil) + if err != nil { + t.Fatal(err) + } + + certPath := filepath.Join(dir, GeneratedCertFileName) + certBefore, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + + if err := m.CheckAndRenew(); err != nil { + t.Fatal(err) + } + + certAfter, err := os.ReadFile(certPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(certBefore, certAfter) { + t.Fatal("did not expect renewal when certificate is still valid") + } +} + +func TestSSLManager_StartDailyRenewalCheck_Managed(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + sched := &fakeScheduler{} + + m, err := NewSSLManager(logger, config.HttpsSSLConfig{CertDir: dir}, sched) + if err != nil { + t.Fatal(err) + } + if err := m.StartDailyRenewalCheck(); err != nil { + t.Fatal(err) + } + if !sched.has(sslRenewalJobName) { + t.Fatal("expected daily renewal job to be scheduled for managed SSL") + } + if err := m.Close(); err != nil { + t.Fatal(err) + } + if sched.has(sslRenewalJobName) { + t.Fatal("expected daily renewal job to be removed on Close") + } +} + +func TestSSLManager_StartDailyRenewalCheck_CustomDisabled(t *testing.T) { + logger := zap.NewNop() + dir := t.TempDir() + writeCustomCert(t, dir) + sched := &fakeScheduler{} + + m, err := NewSSLManager(logger, config.HttpsSSLConfig{CertDir: dir}, sched) + if err != nil { + t.Fatal(err) + } + if err := m.StartDailyRenewalCheck(); err != nil { + t.Fatal(err) + } + if sched.has(sslRenewalJobName) { + t.Fatal("expected no daily renewal job for custom certificates") + } +} + +// fakeScheduler records scheduled jobs for unit tests. +type fakeScheduler struct { + mu sync.Mutex + jobs map[string]string +} + +func (f *fakeScheduler) Name() string { return "fake" } +func (f *fakeScheduler) Start() error { return nil } +func (f *fakeScheduler) Close() error { return nil } +func (f *fakeScheduler) ListNames() []string { + f.mu.Lock() + defer f.mu.Unlock() + names := make([]string, 0, len(f.jobs)) + for n := range f.jobs { + names = append(names, n) + } + return names +} +func (f *fakeScheduler) IsAvailable(id string) bool { return f.has(id) } +func (f *fakeScheduler) RemoveWithPrefix(prefix string) { + f.mu.Lock() + defer f.mu.Unlock() + for n := range f.jobs { + if len(n) >= len(prefix) && n[:len(prefix)] == prefix { + delete(f.jobs, n) + } + } +} +func (f *fakeScheduler) AddFunc(name, spec string, _ func()) error { + f.mu.Lock() + defer f.mu.Unlock() + if f.jobs == nil { + f.jobs = map[string]string{} + } + f.jobs[name] = spec + return nil +} +func (f *fakeScheduler) RemoveFunc(name string) { + f.mu.Lock() + defer f.mu.Unlock() + delete(f.jobs, name) +} +func (f *fakeScheduler) has(name string) bool { + f.mu.Lock() + defer f.mu.Unlock() + _, ok := f.jobs[name] + return ok +} + +// writeCustomCert writes custom.crt / custom.key under dir. +func writeCustomCert(t *testing.T, dir string) { + t.Helper() + certFile, keyFile := writeTestCert(t, dir, time.Now(), time.Now().Add(365*24*time.Hour)) + // move generated names to custom names + customCert := filepath.Join(dir, CustomCertFileName) + customKey := filepath.Join(dir, CustomKeyFileName) + if err := os.Rename(certFile, customCert); err != nil { + t.Fatal(err) + } + if err := os.Rename(keyFile, customKey); err != nil { + t.Fatal(err) + } +} + +// writeTestCert creates a self-signed cert/key pair under dir with the given validity window. +func writeTestCert(t *testing.T, dir string, notBefore, notAfter time.Time) (certFile, keyFile string) { + t.Helper() + + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatal(err) + } + + serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128) + serialNumber, err := rand.Int(rand.Reader, serialNumberLimit) + if err != nil { + t.Fatal(err) + } + + template := x509.Certificate{ + SerialNumber: serialNumber, + Subject: pkix.Name{Organization: []string{OrganizationName}}, + NotBefore: notBefore, + NotAfter: notAfter, + KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + BasicConstraintsValid: true, + } + + derBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey) + if err != nil { + t.Fatal(err) + } + + var certBuf bytes.Buffer + if err := pem.Encode(&certBuf, &pem.Block{Type: "CERTIFICATE", Bytes: derBytes}); err != nil { + t.Fatal(err) + } + + privBytes, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + t.Fatal(err) + } + var keyBuf bytes.Buffer + if err := pem.Encode(&keyBuf, &pem.Block{Type: "PRIVATE KEY", Bytes: privBytes}); err != nil { + t.Fatal(err) + } + + certFile = filepath.Join(dir, GeneratedCertFileName) + keyFile = filepath.Join(dir, GeneratedKeyFileName) + if err := os.WriteFile(certFile, certBuf.Bytes(), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(keyFile, keyBuf.Bytes(), 0o600); err != nil { + t.Fatal(err) + } + return certFile, keyFile +} diff --git a/pkg/service/http_listener/listener.go b/pkg/service/http_listener/listener.go index 752681c..6b25175 100644 --- a/pkg/service/http_listener/listener.go +++ b/pkg/service/http_listener/listener.go @@ -11,6 +11,7 @@ import ( "github.com/mycontroller-org/server/v2/pkg/service/http_listener/https" "github.com/mycontroller-org/server/v2/pkg/types/config" + schedulerTY "github.com/mycontroller-org/server/v2/pkg/types/scheduler" serviceTY "github.com/mycontroller-org/server/v2/pkg/types/service" "github.com/mycontroller-org/server/v2/pkg/utils" loggerUtils "github.com/mycontroller-org/server/v2/pkg/utils/logger" @@ -25,9 +26,11 @@ const ( ) type HttpListener struct { - logger *zap.Logger - config config.WebConfig - handler http.Handler + logger *zap.Logger + config config.WebConfig + handler http.Handler + scheduler schedulerTY.CoreScheduler + sslManager *https.SSLManager } func New(ctx context.Context, cfg config.WebConfig, handler http.Handler) (serviceTY.Service, error) { @@ -36,10 +39,17 @@ func New(ctx context.Context, cfg config.WebConfig, handler http.Handler) (servi return nil, err } + scheduler, err := schedulerTY.FromContext(ctx) + if err != nil { + logger.Error("unable to get the core scheduler", zap.Error(err)) + return nil, err + } + return &HttpListener{ - logger: logger.Named("http_listener"), - config: cfg, - handler: handler, + logger: logger.Named("http_listener"), + config: cfg, + handler: handler, + scheduler: scheduler, }, nil } @@ -82,26 +92,32 @@ func (l *HttpListener) Start() error { // https ssl service if l.config.HttpsSSL.Enabled { + sslManager, err := https.NewSSLManager(l.logger, l.config.HttpsSSL, l.scheduler) + if err != nil { + l.logger.Error("error on getting https/ssl manager", zap.Error(err), zap.Any("sslConfig", l.config.HttpsSSL)) + return err + } + l.sslManager = sslManager + + // daily renewal only when HTTPS/SSL is enabled and MyController manages the certificate + if err := sslManager.StartDailyRenewalCheck(); err != nil { + l.logger.Error("error on starting SSL daily renewal check", zap.Error(err)) + return err + } + go func() { addr := fmt.Sprintf("%s:%d", l.config.HttpsSSL.BindAddress, l.config.HttpsSSL.Port) l.logger.Info("listening HTTPS/SSL service on", zap.String("address", addr)) - tlsConfig, err := https.GetSSLTLSConfig(l.logger, l.config.HttpsSSL) - if err != nil { - l.logger.Error("error on getting https/ssl tlsConfig", zap.Error(err), zap.Any("sslConfig", l.config.HttpsSSL)) - errs <- err - return - } - server := &http.Server{ ReadTimeout: readTimeout, Addr: addr, - TLSConfig: tlsConfig, + TLSConfig: sslManager.TLSConfig(), Handler: l.handler, ErrorLog: log.New(getLogger(LoggerPrefixSSL, l.logger), "", 0), } - err = server.ListenAndServeTLS("", "") + err := server.ListenAndServeTLS("", "") if err != nil { l.logger.Error("error on starting https/ssl handler", zap.Error(err)) errs <- err @@ -145,5 +161,8 @@ func (l *HttpListener) Start() error { } func (l *HttpListener) Close() error { + if l.sslManager != nil { + return l.sslManager.Close() + } return nil } diff --git a/pkg/types/config/config.go b/pkg/types/config/config.go index 4d832f1..8e60967 100644 --- a/pkg/types/config/config.go +++ b/pkg/types/config/config.go @@ -56,6 +56,12 @@ type HttpsSSLConfig struct { BindAddress string `yaml:"bind_address"` Port uint `yaml:"port"` CertDir string `yaml:"cert_dir"` + // ValidityDays is the lifetime of a self-signed certificate in days. + // Optional; defaults to 365 when unset or non-positive. + ValidityDays int `yaml:"validity_days"` + // RenewBeforeDays regenerates the self-signed certificate when remaining + // validity is less than this many days. Optional; defaults to 30 when unset or non-positive. + RenewBeforeDays int `yaml:"renew_before_days"` } // HttpsACMEConfig struct diff --git a/resources/sample-binary-server.yaml b/resources/sample-binary-server.yaml index 9f3dc19..48e638a 100644 --- a/resources/sample-binary-server.yaml +++ b/resources/sample-binary-server.yaml @@ -21,6 +21,9 @@ web: bind_address: "0.0.0.0" port: 8443 cert_dir: mc_home/certs/https_ssl + # optional self-signed certificate settings (defaults shown) + # validity_days: 365 + # renew_before_days: 30 https_acme: enabled: false bind_address: "0.0.0.0" diff --git a/resources/sample-docker-server.yaml b/resources/sample-docker-server.yaml index 5b8284a..a1ec1a7 100644 --- a/resources/sample-docker-server.yaml +++ b/resources/sample-docker-server.yaml @@ -21,6 +21,9 @@ web: bind_address: "0.0.0.0" port: 8443 cert_dir: /mc_home/certs/https_ssl + # optional self-signed certificate settings (defaults shown) + # validity_days: 365 + # renew_before_days: 30 https_acme: enabled: false bind_address: "0.0.0.0"