mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 09:06:36 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -0,0 +1,268 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package acme implements ACME certificate issuance and renewal.
|
||||
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"
|
||||
)
|
||||
|
||||
const dnsChallengePrecheckDelay = 20 * time.Second
|
||||
|
||||
// User implements lego's registration.User interface for ACME account management.
|
||||
type User struct {
|
||||
Email string
|
||||
Registration *registration.Resource
|
||||
key crypto.PrivateKey
|
||||
}
|
||||
|
||||
// GetEmail returns the email address associated with this ACME account.
|
||||
func (u *User) GetEmail() string {
|
||||
return u.Email
|
||||
}
|
||||
|
||||
// GetRegistration returns the ACME account registration resource.
|
||||
func (u *User) GetRegistration() *registration.Resource {
|
||||
return u.Registration
|
||||
}
|
||||
|
||||
// GetPrivateKey returns the private key used to authenticate with the ACME server.
|
||||
func (u *User) 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, *User, 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 := &User{
|
||||
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: %w", 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.DisableAuthoritativeNssPropagationRequirement())
|
||||
}
|
||||
|
||||
if skipDNS {
|
||||
opts = append(opts, dns01.WrapPreCheck(func(_, _, _ string, _ dns01.PreCheckFunc) (bool, error) {
|
||||
time.Sleep(dnsChallengePrecheckDelay)
|
||||
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,166 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/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)
|
||||
|
||||
obtainDone := make(chan struct{})
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
defer close(obtainDone)
|
||||
return updateCertError(ctx, cert, "dns challenge failed")
|
||||
})
|
||||
defer func() {
|
||||
select {
|
||||
case <-obtainDone:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("async certificate obtain did not finish")
|
||||
}
|
||||
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, repository.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 := repository.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, repository.SaveTLSCertificate(ctx, stored))
|
||||
|
||||
err = updateCertError(ctx, stored, "dns challenge failed")
|
||||
require.Error(t, err)
|
||||
|
||||
finalCert, err := repository.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.Contains(t, 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, repository.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, repository.SaveTLSCertificate(ctx, cert))
|
||||
_, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already applying")
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package tls defines shared error messages for certificate management.
|
||||
package tls
|
||||
|
||||
const (
|
||||
errCertificateNameRequired = "certificate name cannot be empty"
|
||||
errCertificateNameExists = "certificate name already exists"
|
||||
errCertificateContentRequired = "certificate content and key content cannot be empty"
|
||||
errCertificateContentInvalid = "certificate or key format is invalid"
|
||||
errCertificateDeleteReferenced = "certificate is still referenced by proxy routes"
|
||||
errCertificateOnlyACME = "only acme certificates can be updated via this endpoint"
|
||||
errCertificateOnlyUploadConvert = "only uploaded certificates can be converted to acme"
|
||||
errCertificateAlreadyApplying = "certificate is already applying"
|
||||
errCertificateOnlyACMERenew = "only acme certificates can be renewed"
|
||||
errCertificateFilesRequired = "certificate file and key file cannot be empty"
|
||||
errCertificatePEMInvalid = "证书 PEM 内容不合法"
|
||||
|
||||
errDNSAccountInUse = "该 DNS 账号已被证书使用,无法删除"
|
||||
)
|
||||
@@ -0,0 +1,45 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func parseLeafCertificate(certPEM string) (*x509.Certificate, error) {
|
||||
certPEMBlock, _ := pem.Decode([]byte(certPEM))
|
||||
if certPEMBlock == nil {
|
||||
return nil, errors.New(errCertificatePEMInvalid)
|
||||
}
|
||||
leaf, err := x509.ParseCertificate(certPEMBlock.Bytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return leaf, nil
|
||||
}
|
||||
|
||||
func readMultipartFile(fileHeader *multipart.FileHeader) (string, error) {
|
||||
file, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
data, err := io.ReadAll(file)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
@@ -0,0 +1,480 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"strings"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/task"
|
||||
)
|
||||
|
||||
// CertificateInput TLS 证书创建/更新请求。
|
||||
type CertificateInput struct {
|
||||
Name string `json:"name"`
|
||||
CertPEM string `json:"cert_pem"`
|
||||
KeyPEM string `json:"key_pem"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
// CertificateContent TLS 证书 PEM 内容(仅 /content 端点返回)。
|
||||
type CertificateContent struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
CertPEM string `json:"cert_pem"`
|
||||
KeyPEM string `json:"key_pem"`
|
||||
Remark string `json:"remark"`
|
||||
Provider string `json:"provider"`
|
||||
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"`
|
||||
ApplyStatus string `json:"apply_status"`
|
||||
ApplyMessage string `json:"apply_message"`
|
||||
}
|
||||
|
||||
// ApplyInput ACME 证书申请/更新请求。
|
||||
type ApplyInput 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"`
|
||||
}
|
||||
|
||||
// DNSAccountInput DNS 账号创建/更新请求。
|
||||
type DNSAccountInput struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Authorization string `json:"authorization"`
|
||||
}
|
||||
|
||||
// ListCertificates 列出全部证书(不含 PEM)。
|
||||
func ListCertificates(ctx context.Context) ([]model.TLSCertificate, error) {
|
||||
return repository.ListTLSCertificates(ctx)
|
||||
}
|
||||
|
||||
// GetCertificate 获取证书详情(不含 PEM)。
|
||||
func GetCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) {
|
||||
return repository.GetTLSCertificateByID(ctx, id)
|
||||
}
|
||||
|
||||
// GetCertificateContent 获取证书 PEM 内容。
|
||||
func GetCertificateContent(ctx context.Context, id uint) (*CertificateContent, error) {
|
||||
certificate, err := repository.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyPEM, err := openSensitive(certificate.KeyPEM)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &CertificateContent{
|
||||
ID: certificate.ID,
|
||||
Name: certificate.Name,
|
||||
CertPEM: certificate.CertPEM,
|
||||
KeyPEM: keyPEM,
|
||||
Remark: certificate.Remark,
|
||||
Provider: certificate.Provider,
|
||||
AcmeAccountID: certificate.AcmeAccountID,
|
||||
DNSAccountID: certificate.DNSAccountID,
|
||||
KeyAlgorithm: certificate.KeyAlgorithm,
|
||||
AutoRenew: certificate.AutoRenew,
|
||||
PrimaryDomain: certificate.PrimaryDomain,
|
||||
OtherDomains: certificate.OtherDomains,
|
||||
DisableCNAME: certificate.DisableCNAME,
|
||||
SkipDNS: certificate.SkipDNS,
|
||||
DNS1: certificate.DNS1,
|
||||
DNS2: certificate.DNS2,
|
||||
ApplyStatus: certificate.ApplyStatus,
|
||||
ApplyMessage: certificate.ApplyMessage,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// CreateCertificate 从 PEM 创建证书。
|
||||
func CreateCertificate(ctx context.Context, input CertificateInput) (*model.TLSCertificate, error) {
|
||||
certificate, err := buildCertificate(ctx, nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = repository.CreateTLSCertificateRecord(ctx, certificate); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errCertificateNameExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeCertificateForResponse(certificate), nil
|
||||
}
|
||||
|
||||
// CreateCertificateFromFiles 从上传文件创建证书。
|
||||
func CreateCertificateFromFiles(ctx context.Context, name string, certFile *multipart.FileHeader, keyFile *multipart.FileHeader, remark string) (*model.TLSCertificate, error) {
|
||||
if certFile == nil || keyFile == nil {
|
||||
return nil, errors.New(errCertificateFilesRequired)
|
||||
}
|
||||
certContent, err := readMultipartFile(certFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyContent, err := readMultipartFile(keyFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return CreateCertificate(ctx, CertificateInput{
|
||||
Name: name,
|
||||
CertPEM: certContent,
|
||||
KeyPEM: keyContent,
|
||||
Remark: remark,
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateCertificate 更新上传证书。
|
||||
func UpdateCertificate(ctx context.Context, id uint, input CertificateInput) (*model.TLSCertificate, error) {
|
||||
existing, err := repository.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certificate, err := buildCertificate(ctx, existing, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = repository.SaveTLSCertificate(ctx, certificate); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errCertificateNameExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeCertificateForResponse(certificate), nil
|
||||
}
|
||||
|
||||
// DeleteCertificate 删除证书。
|
||||
func DeleteCertificate(ctx context.Context, id uint) error {
|
||||
if err := ensureCertificateNotReferenced(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := repository.GetTLSCertificateByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return repository.DeleteTLSCertificateRecord(ctx, id)
|
||||
}
|
||||
|
||||
// ApplyCertificate 申请 ACME 证书。
|
||||
func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertificate, error) {
|
||||
cert := &model.TLSCertificate{
|
||||
Provider: tlsProviderACME,
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
}
|
||||
fillAcmeCertificateFields(cert, input)
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New(errCertificateNameRequired)
|
||||
}
|
||||
if err := repository.CreateTLSCertificateRecord(ctx, cert); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errCertificateNameExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 先取响应快照再启动异步续签:sanitize 会整体拷贝 cert,若与异步 goroutine
|
||||
// 的字段写入并发会构成数据竞争(生产真实问题)。
|
||||
returned := sanitizeCertificateForResponse(cert)
|
||||
|
||||
obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换)
|
||||
go func(c *model.TLSCertificate) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
_ = obtainFn(asyncCtx, c)
|
||||
}(cert)
|
||||
|
||||
return returned, nil
|
||||
}
|
||||
|
||||
// UpdateACMECertificate 更新 ACME 证书配置。
|
||||
func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := repository.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != tlsProviderACME {
|
||||
return nil, errors.New(errCertificateOnlyACME)
|
||||
}
|
||||
fillAcmeCertificateFields(cert, input)
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New(errCertificateNameRequired)
|
||||
}
|
||||
if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errCertificateNameExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
returned := sanitizeCertificateForResponse(cert)
|
||||
|
||||
obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换)
|
||||
go func(c *model.TLSCertificate) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
_ = obtainFn(asyncCtx, c)
|
||||
}(cert)
|
||||
|
||||
return returned, nil
|
||||
}
|
||||
|
||||
// ConvertCertificateToACME 将上传证书转为 ACME 管理。
|
||||
func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := repository.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "upload" {
|
||||
return nil, errors.New(errCertificateOnlyUploadConvert)
|
||||
}
|
||||
if cert.ApplyStatus == tlsApplyStatusApplying {
|
||||
return nil, errors.New(errCertificateAlreadyApplying)
|
||||
}
|
||||
fillAcmeCertificateFields(cert, input)
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New(errCertificateNameRequired)
|
||||
}
|
||||
cert.ApplyMessage = ""
|
||||
if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errCertificateNameExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换)
|
||||
go func(c *model.TLSCertificate) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
if err := obtainFn(asyncCtx, c); err != nil {
|
||||
return
|
||||
}
|
||||
latest, err := repository.GetTLSCertificateByID(asyncCtx, c.ID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
latest.Provider = tlsProviderACME
|
||||
latest.ApplyStatus = tlsApplyStatusReady
|
||||
latest.ApplyMessage = ""
|
||||
_ = repository.SaveTLSCertificate(asyncCtx, latest)
|
||||
}(cert)
|
||||
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
// RenewCertificate 续期 ACME 证书。
|
||||
func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) {
|
||||
cert, err := repository.GetTLSCertificateByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != tlsProviderACME {
|
||||
return nil, errors.New(errCertificateOnlyACMERenew)
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(SSLSingleRenewPayload{ID: id})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
_, err = task.DispatchTask(ctx, TaskTypeSSLSingleRenew, payload, "manual")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cert.ApplyStatus = tlsApplyStatusApplying
|
||||
cert.ApplyMessage = ""
|
||||
if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
}
|
||||
|
||||
// ListDNSAccounts 列出 DNS 账号。
|
||||
func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) {
|
||||
return repository.ListDNSAccounts(ctx)
|
||||
}
|
||||
|
||||
// CreateDNSAccount 创建 DNS 账号。
|
||||
func CreateDNSAccount(ctx context.Context, input DNSAccountInput) (*model.DNSAccount, error) {
|
||||
authorization, err := sealSensitive(strings.TrimSpace(input.Authorization))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
account := &model.DNSAccount{
|
||||
Name: strings.TrimSpace(input.Name),
|
||||
Type: strings.TrimSpace(input.Type),
|
||||
Authorization: authorization,
|
||||
}
|
||||
if account.Name == "" || account.Type == "" || authorization == "" {
|
||||
return nil, errors.New("DNS 账号参数不完整")
|
||||
}
|
||||
if err := repository.CreateDNSAccountRecord(ctx, account); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeDNSAccountForResponse(account), nil
|
||||
}
|
||||
|
||||
// UpdateDNSAccount 更新 DNS 账号。
|
||||
func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*model.DNSAccount, error) {
|
||||
account, err := repository.GetDNSAccountByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
authorization, err := sealSensitive(strings.TrimSpace(input.Authorization))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
account.Name = strings.TrimSpace(input.Name)
|
||||
account.Type = strings.TrimSpace(input.Type)
|
||||
account.Authorization = authorization
|
||||
if account.Name == "" || account.Type == "" || authorization == "" {
|
||||
return nil, errors.New("DNS 账号参数不完整")
|
||||
}
|
||||
if err := repository.SaveDNSAccount(ctx, account); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeDNSAccountForResponse(account), nil
|
||||
}
|
||||
|
||||
// DeleteDNSAccount 删除 DNS 账号。
|
||||
func DeleteDNSAccount(ctx context.Context, id uint) error {
|
||||
if _, err := repository.GetDNSAccountByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
count, err := repository.CountTLSCertificatesByDNSAccountID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(errDNSAccountInUse)
|
||||
}
|
||||
return repository.DeleteDNSAccountRecord(ctx, id)
|
||||
}
|
||||
|
||||
// GetDefaultAcmeAccount 获取默认 ACME 账号。
|
||||
func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) {
|
||||
account, err := repository.GetDefaultAcmeAccount(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sanitizeAcmeAccountForResponse(account), nil
|
||||
}
|
||||
|
||||
func buildCertificate(_ context.Context, existing *model.TLSCertificate, input CertificateInput) (*model.TLSCertificate, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
certPEM := strings.TrimSpace(input.CertPEM)
|
||||
keyPEM := strings.TrimSpace(input.KeyPEM)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
if name == "" {
|
||||
return nil, errors.New(errCertificateNameRequired)
|
||||
}
|
||||
if certPEM == "" || keyPEM == "" {
|
||||
return nil, errors.New(errCertificateContentRequired)
|
||||
}
|
||||
parsed, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", errCertificateContentInvalid, err)
|
||||
}
|
||||
if len(parsed.Certificate) == 0 {
|
||||
return nil, errors.New(errCertificateContentInvalid)
|
||||
}
|
||||
leaf, err := parseLeafCertificate(certPEM)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sealedKey, err := sealSensitive(keyPEM)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.TLSCertificate{
|
||||
Provider: "upload",
|
||||
ApplyStatus: tlsApplyStatusReady,
|
||||
}
|
||||
}
|
||||
existing.Name = name
|
||||
existing.CertPEM = certPEM
|
||||
existing.KeyPEM = sealedKey
|
||||
existing.NotBefore = leaf.NotBefore
|
||||
existing.NotAfter = leaf.NotAfter
|
||||
existing.Remark = remark
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
|
||||
cert.Name = strings.TrimSpace(input.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 = tlsApplyStatusApplying
|
||||
}
|
||||
|
||||
func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
|
||||
count, err := repository.CountZoneDomainsByCertificateID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(errCertificateDeleteReferenced)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func sanitizeCertificateForResponse(certificate *model.TLSCertificate) *model.TLSCertificate {
|
||||
if certificate == nil {
|
||||
return nil
|
||||
}
|
||||
certCopy := *certificate
|
||||
certCopy.CertPEM = ""
|
||||
certCopy.KeyPEM = ""
|
||||
return &certCopy
|
||||
}
|
||||
|
||||
func sanitizeDNSAccountForResponse(account *model.DNSAccount) *model.DNSAccount {
|
||||
if account == nil {
|
||||
return nil
|
||||
}
|
||||
certCopy := *account
|
||||
certCopy.Authorization = ""
|
||||
return &certCopy
|
||||
}
|
||||
|
||||
func sanitizeAcmeAccountForResponse(account *model.AcmeAccount) *model.AcmeAccount {
|
||||
if account == nil {
|
||||
return nil
|
||||
}
|
||||
certCopy := *account
|
||||
certCopy.PrivateKey = ""
|
||||
return &certCopy
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/credential"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
oftask "Wavelet/openflare/plugins/server/kernel/task"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
|
||||
"Wavelet/pkg/idgen"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"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,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.TLSCertificate{},
|
||||
&model.Zone{},
|
||||
&model.ZoneDomain{},
|
||||
&model.DNSAccount{},
|
||||
&model.AcmeAccount{},
|
||||
&model.TaskExecution{}, // 异步任务执行记录也需要 migrate
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
require.NoError(t, idgen.Init(1))
|
||||
previous := runtimeconfig.Get()
|
||||
runtimeconfig.SetSessionSecret("test_session_secret_for_tls_encryption")
|
||||
credential.SetSessionSecret("test_session_secret_for_tls_encryption")
|
||||
oftask.SetService(&testhelper.NoopTaskService{})
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
runtimeconfig.Set(previous)
|
||||
credential.SetSessionSecret(previous.SessionSecret)
|
||||
tlsTestDBMu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteCertificateRejectsZoneDomainReference(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
certPEM, keyPEM := generateTestCertificatePair(t, []string{"api.example.com"})
|
||||
certificate, err := CreateCertificate(ctx, CertificateInput{Name: "api-cert", CertPEM: certPEM, KeyPEM: keyPEM})
|
||||
require.NoError(t, err)
|
||||
zone := &model.Zone{Domain: "example.com"}
|
||||
require.NoError(t, db.DB(ctx).Create(zone).Error)
|
||||
require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, Domain: "api.example.com", CertID: &certificate.ID}).Error)
|
||||
|
||||
err = DeleteCertificate(ctx, certificate.ID)
|
||||
require.EqualError(t, err, errCertificateDeleteReferenced)
|
||||
}
|
||||
|
||||
func generateTestCertificatePair(t *testing.T, dnsNames []string) (string, string) {
|
||||
t.Helper()
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
require.NoError(t, err)
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(time.Now().UnixNano()),
|
||||
Subject: pkix.Name{
|
||||
CommonName: dnsNames[0],
|
||||
},
|
||||
DNSNames: dnsNames,
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
}
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
|
||||
require.NoError(t, err)
|
||||
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
|
||||
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
|
||||
return string(certPEM), string(keyPEM)
|
||||
}
|
||||
|
||||
func TestCreateCertificateEncryptsPrivateKey(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
certPEM, keyPEM := generateTestCertificatePair(t, []string{"secure.example.com"})
|
||||
certificate, err := CreateCertificate(ctx, CertificateInput{
|
||||
Name: "secure-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := repository.GetTLSCertificateByID(ctx, certificate.ID)
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, keyPEM, stored.KeyPEM)
|
||||
assert.Contains(t, stored.KeyPEM, credential.Prefix)
|
||||
|
||||
content, err := GetCertificateContent(ctx, certificate.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, strings.TrimSpace(keyPEM), strings.TrimSpace(content.KeyPEM))
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/tls/acme"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/task"
|
||||
)
|
||||
|
||||
const (
|
||||
acmeRenewLeadTime = 7 * 24 * time.Hour
|
||||
tlsProviderACME = "acme"
|
||||
tlsApplyStatusApplying = "applying"
|
||||
tlsApplyStatusReady = "ready"
|
||||
)
|
||||
|
||||
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 {
|
||||
task.AppendLog(ctx, "【续签任务】开始续签,设置申请状态为 applying...")
|
||||
cert.ApplyStatus = tlsApplyStatusApplying
|
||||
if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "【续签任务】正在解析 ACME 账户...")
|
||||
acmeAccount, err := resolveAcmeAccount(ctx, cert)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get ACME account: %v", err))
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "【续签任务】正在解析 DNS 账户信息 (ID=%d)...", cert.DNSAccountID)
|
||||
dnsAccount, err := repository.GetDNSAccountByID(ctx, cert.DNSAccountID)
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get DNS account: %v", err))
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "【续签任务】正在解密 DNS 账号凭据及 ACME 账户私钥...")
|
||||
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)
|
||||
task.AppendLog(ctx, "【续签任务】待申请的域名列表: %v", domains)
|
||||
|
||||
task.AppendLog(ctx, "【续签任务】正在调用 ACME 客户端(通过 DNS-01 挑战)发起 SSL 证书签发请求,请稍候...")
|
||||
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,
|
||||
)
|
||||
|
||||
task.AppendLog(ctx, "【续签任务】正在保存 ACME 账户可能的变更...")
|
||||
if err := persistAcmeAccountUpdates(ctx, cert, acmeAccount, newAccountURL, newPrivateKeyPEM, acmePrivateKey); err != nil {
|
||||
return updateCertError(ctx, cert, err.Error())
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return updateCertError(ctx, cert, err.Error())
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "【续签任务】证书签发成功,正在将证书内容与私钥安全写入数据库...")
|
||||
if err := saveObtainedCertificate(ctx, cert, result); err != nil {
|
||||
return updateCertError(ctx, cert, err.Error())
|
||||
}
|
||||
task.AppendLog(ctx, "【续签任务】证书数据存储完成!")
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateCertError(ctx context.Context, cert *model.TLSCertificate, message string) error {
|
||||
cert.ApplyStatus = "error"
|
||||
cert.ApplyMessage = message
|
||||
if err := repository.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.SplitSeq(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 != tlsProviderACME || cert.ApplyStatus == tlsApplyStatusApplying {
|
||||
continue
|
||||
}
|
||||
if cert.NotAfter.IsZero() {
|
||||
continue
|
||||
}
|
||||
if cert.NotAfter.Sub(now) < acmeRenewLeadTime {
|
||||
due = append(due, cert)
|
||||
}
|
||||
}
|
||||
return due
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/tls/acme"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
)
|
||||
|
||||
func resolveAcmeAccount(ctx context.Context, cert *model.TLSCertificate) (*model.AcmeAccount, error) {
|
||||
acmeAccount, err := repository.GetAcmeAccountByID(ctx, cert.AcmeAccountID)
|
||||
if err == nil {
|
||||
return acmeAccount, nil
|
||||
}
|
||||
acmeAccount, err = repository.GetDefaultAcmeAccount(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get ACME account: %w", err)
|
||||
}
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return acmeAccount, nil
|
||||
}
|
||||
|
||||
func persistAcmeAccountUpdates(
|
||||
ctx context.Context,
|
||||
cert *model.TLSCertificate,
|
||||
acmeAccount *model.AcmeAccount,
|
||||
newAccountURL, newPrivateKeyPEM, acmePrivateKey string,
|
||||
) error {
|
||||
accountChanged := (newPrivateKeyPEM != "" && acmePrivateKey != newPrivateKeyPEM) ||
|
||||
(newAccountURL != "" && acmeAccount.URL != newAccountURL)
|
||||
if !accountChanged {
|
||||
return nil
|
||||
}
|
||||
if newPrivateKeyPEM != "" && acmePrivateKey != newPrivateKeyPEM {
|
||||
sealedKey, sealErr := sealSensitive(newPrivateKeyPEM)
|
||||
if sealErr != nil {
|
||||
return fmt.Errorf("failed to seal ACME account key: %w", sealErr)
|
||||
}
|
||||
acmeAccount.PrivateKey = sealedKey
|
||||
}
|
||||
if newAccountURL != "" {
|
||||
acmeAccount.URL = newAccountURL
|
||||
}
|
||||
if acmeAccount.ID == 0 {
|
||||
if dbErr := repository.CreateAcmeAccountRecord(ctx, acmeAccount); dbErr != nil {
|
||||
return fmt.Errorf("failed to create ACME account: %w", dbErr)
|
||||
}
|
||||
} else if dbErr := repository.SaveAcmeAccount(ctx, acmeAccount); dbErr != nil {
|
||||
return fmt.Errorf("failed to save ACME account: %w", dbErr)
|
||||
}
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
return repository.SaveTLSCertificate(ctx, cert)
|
||||
}
|
||||
|
||||
func saveObtainedCertificate(ctx context.Context, cert *model.TLSCertificate, result *acme.CertificateResult) error {
|
||||
sealedKey, err := sealSensitive(result.KeyPEM)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to seal certificate key: %w", err)
|
||||
}
|
||||
cert.CertPEM = result.CertPEM
|
||||
cert.KeyPEM = sealedKey
|
||||
cert.NotBefore = result.NotBefore
|
||||
cert.NotAfter = result.NotAfter
|
||||
cert.ApplyStatus = tlsApplyStatusReady
|
||||
cert.ApplyMessage = ""
|
||||
return repository.SaveTLSCertificate(ctx, cert)
|
||||
}
|
||||
@@ -0,0 +1,452 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/apiutil"
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, "记录不存在")
|
||||
}
|
||||
|
||||
// GetCertificates 列出 TLS 证书。
|
||||
// @Summary 列出 TLS 证书
|
||||
// @Description 返回全部 TLS 证书(不含 PEM),需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.TLSCertificate} "证书列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates [get]
|
||||
func GetCertificates(c *gin.Context) {
|
||||
certificates, err := ListCertificates(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificates))
|
||||
}
|
||||
|
||||
// GetCertificateDetail 获取 TLS 证书详情。
|
||||
// @Summary 获取 TLS 证书详情
|
||||
// @Description 按 ID 返回 TLS 证书详情(不含 PEM),需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "证书详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id} [get]
|
||||
func GetCertificateDetail(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
certificate, err := GetCertificate(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// GetCertificateContentHandler 获取 TLS 证书 PEM 内容。
|
||||
// @Summary 获取 TLS 证书 PEM 内容
|
||||
// @Description 按 ID 返回证书与私钥 PEM 内容,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Success 200 {object} response.Any{data=tls.CertificateContent} "证书 PEM 内容"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id}/content [get]
|
||||
func GetCertificateContentHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
content, err := GetCertificateContent(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(content))
|
||||
}
|
||||
|
||||
// CreateCertificateHandler 从 PEM 创建证书。
|
||||
// @Summary 创建 TLS 证书
|
||||
// @Description 从 PEM 文本创建 TLS 证书,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body tls.CertificateInput true "证书参数"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "创建成功的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates [post]
|
||||
func CreateCertificateHandler(c *gin.Context) {
|
||||
var input CertificateInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
certificate, err := CreateCertificate(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// UpdateCertificateHandler 更新证书。
|
||||
// @Summary 更新 TLS 证书
|
||||
// @Description 按 ID 更新 TLS 证书 PEM 信息,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Param request body tls.CertificateInput true "证书参数"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "更新后的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id}/update [post]
|
||||
func UpdateCertificateHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input CertificateInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
certificate, err := UpdateCertificate(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// ImportCertificateFile 从文件导入证书。
|
||||
// @Summary 从文件导入 TLS 证书
|
||||
// @Description 上传证书与私钥文件创建 TLS 证书,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept multipart/form-data
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param name formData string false "证书名称"
|
||||
// @Param remark formData string false "备注"
|
||||
// @Param cert_file formData file true "证书文件"
|
||||
// @Param key_file formData file true "私钥文件"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "导入成功的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/import-file [post]
|
||||
func ImportCertificateFile(c *gin.Context) {
|
||||
name := c.PostForm("name")
|
||||
remark := c.PostForm("remark")
|
||||
certFile, err := c.FormFile("cert_file")
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "缺少证书文件")
|
||||
return
|
||||
}
|
||||
keyFile, err := c.FormFile("key_file")
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "缺少私钥文件")
|
||||
return
|
||||
}
|
||||
certificate, err := CreateCertificateFromFiles(c.Request.Context(), name, certFile, keyFile, remark)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// DeleteCertificateHandler 删除证书。
|
||||
// @Summary 删除 TLS 证书
|
||||
// @Description 按 ID 删除 TLS 证书,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Success 200 {object} response.Any "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id}/delete [post]
|
||||
func DeleteCertificateHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteCertificate(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ApplyCertificateHandler 申请 ACME 证书。
|
||||
// @Summary 申请 ACME 证书
|
||||
// @Description 通过 ACME 申请新的 TLS 证书,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body tls.ApplyInput true "ACME 申请参数"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "申请中的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/apply [post]
|
||||
func ApplyCertificateHandler(c *gin.Context) {
|
||||
var input ApplyInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
certificate, err := ApplyCertificate(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// UpdateACMECertificateHandler 更新 ACME 证书配置。
|
||||
// @Summary 更新 ACME 证书配置
|
||||
// @Description 按 ID 更新 ACME 证书申请配置,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Param request body tls.ApplyInput true "ACME 申请参数"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "更新后的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id}/update-acme [post]
|
||||
func UpdateACMECertificateHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input ApplyInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
certificate, err := UpdateACMECertificate(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// ConvertCertificateToACMEHandler 将上传证书转为 ACME。
|
||||
// @Summary 将证书转为 ACME 管理
|
||||
// @Description 将已上传证书转换为 ACME 自动续期模式,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Param request body tls.ApplyInput true "ACME 申请参数"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "转换后的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id}/convert-acme [post]
|
||||
func ConvertCertificateToACMEHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input ApplyInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
certificate, err := ConvertCertificateToACME(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// RenewCertificateHandler 续期 ACME 证书。
|
||||
// @Summary 续期 ACME 证书
|
||||
// @Description 手动触发 ACME 证书续期,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "证书 ID"
|
||||
// @Success 200 {object} response.Any{data=model.TLSCertificate} "续期后的证书"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/tls-certificates/{id}/renew [post]
|
||||
func RenewCertificateHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
certificate, err := RenewCertificate(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(certificate))
|
||||
}
|
||||
|
||||
// GetDNSAccounts 列出 DNS 账号。
|
||||
// @Summary 列出 DNS 账号
|
||||
// @Description 返回全部 DNS 提供商账号,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.DNSAccount} "DNS 账号列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/dns-accounts [get]
|
||||
func GetDNSAccounts(c *gin.Context) {
|
||||
accounts, err := ListDNSAccounts(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(accounts))
|
||||
}
|
||||
|
||||
// CreateDNSAccountHandler 创建 DNS 账号。
|
||||
// @Summary 创建 DNS 账号
|
||||
// @Description 创建新的 DNS 提供商账号,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body tls.DNSAccountInput true "DNS 账号参数"
|
||||
// @Success 200 {object} response.Any{data=model.DNSAccount} "创建成功的 DNS 账号"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/dns-accounts [post]
|
||||
func CreateDNSAccountHandler(c *gin.Context) {
|
||||
var input DNSAccountInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
account, err := CreateDNSAccount(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(account))
|
||||
}
|
||||
|
||||
// UpdateDNSAccountHandler 更新 DNS 账号。
|
||||
// @Summary 更新 DNS 账号
|
||||
// @Description 按 ID 更新 DNS 提供商账号,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "DNS 账号 ID"
|
||||
// @Param request body tls.DNSAccountInput true "DNS 账号参数"
|
||||
// @Success 200 {object} response.Any{data=model.DNSAccount} "更新后的 DNS 账号"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/dns-accounts/{id}/update [post]
|
||||
func UpdateDNSAccountHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input DNSAccountInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
account, err := UpdateDNSAccount(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(account))
|
||||
}
|
||||
|
||||
// DeleteDNSAccountHandler 删除 DNS 账号。
|
||||
// @Summary 删除 DNS 账号
|
||||
// @Description 按 ID 删除 DNS 提供商账号,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "DNS 账号 ID"
|
||||
// @Success 200 {object} response.Any "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/dns-accounts/{id}/delete [post]
|
||||
func DeleteDNSAccountHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteDNSAccount(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// GetDefaultAcmeAccountHandler 获取默认 ACME 账号。
|
||||
// @Summary 获取默认 ACME 账号
|
||||
// @Description 返回系统默认 ACME 账号配置,需要管理员权限
|
||||
// @Tags openflare-tls
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=model.AcmeAccount} "默认 ACME 账号"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/acme-accounts/default [get]
|
||||
func GetDefaultAcmeAccountHandler(c *gin.Context) {
|
||||
account, err := GetDefaultAcmeAccount(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(account))
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import "Wavelet/openflare/plugins/server/kernel/credential"
|
||||
|
||||
func sealSensitive(plaintext string) (string, error) {
|
||||
return credential.Seal(plaintext)
|
||||
}
|
||||
|
||||
// OpenKeyPEM decrypts a stored certificate private key for runtime distribution.
|
||||
func OpenKeyPEM(stored string) (string, error) {
|
||||
return openSensitive(stored)
|
||||
}
|
||||
|
||||
func openSensitive(stored string) (string, error) {
|
||||
return credential.Open(stored)
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/pkg/logger"
|
||||
)
|
||||
|
||||
// RunSSLRenewJob renews all TLS certificates that are due for renewal.
|
||||
func RunSSLRenewJob(ctx context.Context) error {
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job started")
|
||||
|
||||
certificates, err := repository.ListTLSCertificates(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "[OpenFlareTasks] list certificates failed: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
due := CertificatesDueForRenewal(certificates, now)
|
||||
if len(due) == 0 {
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job completed: no certificates due")
|
||||
return nil
|
||||
}
|
||||
|
||||
var triggered int
|
||||
for _, cert := range due {
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] renewing certificate id=%d domain=%s", cert.ID, cert.PrimaryDomain)
|
||||
if _, err := RenewCertificate(ctx, cert.ID); err != nil {
|
||||
logger.ErrorF(ctx, "[OpenFlareTasks] renew certificate id=%d domain=%s failed: %v", cert.ID, cert.PrimaryDomain, err)
|
||||
continue
|
||||
}
|
||||
triggered++
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job completed: triggered=%d eligible=%d", triggered, len(due))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
|
||||
oftask "Wavelet/openflare/plugins/server/kernel/task"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func setupSSLRenewTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
require.NoError(t, db.DB(context.Background()).AutoMigrate(&model.TLSCertificate{}, &model.TaskExecution{}))
|
||||
previous := runtimeconfig.Get()
|
||||
runtimeconfig.SetSessionSecret("test_session_secret_for_ssl_renew")
|
||||
oftask.SetService(&testhelper.NoopTaskService{})
|
||||
return func() {
|
||||
runtimeconfig.Set(previous)
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunSSLRenewJobTriggersDueCertificates(t *testing.T) {
|
||||
cleanup := setupSSLRenewTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
return nil
|
||||
})
|
||||
defer restore()
|
||||
|
||||
now := time.Now().UTC()
|
||||
due := &model.TLSCertificate{
|
||||
Name: "due-cert",
|
||||
Provider: "acme",
|
||||
AutoRenew: true,
|
||||
ApplyStatus: "ready",
|
||||
PrimaryDomain: "due.example.com",
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
NotAfter: now.Add(2 * 24 * time.Hour),
|
||||
}
|
||||
fresh := &model.TLSCertificate{
|
||||
Name: "fresh-cert",
|
||||
Provider: "acme",
|
||||
AutoRenew: true,
|
||||
ApplyStatus: "ready",
|
||||
PrimaryDomain: "fresh.example.com",
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
NotAfter: now.Add(30 * 24 * time.Hour),
|
||||
}
|
||||
require.NoError(t, repository.CreateTLSCertificateRecord(ctx, due))
|
||||
require.NoError(t, repository.CreateTLSCertificateRecord(ctx, fresh))
|
||||
|
||||
require.NoError(t, RunSSLRenewJob(ctx))
|
||||
|
||||
renewed, err := repository.GetTLSCertificateByID(ctx, due.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "applying", renewed.ApplyStatus)
|
||||
|
||||
unchanged, err := repository.GetTLSCertificateByID(ctx, fresh.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ready", unchanged.ApplyStatus)
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/task"
|
||||
)
|
||||
|
||||
const (
|
||||
// SSLSingleRenewTask renews a single ACME TLS certificate.
|
||||
SSLSingleRenewTask = "openflare:ssl_single_renew"
|
||||
// TaskTypeSSLSingleRenew is the admin task type for single SSL renewal.
|
||||
TaskTypeSSLSingleRenew = "of_ssl_single_renew"
|
||||
)
|
||||
|
||||
// SSLSingleRenewMeta describes the single SSL renewal task.
|
||||
var SSLSingleRenewMeta = task.TaskMeta{
|
||||
Type: TaskTypeSSLSingleRenew,
|
||||
AsynqTask: SSLSingleRenewTask,
|
||||
Name: "OpenFlare 单证书 SSL 续期",
|
||||
Description: "对单个指定的 ACME 证书执行续期",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
{
|
||||
Name: "id",
|
||||
Label: "证书 ID",
|
||||
Type: "number",
|
||||
Required: true,
|
||||
Placeholder: "请输入证书 ID",
|
||||
Description: "待续期的 TLS 证书 ID",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// SSLSingleRenewPayload is the payload structure for SSLSingleRenewTask.
|
||||
type SSLSingleRenewPayload struct {
|
||||
ID uint `json:"id"`
|
||||
}
|
||||
|
||||
// SSLSingleRenewHandler renews a specific TLS certificate.
|
||||
type SSLSingleRenewHandler struct{}
|
||||
|
||||
// ValidatePayload validates and normalizes the task payload.
|
||||
func (h *SSLSingleRenewHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
if len(payload) == 0 {
|
||||
return nil, errors.New("任务参数不能为空")
|
||||
}
|
||||
|
||||
var req SSLSingleRenewPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
return nil, fmt.Errorf("无效的 JSON 格式: %w", err)
|
||||
}
|
||||
|
||||
if req.ID == 0 {
|
||||
return nil, errors.New("证书 ID 不能为空或零")
|
||||
}
|
||||
|
||||
return json.Marshal(req)
|
||||
}
|
||||
|
||||
// Execute runs the certificate renewal for the specified ID.
|
||||
func (h *SSLSingleRenewHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
var req SSLSingleRenewPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
return nil, fmt.Errorf("解析任务参数: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始续期证书,ID: %d", req.ID)
|
||||
|
||||
cert, err := repository.GetTLSCertificateByID(ctx, req.ID)
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "获取证书记录失败 ID=%d: %v", req.ID, err)
|
||||
return nil, fmt.Errorf("获取证书记录失败: %w", err)
|
||||
}
|
||||
|
||||
if cert.Provider != tlsProviderACME {
|
||||
task.AppendLog(ctx, "证书 %s (ID=%d) 不是 ACME 托管证书,无法自动续期 (Provider: %s)", cert.PrimaryDomain, req.ID, cert.Provider)
|
||||
return nil, fmt.Errorf("证书 %s 不是 ACME 托管证书,无法自动续期", cert.PrimaryDomain)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "准备为域名 [%s] 申请/续期证书", cert.PrimaryDomain)
|
||||
if err := obtainTLSCertificate(ctx, cert); err != nil {
|
||||
task.AppendLog(ctx, "申请证书失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("证书 %s 续签成功", cert.PrimaryDomain)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSSLSingleRenewHandler_ValidatePayload(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
payload []byte
|
||||
wantErr bool
|
||||
errMsg string
|
||||
}{
|
||||
{
|
||||
name: "empty payload",
|
||||
payload: nil,
|
||||
wantErr: true,
|
||||
errMsg: "任务参数不能为空",
|
||||
},
|
||||
{
|
||||
name: "invalid JSON",
|
||||
payload: []byte(`{`),
|
||||
wantErr: true,
|
||||
errMsg: "无效的 JSON 格式",
|
||||
},
|
||||
{
|
||||
name: "zero ID",
|
||||
payload: []byte(`{"id":0}`),
|
||||
wantErr: true,
|
||||
errMsg: "证书 ID 不能为空或零",
|
||||
},
|
||||
{
|
||||
name: "valid ID",
|
||||
payload: []byte(`{"id":123}`),
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
handler := &SSLSingleRenewHandler{}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := handler.ValidatePayload(tt.payload)
|
||||
if tt.wantErr {
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tt.errMsg)
|
||||
assert.Nil(t, got)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
var payload SSLSingleRenewPayload
|
||||
err = json.Unmarshal(got, &payload)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint(123), payload.ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSSLSingleRenewHandler_Execute(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
handler := &SSLSingleRenewHandler{}
|
||||
|
||||
t.Run("invalid payload", func(t *testing.T) {
|
||||
_, err := handler.Execute(ctx, []byte(`{`))
|
||||
require.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("certificate not found", func(t *testing.T) {
|
||||
payload, err := json.Marshal(SSLSingleRenewPayload{ID: 999})
|
||||
require.NoError(t, err)
|
||||
_, err = handler.Execute(ctx, payload)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "获取证书记录失败")
|
||||
})
|
||||
|
||||
t.Run("certificate provider is not ACME", func(t *testing.T) {
|
||||
cert := &model.TLSCertificate{
|
||||
Name: "custom-cert",
|
||||
Provider: "custom",
|
||||
PrimaryDomain: "example.com",
|
||||
}
|
||||
err := repository.CreateTLSCertificateRecord(ctx, cert)
|
||||
require.NoError(t, err)
|
||||
|
||||
payload, err := json.Marshal(SSLSingleRenewPayload{ID: cert.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = handler.Execute(ctx, payload)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "不是 ACME 托管证书")
|
||||
})
|
||||
|
||||
t.Run("successful renewal", func(t *testing.T) {
|
||||
cert := &model.TLSCertificate{
|
||||
Name: "acme-cert-success",
|
||||
Provider: tlsProviderACME,
|
||||
PrimaryDomain: "success.example.com",
|
||||
}
|
||||
err := repository.CreateTLSCertificateRecord(ctx, cert)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Mock obtainCertificate to succeed
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error {
|
||||
c.ApplyStatus = tlsApplyStatusReady
|
||||
return repository.SaveTLSCertificate(ctx, c)
|
||||
})
|
||||
defer restore()
|
||||
|
||||
payload, err := json.Marshal(SSLSingleRenewPayload{ID: cert.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
res, err := handler.Execute(ctx, payload)
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, res.Message, "续签成功")
|
||||
|
||||
updated, err := repository.GetTLSCertificateByID(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tlsApplyStatusReady, updated.ApplyStatus)
|
||||
})
|
||||
|
||||
t.Run("failed renewal in obtain", func(t *testing.T) {
|
||||
cert := &model.TLSCertificate{
|
||||
Name: "acme-cert-fail",
|
||||
Provider: tlsProviderACME,
|
||||
PrimaryDomain: "fail.example.com",
|
||||
}
|
||||
err := repository.CreateTLSCertificateRecord(ctx, cert)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Mock obtainCertificate to fail
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error {
|
||||
return errors.New("ACME server timeout")
|
||||
})
|
||||
defer restore()
|
||||
|
||||
payload, err := json.Marshal(SSLSingleRenewPayload{ID: cert.ID})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = handler.Execute(ctx, payload)
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, "ACME server timeout", err.Error())
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user