mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
[新增] 添加转换上传的 TLS 证书为 ACME 管理证书的功能
This commit is contained in:
@@ -1,7 +1,9 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"openflare/model"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
@@ -84,3 +86,186 @@ func TestAcmeAndDnsIntegration(t *testing.T) {
|
||||
t.Fatalf("Failed to delete DNS Account after cert cleanup: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertTLSCertificateToAcmePreservesUploadUntilSuccess(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
originalCertPEM, originalKeyPEM := generateCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: originalCertPEM,
|
||||
KeyPEM: originalKeyPEM,
|
||||
Remark: "manual upload",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
originalCertPEM = cert.CertPEM
|
||||
originalKeyPEM = cert.KeyPEM
|
||||
|
||||
newCertPEM, newKeyPEM := generateCertificatePair(t, []string{"managed.example.com"})
|
||||
started := make(chan struct{}, 1)
|
||||
release := make(chan struct{})
|
||||
restore := SetTLSCertificateObtainFuncForTest(func(c *model.TLSCertificate) error {
|
||||
started <- struct{}{}
|
||||
<-release
|
||||
c.CertPEM = newCertPEM
|
||||
c.KeyPEM = newKeyPEM
|
||||
c.NotBefore = time.Now().Add(-time.Hour)
|
||||
c.NotAfter = time.Now().Add(90 * 24 * time.Hour)
|
||||
c.ApplyStatus = "ready"
|
||||
c.ApplyMessage = ""
|
||||
return model.DB.Save(c).Error
|
||||
})
|
||||
t.Cleanup(restore)
|
||||
|
||||
converted, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{
|
||||
Name: "managed-cert",
|
||||
Remark: "converted",
|
||||
AcmeAccountID: 1,
|
||||
DnsAccountID: 2,
|
||||
KeyAlgorithm: "EC256",
|
||||
AutoRenew: true,
|
||||
PrimaryDomain: "managed.example.com",
|
||||
OtherDomains: "www.managed.example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ConvertTLSCertificateToAcme failed: %v", err)
|
||||
}
|
||||
if converted.ID != cert.ID {
|
||||
t.Fatalf("expected converted certificate to keep id %d, got %d", cert.ID, converted.ID)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected conversion obtain task to start")
|
||||
}
|
||||
|
||||
applying, err := model.GetTLSCertificateByID(cert.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("reload applying certificate failed: %v", err)
|
||||
}
|
||||
if applying.Provider != "upload" {
|
||||
t.Fatalf("expected provider to remain upload while applying, got %s", applying.Provider)
|
||||
}
|
||||
if applying.ApplyStatus != "applying" {
|
||||
t.Fatalf("expected applying status, got %s", applying.ApplyStatus)
|
||||
}
|
||||
if applying.CertPEM != originalCertPEM || applying.KeyPEM != originalKeyPEM {
|
||||
t.Fatal("expected original PEM payloads to be preserved while applying")
|
||||
}
|
||||
|
||||
close(release)
|
||||
|
||||
finalCert := waitForCertificateState(t, cert.ID, func(c *model.TLSCertificate) bool {
|
||||
return c.Provider == "acme" && c.ApplyStatus == "ready"
|
||||
})
|
||||
if finalCert.CertPEM != newCertPEM || finalCert.KeyPEM != newKeyPEM {
|
||||
t.Fatal("expected successful conversion to replace PEM payloads")
|
||||
}
|
||||
if !finalCert.AutoRenew {
|
||||
t.Fatal("expected converted certificate to keep auto renew enabled")
|
||||
}
|
||||
if finalCert.PrimaryDomain != "managed.example.com" || finalCert.OtherDomains != "www.managed.example.com" {
|
||||
t.Fatalf("expected converted certificate to persist ACME domains, got %+v", finalCert)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertTLSCertificateToAcmePreservesUploadOnFailure(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
originalCertPEM, originalKeyPEM := generateCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: originalCertPEM,
|
||||
KeyPEM: originalKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
originalCertPEM = cert.CertPEM
|
||||
originalKeyPEM = cert.KeyPEM
|
||||
|
||||
restore := SetTLSCertificateObtainFuncForTest(func(c *model.TLSCertificate) error {
|
||||
err := errors.New("dns challenge failed")
|
||||
updateCertError(c, err.Error())
|
||||
return err
|
||||
})
|
||||
t.Cleanup(restore)
|
||||
|
||||
if _, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{
|
||||
Name: "manual-cert",
|
||||
DnsAccountID: 1,
|
||||
PrimaryDomain: "manual.example.com",
|
||||
}); err != nil {
|
||||
t.Fatalf("ConvertTLSCertificateToAcme failed: %v", err)
|
||||
}
|
||||
|
||||
finalCert := waitForCertificateState(t, cert.ID, func(c *model.TLSCertificate) bool {
|
||||
return c.ApplyStatus == "error"
|
||||
})
|
||||
if finalCert.Provider != "upload" {
|
||||
t.Fatalf("expected failed conversion to keep upload provider, got %s", finalCert.Provider)
|
||||
}
|
||||
if finalCert.CertPEM != originalCertPEM || finalCert.KeyPEM != originalKeyPEM {
|
||||
t.Fatal("expected failed conversion to preserve original PEM payloads")
|
||||
}
|
||||
if !strings.Contains(finalCert.ApplyMessage, "dns challenge failed") {
|
||||
t.Fatalf("expected conversion error message, got %q", finalCert.ApplyMessage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertTLSCertificateToAcmeRejectsInvalidStates(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
certPEM, keyPEM := generateCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
cert.Provider = "acme"
|
||||
if err := cert.Update(); err != nil {
|
||||
t.Fatalf("failed to mark certificate acme: %v", err)
|
||||
}
|
||||
if _, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{Name: "manual-cert"}); err == nil || !strings.Contains(err.Error(), "only uploaded") {
|
||||
t.Fatalf("expected non-upload conversion to fail, got %v", err)
|
||||
}
|
||||
|
||||
cert.Provider = "upload"
|
||||
cert.ApplyStatus = "applying"
|
||||
if err := cert.Update(); err != nil {
|
||||
t.Fatalf("failed to mark certificate applying: %v", err)
|
||||
}
|
||||
if _, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{Name: "manual-cert"}); err == nil || !strings.Contains(err.Error(), "already applying") {
|
||||
t.Fatalf("expected applying conversion to fail, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func waitForCertificateState(t *testing.T, id uint, matches func(*model.TLSCertificate) bool) *model.TLSCertificate {
|
||||
t.Helper()
|
||||
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
t.Fatalf("reload certificate %d failed: %v", id, err)
|
||||
}
|
||||
if matches(cert) {
|
||||
return cert
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
t.Fatalf("reload certificate %d failed: %v", id, err)
|
||||
}
|
||||
t.Fatalf("certificate %d did not reach expected state: %+v", id, cert)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -40,6 +40,16 @@ type TLSApplyInput struct {
|
||||
DNS2 string `json:"dns2"`
|
||||
}
|
||||
|
||||
var obtainTLSCertificate = ObtainSSL
|
||||
|
||||
func SetTLSCertificateObtainFuncForTest(fn func(*model.TLSCertificate) error) func() {
|
||||
previous := obtainTLSCertificate
|
||||
obtainTLSCertificate = fn
|
||||
return func() {
|
||||
obtainTLSCertificate = previous
|
||||
}
|
||||
}
|
||||
|
||||
func ListTLSCertificates() ([]*model.TLSCertificate, error) {
|
||||
return model.ListTLSCertificates()
|
||||
}
|
||||
@@ -191,7 +201,7 @@ func ApplyTLSCertificate(input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
|
||||
// Async obtain SSL
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = ObtainSSL(c)
|
||||
_ = obtainTLSCertificate(c)
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
@@ -233,7 +243,64 @@ func UpdateAcmeCertificate(id uint, input TLSApplyInput) (*model.TLSCertificate,
|
||||
|
||||
// Async obtain SSL with updated config
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = ObtainSSL(c)
|
||||
_ = obtainTLSCertificate(c)
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func ConvertTLSCertificateToAcme(id uint, input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "upload" {
|
||||
return nil, errors.New("only uploaded certificates can be converted to acme")
|
||||
}
|
||||
if cert.ApplyStatus == "applying" {
|
||||
return nil, errors.New("certificate is already applying")
|
||||
}
|
||||
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, errors.New("certificate name cannot be empty")
|
||||
}
|
||||
|
||||
cert.Name = name
|
||||
cert.Remark = strings.TrimSpace(input.Remark)
|
||||
cert.AcmeAccountID = input.AcmeAccountID
|
||||
cert.DnsAccountID = input.DnsAccountID
|
||||
cert.KeyAlgorithm = input.KeyAlgorithm
|
||||
cert.AutoRenew = input.AutoRenew
|
||||
cert.PrimaryDomain = strings.TrimSpace(input.PrimaryDomain)
|
||||
cert.OtherDomains = strings.TrimSpace(input.OtherDomains)
|
||||
cert.DisableCNAME = input.DisableCNAME
|
||||
cert.SkipDNS = input.SkipDNS
|
||||
cert.DNS1 = strings.TrimSpace(input.DNS1)
|
||||
cert.DNS2 = strings.TrimSpace(input.DNS2)
|
||||
cert.ApplyStatus = "applying"
|
||||
cert.ApplyMessage = ""
|
||||
|
||||
if err := cert.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
go func(c *model.TLSCertificate) {
|
||||
if err := obtainTLSCertificate(c); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
latest, err := model.GetTLSCertificateByID(c.ID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
latest.Provider = "acme"
|
||||
latest.ApplyStatus = "ready"
|
||||
latest.ApplyMessage = ""
|
||||
_ = latest.Update()
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
@@ -250,7 +317,7 @@ func RenewTLSCertificate(id uint) (*model.TLSCertificate, error) {
|
||||
|
||||
// Async obtain SSL
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = ObtainSSL(c)
|
||||
_ = obtainTLSCertificate(c)
|
||||
}(cert)
|
||||
|
||||
cert.ApplyStatus = "applying"
|
||||
|
||||
Reference in New Issue
Block a user