mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 15:46:37 +08:00
[新增] 添加 ACME 和 DNS 账号管理功能,支持证书申请与续期
This commit is contained in:
@@ -0,0 +1,298 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"openflare/model"
|
||||
|
||||
"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"
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func GetOrCreateLegoClient(account *model.AcmeAccount, keyAlgorithm string) (*lego.Client, *AcmeUser, error) {
|
||||
var privateKey crypto.PrivateKey
|
||||
var err error
|
||||
|
||||
if account.PrivateKey == "" {
|
||||
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
|
||||
}
|
||||
account.PrivateKey = pemStr
|
||||
// Don't save it to DB yet, wait for successful registration.
|
||||
} else {
|
||||
privateKey, err = parsePrivateKey(account.PrivateKey)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
user := &AcmeUser{
|
||||
Email: account.Email,
|
||||
key: privateKey,
|
||||
}
|
||||
|
||||
if account.URL != "" {
|
||||
user.Registration = ®istration.Resource{
|
||||
Body: acme.Account{
|
||||
Status: "valid",
|
||||
Contact: []string{"mailto:" + account.Email},
|
||||
},
|
||||
URI: account.URL,
|
||||
}
|
||||
}
|
||||
|
||||
config := lego.NewConfig(user)
|
||||
// Use Let's Encrypt production environment by default
|
||||
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 account.URL == "" {
|
||||
reg, err := client.Registration.Register(registration.RegisterOptions{TermsOfServiceAgreed: true})
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
user.Registration = reg
|
||||
account.URL = reg.URI
|
||||
if account.ID == 0 {
|
||||
err = model.DB.Create(account).Error
|
||||
} else {
|
||||
err = model.DB.Save(account).Error
|
||||
}
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return client, user, nil
|
||||
}
|
||||
|
||||
func SetupDNSProvider(client *lego.Client, dnsAccount *model.DnsAccount, dns1, dns2 string, disableCNAME, skipDNS bool) error {
|
||||
var provider challengeProvider
|
||||
|
||||
switch dnsAccount.Type {
|
||||
case "cloudflare":
|
||||
var creds map[string]string
|
||||
if err := json.Unmarshal([]byte(dnsAccount.Authorization), &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", dnsAccount.Type)
|
||||
}
|
||||
|
||||
// We can use custom DNS servers to verify challenges if provided
|
||||
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) {
|
||||
// If we skip the local DNS check entirely, we might trigger Let's Encrypt to verify
|
||||
// BEFORE Cloudflare's edge servers have actually synced the TXT record (which takes 5-15 seconds).
|
||||
// So we add a safe 20-second artificial delay before forcing the true return.
|
||||
time.Sleep(20 * time.Second)
|
||||
return true, nil
|
||||
}))
|
||||
}
|
||||
|
||||
return client.Challenge.SetDNS01Provider(provider, opts...)
|
||||
}
|
||||
|
||||
// challengeProvider interface helps to bypass the strict type definition of SetDNS01Provider
|
||||
type challengeProvider interface {
|
||||
Present(domain, token, keyAuth string) error
|
||||
CleanUp(domain, token, keyAuth string) error
|
||||
}
|
||||
|
||||
func ObtainSSL(cert *model.TLSCertificate) error {
|
||||
cert.ApplyStatus = "applying"
|
||||
model.DB.Save(cert)
|
||||
|
||||
acmeAccount, err := model.GetAcmeAccountByID(cert.AcmeAccountID)
|
||||
if err != nil {
|
||||
// Fallback to default ACME account if the specified one is not found (e.g. ID 0 during testing)
|
||||
acmeAccount, err = model.GetDefaultAcmeAccount()
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to get ACME account: %v", err))
|
||||
return err
|
||||
}
|
||||
// Self-heal the certificate
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
model.DB.Save(cert)
|
||||
}
|
||||
|
||||
dnsAccount, err := model.GetDnsAccountByID(cert.DnsAccountID)
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to get DNS account: %v", err))
|
||||
return err
|
||||
}
|
||||
|
||||
client, _, err := GetOrCreateLegoClient(acmeAccount, cert.KeyAlgorithm)
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to create ACME client: %v", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = SetupDNSProvider(client, dnsAccount, cert.DNS1, cert.DNS2, cert.DisableCNAME, cert.SkipDNS)
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to setup DNS provider: %v", err))
|
||||
return err
|
||||
}
|
||||
|
||||
domains := []string{cert.PrimaryDomain}
|
||||
if cert.OtherDomains != "" {
|
||||
for _, d := range strings.Split(cert.OtherDomains, "\n") {
|
||||
d = strings.TrimSpace(d)
|
||||
if d != "" {
|
||||
domains = append(domains, d)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
request := certificate.ObtainRequest{
|
||||
Domains: domains,
|
||||
Bundle: true,
|
||||
}
|
||||
|
||||
certificates, err := client.Certificate.Obtain(request)
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to obtain certificate: %v", err))
|
||||
return err
|
||||
}
|
||||
|
||||
cert.CertPEM = string(certificates.Certificate)
|
||||
cert.KeyPEM = string(certificates.PrivateKey)
|
||||
|
||||
// Parse validity dates
|
||||
certBlock, _ := pem.Decode(certificates.Certificate)
|
||||
if certBlock != nil {
|
||||
parsedCert, err := x509.ParseCertificate(certBlock.Bytes)
|
||||
if err == nil {
|
||||
cert.NotBefore = parsedCert.NotBefore
|
||||
cert.NotAfter = parsedCert.NotAfter
|
||||
}
|
||||
}
|
||||
|
||||
cert.ApplyStatus = "ready"
|
||||
cert.ApplyMessage = ""
|
||||
return model.DB.Save(cert).Error
|
||||
}
|
||||
|
||||
func updateCertError(cert *model.TLSCertificate, message string) {
|
||||
cert.ApplyStatus = "error"
|
||||
cert.ApplyMessage = message
|
||||
model.DB.Save(cert)
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"openflare/model"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestAcmeAndDnsIntegration(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// 1. Create a DNS Account
|
||||
dnsAccount := &model.DnsAccount{
|
||||
Name: "Test Cloudflare",
|
||||
Type: "cloudflare",
|
||||
Authorization: `{"api_token": "dummy_token"}`,
|
||||
}
|
||||
if err := dnsAccount.Insert(); err != nil {
|
||||
t.Fatalf("Failed to insert DNS Account: %v", err)
|
||||
}
|
||||
|
||||
// 2. Apply for TLS Certificate (using the new ApplyTLSCertificate function)
|
||||
certInput := TLSApplyInput{
|
||||
Name: "Test ACME Cert",
|
||||
PrimaryDomain: "example.com",
|
||||
OtherDomains: "*.example.com",
|
||||
DnsAccountID: dnsAccount.ID,
|
||||
KeyAlgorithm: "RSA2048",
|
||||
AutoRenew: true,
|
||||
}
|
||||
|
||||
cert, err := ApplyTLSCertificate(certInput)
|
||||
if err != nil {
|
||||
t.Fatalf("ApplyTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
if cert.ApplyStatus != "applying" {
|
||||
t.Fatalf("Expected cert ApplyStatus to be applying, got %s", cert.ApplyStatus)
|
||||
}
|
||||
|
||||
if cert.Provider != "acme" {
|
||||
t.Fatalf("Expected cert Provider to be acme, got %s", cert.Provider)
|
||||
}
|
||||
|
||||
// 3. Try to delete the DNS account (should fail since it's used by the cert)
|
||||
// Actually, the delete logic is in the controller for the foreign key check.
|
||||
// But let's check if the controller logic can be tested here, or we just trust the DB setup.
|
||||
var count int64
|
||||
model.DB.Model(&model.TLSCertificate{}).Where("dns_account_id = ?", dnsAccount.ID).Count(&count)
|
||||
if count != 1 {
|
||||
t.Fatalf("Expected 1 certificate associated with DNS account, got %d", count)
|
||||
}
|
||||
|
||||
// 4. Test RenewTLSCertificate
|
||||
renewedCert, err := RenewTLSCertificate(cert.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("RenewTLSCertificate failed: %v", err)
|
||||
}
|
||||
if renewedCert.ApplyStatus != "applying" {
|
||||
t.Fatalf("Expected renewed cert ApplyStatus to be applying, got %s", renewedCert.ApplyStatus)
|
||||
}
|
||||
|
||||
// Wait for the async goroutine to fail (it now registers an LE account, which takes longer)
|
||||
time.Sleep(5 * time.Second)
|
||||
|
||||
// Reload cert and verify error status
|
||||
finalCert, err := model.GetTLSCertificateByID(renewedCert.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to reload cert: %v", err)
|
||||
}
|
||||
if finalCert.ApplyStatus != "error" {
|
||||
t.Fatalf("Expected final cert ApplyStatus to be error, got %s", finalCert.ApplyStatus)
|
||||
}
|
||||
if finalCert.ApplyMessage == "" {
|
||||
t.Fatalf("Expected final cert ApplyMessage to be populated, got empty")
|
||||
}
|
||||
|
||||
// Clean up
|
||||
if err := DeleteTLSCertificate(cert.ID); err != nil {
|
||||
t.Fatalf("DeleteTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
if err := dnsAccount.Delete(); err != nil {
|
||||
t.Fatalf("Failed to delete DNS Account after cert cleanup: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -25,6 +25,21 @@ type TLSCertificateContent struct {
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
type TLSApplyInput struct {
|
||||
Name string `json:"name"`
|
||||
Remark string `json:"remark"`
|
||||
AcmeAccountID uint `json:"acme_account_id"`
|
||||
DnsAccountID uint `json:"dns_account_id"`
|
||||
KeyAlgorithm string `json:"key_algorithm"`
|
||||
AutoRenew bool `json:"auto_renew"`
|
||||
PrimaryDomain string `json:"primary_domain"`
|
||||
OtherDomains string `json:"other_domains"`
|
||||
DisableCNAME bool `json:"disable_cname"`
|
||||
SkipDNS bool `json:"skip_dns"`
|
||||
DNS1 string `json:"dns1"`
|
||||
DNS2 string `json:"dns2"`
|
||||
}
|
||||
|
||||
func ListTLSCertificates() ([]*model.TLSCertificate, error) {
|
||||
return model.ListTLSCertificates()
|
||||
}
|
||||
@@ -143,6 +158,107 @@ func DeleteTLSCertificate(id uint) error {
|
||||
return certificate.Delete()
|
||||
}
|
||||
|
||||
func ApplyTLSCertificate(input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
cert := &model.TLSCertificate{
|
||||
Name: strings.TrimSpace(input.Name),
|
||||
Remark: strings.TrimSpace(input.Remark),
|
||||
Provider: "acme",
|
||||
AcmeAccountID: input.AcmeAccountID,
|
||||
DnsAccountID: input.DnsAccountID,
|
||||
KeyAlgorithm: input.KeyAlgorithm,
|
||||
AutoRenew: input.AutoRenew,
|
||||
PrimaryDomain: strings.TrimSpace(input.PrimaryDomain),
|
||||
OtherDomains: strings.TrimSpace(input.OtherDomains),
|
||||
DisableCNAME: input.DisableCNAME,
|
||||
SkipDNS: input.SkipDNS,
|
||||
DNS1: strings.TrimSpace(input.DNS1),
|
||||
DNS2: strings.TrimSpace(input.DNS2),
|
||||
ApplyStatus: "applying",
|
||||
CertPEM: " ", // Temporary empty value, since gorm may prevent empty insert
|
||||
KeyPEM: " ", // Temporary empty value
|
||||
}
|
||||
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New("certificate name cannot be empty")
|
||||
}
|
||||
|
||||
if err := cert.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Async obtain SSL
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = ObtainSSL(c)
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func UpdateAcmeCertificate(id uint, input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "acme" {
|
||||
return nil, errors.New("only acme certificates can be updated via this endpoint")
|
||||
}
|
||||
|
||||
cert.Name = strings.TrimSpace(input.Name)
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New("certificate name cannot be empty")
|
||||
}
|
||||
|
||||
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"
|
||||
|
||||
if err := cert.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Async obtain SSL with updated config
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = ObtainSSL(c)
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func RenewTLSCertificate(id uint) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "acme" {
|
||||
return nil, errors.New("only acme certificates can be renewed")
|
||||
}
|
||||
|
||||
// Async obtain SSL
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = ObtainSSL(c)
|
||||
}(cert)
|
||||
|
||||
cert.ApplyStatus = "applying"
|
||||
cert.Update()
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func buildTLSCertificate(existing *model.TLSCertificate, input TLSCertificateInput) (*model.TLSCertificate, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
certPEM := strings.TrimSpace(input.CertPEM)
|
||||
|
||||
Reference in New Issue
Block a user