mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 07:36:37 +08:00
migrate
This commit is contained in:
@@ -0,0 +1,262 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package acme
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/go-acme/lego/v4/acme"
|
||||
"github.com/go-acme/lego/v4/certcrypto"
|
||||
"github.com/go-acme/lego/v4/certificate"
|
||||
"github.com/go-acme/lego/v4/challenge/dns01"
|
||||
"github.com/go-acme/lego/v4/lego"
|
||||
"github.com/go-acme/lego/v4/providers/dns/cloudflare"
|
||||
"github.com/go-acme/lego/v4/registration"
|
||||
)
|
||||
|
||||
// AcmeUser implements lego's user interface.
|
||||
type AcmeUser struct {
|
||||
Email string
|
||||
Registration *registration.Resource
|
||||
key crypto.PrivateKey
|
||||
}
|
||||
|
||||
func (u *AcmeUser) GetEmail() string {
|
||||
return u.Email
|
||||
}
|
||||
|
||||
func (u *AcmeUser) GetRegistration() *registration.Resource {
|
||||
return u.Registration
|
||||
}
|
||||
|
||||
func (u *AcmeUser) GetPrivateKey() crypto.PrivateKey {
|
||||
return u.key
|
||||
}
|
||||
|
||||
// CertificateResult holds obtained certificate material.
|
||||
type CertificateResult struct {
|
||||
CertPEM string
|
||||
KeyPEM string
|
||||
NotBefore time.Time
|
||||
NotAfter time.Time
|
||||
}
|
||||
|
||||
func parsePrivateKey(pemData string) (crypto.PrivateKey, error) {
|
||||
block, _ := pem.Decode([]byte(pemData))
|
||||
if block == nil {
|
||||
return nil, errors.New("failed to parse PEM block containing the key")
|
||||
}
|
||||
|
||||
if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
return nil, errors.New("failed to parse private key")
|
||||
}
|
||||
|
||||
func encodePrivateKey(key crypto.PrivateKey) (string, error) {
|
||||
var pemBlock *pem.Block
|
||||
switch k := key.(type) {
|
||||
case *rsa.PrivateKey:
|
||||
pemBlock = &pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(k)}
|
||||
case *ecdsa.PrivateKey:
|
||||
b, err := x509.MarshalECPrivateKey(k)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
pemBlock = &pem.Block{Type: "EC PRIVATE KEY", Bytes: b}
|
||||
default:
|
||||
return "", errors.New("unsupported key type")
|
||||
}
|
||||
return string(pem.EncodeToMemory(pemBlock)), nil
|
||||
}
|
||||
|
||||
// GetOrCreateLegoClient returns a configured lego client and optional new account credentials.
|
||||
func GetOrCreateLegoClient(acmeEmail, privateKeyPEM, accountURL string, keyAlgorithm string) (*lego.Client, *AcmeUser, string, string, error) {
|
||||
var privateKey crypto.PrivateKey
|
||||
var err error
|
||||
var newPrivateKeyPEM string
|
||||
var newAccountURL string
|
||||
|
||||
if privateKeyPEM == "" {
|
||||
privateKey, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
pemStr, err := encodePrivateKey(privateKey)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
newPrivateKeyPEM = pemStr
|
||||
} else {
|
||||
privateKey, err = parsePrivateKey(privateKeyPEM)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
}
|
||||
|
||||
user := &AcmeUser{
|
||||
Email: acmeEmail,
|
||||
key: privateKey,
|
||||
}
|
||||
|
||||
if accountURL != "" {
|
||||
user.Registration = ®istration.Resource{
|
||||
Body: acme.Account{
|
||||
Status: "valid",
|
||||
Contact: []string{"mailto:" + acmeEmail},
|
||||
},
|
||||
URI: accountURL,
|
||||
}
|
||||
}
|
||||
|
||||
config := lego.NewConfig(user)
|
||||
config.CADirURL = lego.LEDirectoryProduction
|
||||
|
||||
switch keyAlgorithm {
|
||||
case "RSA2048":
|
||||
config.Certificate.KeyType = certcrypto.RSA2048
|
||||
case "RSA4096":
|
||||
config.Certificate.KeyType = certcrypto.RSA4096
|
||||
case "EC256":
|
||||
config.Certificate.KeyType = certcrypto.EC256
|
||||
case "EC384":
|
||||
config.Certificate.KeyType = certcrypto.EC384
|
||||
default:
|
||||
config.Certificate.KeyType = certcrypto.RSA2048
|
||||
}
|
||||
|
||||
client, err := lego.NewClient(config)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
|
||||
if accountURL == "" {
|
||||
reg, err := client.Registration.Register(registration.RegisterOptions{TermsOfServiceAgreed: true})
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
user.Registration = reg
|
||||
newAccountURL = reg.URI
|
||||
}
|
||||
|
||||
return client, user, newPrivateKeyPEM, newAccountURL, nil
|
||||
}
|
||||
|
||||
// SetupDNSProvider configures DNS-01 challenge for the lego client.
|
||||
func SetupDNSProvider(client *lego.Client, dnsType, dnsAuth string, dns1, dns2 string, disableCNAME, skipDNS bool) error {
|
||||
var provider challengeProvider
|
||||
|
||||
switch dnsType {
|
||||
case "cloudflare":
|
||||
var creds map[string]string
|
||||
if err := json.Unmarshal([]byte(dnsAuth), &creds); err != nil {
|
||||
return fmt.Errorf("failed to parse cloudflare credentials: %v", err)
|
||||
}
|
||||
|
||||
config := cloudflare.NewDefaultConfig()
|
||||
config.AuthToken = creds["api_token"]
|
||||
|
||||
p, err := cloudflare.NewDNSProviderConfig(config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
provider = p
|
||||
default:
|
||||
return fmt.Errorf("unsupported DNS provider: %s", dnsType)
|
||||
}
|
||||
|
||||
var resolvers []string
|
||||
if dns1 != "" {
|
||||
resolvers = append(resolvers, dns1+":53")
|
||||
}
|
||||
if dns2 != "" {
|
||||
resolvers = append(resolvers, dns2+":53")
|
||||
}
|
||||
|
||||
var opts []dns01.ChallengeOption
|
||||
|
||||
if len(resolvers) > 0 {
|
||||
opts = append(opts, dns01.AddRecursiveNameservers(resolvers))
|
||||
}
|
||||
|
||||
if disableCNAME {
|
||||
opts = append(opts, dns01.DisableCompletePropagationRequirement())
|
||||
}
|
||||
|
||||
if skipDNS {
|
||||
opts = append(opts, dns01.WrapPreCheck(func(domain, fqdn, value string, check dns01.PreCheckFunc) (bool, error) {
|
||||
time.Sleep(20 * time.Second)
|
||||
return true, nil
|
||||
}))
|
||||
}
|
||||
|
||||
return client.Challenge.SetDNS01Provider(provider, opts...)
|
||||
}
|
||||
|
||||
type challengeProvider interface {
|
||||
Present(domain, token, keyAuth string) error
|
||||
CleanUp(domain, token, keyAuth string) error
|
||||
}
|
||||
|
||||
// ObtainSSL obtains a certificate via ACME DNS-01 challenge.
|
||||
func ObtainSSL(
|
||||
acmeEmail, acmePrivateKeyPEM, acmeURL string,
|
||||
dnsType, dnsAuth string,
|
||||
dns1, dns2 string,
|
||||
disableCNAME, skipDNS bool,
|
||||
keyAlgorithm string,
|
||||
domains []string,
|
||||
) (string, string, *CertificateResult, error) {
|
||||
client, _, newPrivateKeyPEM, newAccountURL, err := GetOrCreateLegoClient(acmeEmail, acmePrivateKeyPEM, acmeURL, keyAlgorithm)
|
||||
if err != nil {
|
||||
return "", "", nil, fmt.Errorf("failed to create ACME client: %w", err)
|
||||
}
|
||||
|
||||
err = SetupDNSProvider(client, dnsType, dnsAuth, dns1, dns2, disableCNAME, skipDNS)
|
||||
if err != nil {
|
||||
return newAccountURL, newPrivateKeyPEM, nil, fmt.Errorf("failed to setup DNS provider: %w", err)
|
||||
}
|
||||
|
||||
request := certificate.ObtainRequest{
|
||||
Domains: domains,
|
||||
Bundle: true,
|
||||
}
|
||||
|
||||
certificates, err := client.Certificate.Obtain(request)
|
||||
if err != nil {
|
||||
return newAccountURL, newPrivateKeyPEM, nil, fmt.Errorf("failed to obtain certificate: %w", err)
|
||||
}
|
||||
|
||||
result := &CertificateResult{
|
||||
CertPEM: string(certificates.Certificate),
|
||||
KeyPEM: string(certificates.PrivateKey),
|
||||
}
|
||||
|
||||
certBlock, _ := pem.Decode(certificates.Certificate)
|
||||
if certBlock != nil {
|
||||
parsedCert, err := x509.ParseCertificate(certBlock.Bytes)
|
||||
if err == nil {
|
||||
result.NotBefore = parsedCert.NotBefore
|
||||
result.NotAfter = parsedCert.NotAfter
|
||||
}
|
||||
}
|
||||
|
||||
return newAccountURL, newPrivateKeyPEM, result, nil
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSplitAcmeDomains(t *testing.T) {
|
||||
assert.Equal(t, []string{"example.com"}, splitAcmeDomains("example.com", ""))
|
||||
assert.Equal(t, []string{"example.com", "*.example.com"}, splitAcmeDomains("example.com", "*.example.com"))
|
||||
assert.Equal(t, []string{"example.com", "www.example.com", "api.example.com"}, splitAcmeDomains("example.com", "www.example.com\napi.example.com"))
|
||||
assert.Equal(t, []string{"example.com", "www.example.com", "api.example.com"}, splitAcmeDomains("example.com", "www.example.com, api.example.com"))
|
||||
}
|
||||
|
||||
func TestCertificatesDueForRenewal(t *testing.T) {
|
||||
now := time.Date(2026, 6, 18, 12, 0, 0, 0, time.UTC)
|
||||
certificates := []model.TLSCertificate{
|
||||
{ID: 1, Provider: "acme", AutoRenew: true, ApplyStatus: "ready", PrimaryDomain: "due.example.com", NotAfter: now.Add(3 * 24 * time.Hour)},
|
||||
{ID: 2, Provider: "acme", AutoRenew: true, ApplyStatus: "ready", PrimaryDomain: "fresh.example.com", NotAfter: now.Add(30 * 24 * time.Hour)},
|
||||
{ID: 3, Provider: "upload", AutoRenew: true, ApplyStatus: "ready", PrimaryDomain: "upload.example.com", NotAfter: now.Add(24 * time.Hour)},
|
||||
{ID: 4, Provider: "acme", AutoRenew: false, ApplyStatus: "ready", PrimaryDomain: "manual.example.com", NotAfter: now.Add(24 * time.Hour)},
|
||||
{ID: 5, Provider: "acme", AutoRenew: true, ApplyStatus: "applying", PrimaryDomain: "busy.example.com", NotAfter: now.Add(24 * time.Hour)},
|
||||
}
|
||||
|
||||
due := CertificatesDueForRenewal(certificates, now)
|
||||
require.Len(t, due, 1)
|
||||
assert.Equal(t, uint(1), due[0].ID)
|
||||
assert.Equal(t, "due.example.com", due[0].PrimaryDomain)
|
||||
}
|
||||
|
||||
func TestApplyCertificateReturnsApplying(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
dnsAccount, err := CreateDNSAccount(ctx, DNSAccountInput{
|
||||
Name: "Test Cloudflare",
|
||||
Type: "cloudflare",
|
||||
Authorization: `{"api_token": "dummy_token"}`,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
return updateCertError(ctx, cert, "dns challenge failed")
|
||||
})
|
||||
defer restore()
|
||||
|
||||
cert, err := ApplyCertificate(ctx, ApplyInput{
|
||||
Name: "Test ACME Cert",
|
||||
PrimaryDomain: "example.com",
|
||||
OtherDomains: "*.example.com",
|
||||
DnsAccountID: dnsAccount.ID,
|
||||
KeyAlgorithm: "RSA2048",
|
||||
AutoRenew: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "applying", cert.ApplyStatus)
|
||||
assert.Equal(t, "acme", cert.Provider)
|
||||
}
|
||||
|
||||
func TestRenewCertificateSetsApplying(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
cert := &model.TLSCertificate{
|
||||
Name: "renew-cert",
|
||||
Provider: "acme",
|
||||
AutoRenew: true,
|
||||
ApplyStatus: "ready",
|
||||
PrimaryDomain: "renew.example.com",
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
}
|
||||
require.NoError(t, model.CreateTLSCertificateRecord(ctx, cert))
|
||||
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error {
|
||||
return nil
|
||||
})
|
||||
defer restore()
|
||||
|
||||
renewed, err := RenewCertificate(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "applying", renewed.ApplyStatus)
|
||||
}
|
||||
|
||||
func TestConvertCertificateToACMEPreservesUploadOnFailure(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
originalCertPEM, originalKeyPEM := generateTestCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateCertificate(ctx, CertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: originalCertPEM,
|
||||
KeyPEM: originalKeyPEM,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := model.GetTLSCertificateByID(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
originalStoredCertPEM := stored.CertPEM
|
||||
originalStoredKeyPEM := stored.KeyPEM
|
||||
|
||||
stored.ApplyStatus = "applying"
|
||||
stored.PrimaryDomain = "manual.example.com"
|
||||
require.NoError(t, model.SaveTLSCertificate(ctx, stored))
|
||||
|
||||
err = updateCertError(ctx, stored, "dns challenge failed")
|
||||
require.Error(t, err)
|
||||
|
||||
finalCert, err := model.GetTLSCertificateByID(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "upload", finalCert.Provider)
|
||||
assert.Equal(t, "error", finalCert.ApplyStatus)
|
||||
assert.Equal(t, originalStoredCertPEM, finalCert.CertPEM)
|
||||
assert.Equal(t, originalStoredKeyPEM, finalCert.KeyPEM)
|
||||
assert.True(t, strings.Contains(finalCert.ApplyMessage, "dns challenge failed"))
|
||||
}
|
||||
|
||||
func TestConvertCertificateToACMERejectsInvalidStates(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
certPEM, keyPEM := generateTestCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateCertificate(ctx, CertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
cert.Provider = "acme"
|
||||
require.NoError(t, model.SaveTLSCertificate(ctx, cert))
|
||||
_, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "only uploaded")
|
||||
|
||||
cert.Provider = "upload"
|
||||
cert.ApplyStatus = "applying"
|
||||
require.NoError(t, model.SaveTLSCertificate(ctx, cert))
|
||||
_, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already applying")
|
||||
}
|
||||
@@ -23,6 +23,4 @@ const (
|
||||
errManagedDomainCertNotFound = "所选证书不存在"
|
||||
|
||||
errDNSAccountInUse = "该 DNS 账号已被证书使用,无法删除"
|
||||
|
||||
errACMENotImplemented = "ACME certificate obtain is not implemented yet"
|
||||
)
|
||||
|
||||
@@ -176,7 +176,7 @@ func DeleteCertificate(ctx context.Context, id uint) error {
|
||||
return model.DeleteTLSCertificateRecord(ctx, id)
|
||||
}
|
||||
|
||||
// ApplyCertificate 申请 ACME 证书(当前为占位实现)。
|
||||
// ApplyCertificate 申请 ACME 证书。
|
||||
func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertificate, error) {
|
||||
cert := &model.TLSCertificate{
|
||||
Provider: "acme",
|
||||
@@ -193,10 +193,15 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return markACMEStubFailure(ctx, cert)
|
||||
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = obtainTLSCertificate(context.Background(), c)
|
||||
}(cert)
|
||||
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
// UpdateACMECertificate 更新 ACME 证书配置(当前为占位实现)。
|
||||
// UpdateACMECertificate 更新 ACME 证书配置。
|
||||
func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
@@ -215,10 +220,15 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return markACMEStubFailure(ctx, cert)
|
||||
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = obtainTLSCertificate(context.Background(), c)
|
||||
}(cert)
|
||||
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
// ConvertCertificateToACME 将上传证书转为 ACME 管理(当前为占位实现)。
|
||||
// ConvertCertificateToACME 将上传证书转为 ACME 管理。
|
||||
func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
@@ -241,10 +251,25 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return markACMEStubFailure(ctx, cert)
|
||||
|
||||
go func(c *model.TLSCertificate) {
|
||||
if err := obtainTLSCertificate(context.Background(), c); err != nil {
|
||||
return
|
||||
}
|
||||
latest, err := model.GetTLSCertificateByID(context.Background(), c.ID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
latest.Provider = "acme"
|
||||
latest.ApplyStatus = "ready"
|
||||
latest.ApplyMessage = ""
|
||||
_ = model.SaveTLSCertificate(context.Background(), latest)
|
||||
}(cert)
|
||||
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
// RenewCertificate 续期 ACME 证书(当前为占位实现)。
|
||||
// RenewCertificate 续期 ACME 证书。
|
||||
func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
@@ -253,12 +278,17 @@ func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, erro
|
||||
if cert.Provider != "acme" {
|
||||
return nil, errors.New(errCertificateOnlyACMERenew)
|
||||
}
|
||||
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = obtainTLSCertificate(context.Background(), c)
|
||||
}(cert)
|
||||
|
||||
cert.ApplyStatus = "applying"
|
||||
cert.ApplyMessage = ""
|
||||
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return markACMEStubFailure(ctx, cert)
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
// ListDNSAccounts 列出 DNS 账号。
|
||||
@@ -386,19 +416,9 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
|
||||
cert.SkipDNS = input.SkipDNS
|
||||
cert.DNS1 = strings.TrimSpace(input.DNS1)
|
||||
cert.DNS2 = strings.TrimSpace(input.DNS2)
|
||||
cert.Provider = "acme"
|
||||
cert.ApplyStatus = "applying"
|
||||
}
|
||||
|
||||
func markACMEStubFailure(ctx context.Context, cert *model.TLSCertificate) (*model.TLSCertificate, error) {
|
||||
cert.ApplyStatus = "failed"
|
||||
cert.ApplyMessage = errACMENotImplemented
|
||||
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
|
||||
routes, err := model.ListTLSProxyRouteRefs(ctx)
|
||||
if err != nil {
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -24,8 +25,11 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var tlsTestDBMu sync.Mutex
|
||||
|
||||
func setupTLSTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
tlsTestDBMu.Lock()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
@@ -44,6 +48,7 @@ func setupTLSTestDB(t *testing.T) func() {
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
tlsTestDBMu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const acmeRenewLeadTime = 7 * 24 * time.Hour
|
||||
|
||||
var obtainTLSCertificate = obtainCertificate
|
||||
|
||||
// SetObtainCertificateFuncForTest swaps the async obtain implementation for tests.
|
||||
func SetObtainCertificateFuncForTest(fn func(context.Context, *model.TLSCertificate) error) func() {
|
||||
previous := obtainTLSCertificate
|
||||
obtainTLSCertificate = fn
|
||||
return func() {
|
||||
obtainTLSCertificate = previous
|
||||
}
|
||||
}
|
||||
|
||||
func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
cert.ApplyStatus = "applying"
|
||||
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
acmeAccount, err := model.GetAcmeAccountByID(ctx, cert.AcmeAccountID)
|
||||
if err != nil {
|
||||
acmeAccount, err = model.GetDefaultAcmeAccount(ctx)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get ACME account: %v", err))
|
||||
}
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
dnsAccount, err := model.GetDNSAccountByID(ctx, cert.DnsAccountID)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get DNS account: %v", err))
|
||||
}
|
||||
|
||||
dnsAuth, err := openSensitive(dnsAccount.Authorization)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to decrypt DNS credentials: %v", err))
|
||||
}
|
||||
|
||||
acmePrivateKey, err := openSensitive(acmeAccount.PrivateKey)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to decrypt ACME account key: %v", err))
|
||||
}
|
||||
|
||||
domains := splitAcmeDomains(cert.PrimaryDomain, cert.OtherDomains)
|
||||
|
||||
newAccountURL, newPrivateKeyPEM, result, err := acme.ObtainSSL(
|
||||
acmeAccount.Email,
|
||||
acmePrivateKey,
|
||||
acmeAccount.URL,
|
||||
dnsAccount.Type,
|
||||
dnsAuth,
|
||||
cert.DNS1,
|
||||
cert.DNS2,
|
||||
cert.DisableCNAME,
|
||||
cert.SkipDNS,
|
||||
cert.KeyAlgorithm,
|
||||
domains,
|
||||
)
|
||||
|
||||
if (newPrivateKeyPEM != "" && acmePrivateKey != newPrivateKeyPEM) || (newAccountURL != "" && acmeAccount.URL != newAccountURL) {
|
||||
if newPrivateKeyPEM != "" {
|
||||
sealedKey, sealErr := sealSensitive(newPrivateKeyPEM)
|
||||
if sealErr != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to seal ACME account key: %v", sealErr))
|
||||
}
|
||||
acmeAccount.PrivateKey = sealedKey
|
||||
}
|
||||
if newAccountURL != "" {
|
||||
acmeAccount.URL = newAccountURL
|
||||
}
|
||||
if acmeAccount.ID == 0 {
|
||||
if dbErr := model.CreateAcmeAccountRecord(ctx, acmeAccount); dbErr != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to create ACME account: %v", dbErr))
|
||||
}
|
||||
} else if dbErr := model.SaveAcmeAccount(ctx, acmeAccount); dbErr != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to save ACME account: %v", dbErr))
|
||||
}
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, err.Error())
|
||||
}
|
||||
|
||||
sealedKey, err := sealSensitive(result.KeyPEM)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to seal certificate key: %v", err))
|
||||
}
|
||||
|
||||
cert.CertPEM = result.CertPEM
|
||||
cert.KeyPEM = sealedKey
|
||||
cert.NotBefore = result.NotBefore
|
||||
cert.NotAfter = result.NotAfter
|
||||
cert.ApplyStatus = "ready"
|
||||
cert.ApplyMessage = ""
|
||||
|
||||
return model.SaveTLSCertificate(ctx, cert)
|
||||
}
|
||||
|
||||
func updateCertError(ctx context.Context, cert *model.TLSCertificate, message string) error {
|
||||
cert.ApplyStatus = "error"
|
||||
cert.ApplyMessage = message
|
||||
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("%s", message)
|
||||
}
|
||||
|
||||
func splitAcmeDomains(primaryDomain, otherDomains string) []string {
|
||||
primaryDomain = strings.TrimSpace(primaryDomain)
|
||||
domains := []string{}
|
||||
if primaryDomain != "" {
|
||||
domains = append(domains, primaryDomain)
|
||||
}
|
||||
otherDomains = strings.TrimSpace(otherDomains)
|
||||
if otherDomains == "" {
|
||||
return domains
|
||||
}
|
||||
|
||||
separator := "\n"
|
||||
if !strings.Contains(otherDomains, "\n") && strings.Contains(otherDomains, ",") {
|
||||
separator = ","
|
||||
}
|
||||
for _, domain := range strings.Split(otherDomains, separator) {
|
||||
domain = strings.TrimSpace(domain)
|
||||
if domain != "" {
|
||||
domains = append(domains, domain)
|
||||
}
|
||||
}
|
||||
return domains
|
||||
}
|
||||
|
||||
// CertificatesDueForRenewal returns ACME certificates that should be renewed at the given time.
|
||||
func CertificatesDueForRenewal(certificates []model.TLSCertificate, now time.Time) []model.TLSCertificate {
|
||||
due := make([]model.TLSCertificate, 0)
|
||||
for _, cert := range certificates {
|
||||
if !cert.AutoRenew || cert.Provider != "acme" || cert.ApplyStatus == "applying" {
|
||||
continue
|
||||
}
|
||||
if cert.NotAfter.IsZero() {
|
||||
continue
|
||||
}
|
||||
if cert.NotAfter.Sub(now) < acmeRenewLeadTime {
|
||||
due = append(due, cert)
|
||||
}
|
||||
}
|
||||
return due
|
||||
}
|
||||
Reference in New Issue
Block a user