refactor(proxy): bind routes through zone domains

This commit is contained in:
ryan
2026-07-12 14:52:19 +08:00
parent e51f1e583d
commit d0536fcdd5
13 changed files with 332 additions and 638 deletions
-17
View File
@@ -5,10 +5,8 @@ package tls
import (
"crypto/x509"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"io"
"mime/multipart"
"strings"
@@ -45,18 +43,3 @@ func isUniqueConstraintError(err error) bool {
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) {
text := strings.TrimSpace(raw)
if text == "" {
return nil, nil
}
var domainCertIDs []uint
if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil {
return nil, err
}
if domainCount > 0 && len(domainCertIDs) != domainCount {
return nil, fmt.Errorf("domain_cert_ids length mismatch")
}
return domainCertIDs, nil
}
+3 -26
View File
@@ -430,35 +430,12 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
}
func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
routes, err := model.ListTLSProxyRouteRefs(ctx)
count, err := model.CountZoneDomainsByCertificateID(ctx, id)
if err != nil {
return err
}
for _, route := range routes {
if route.CertID != nil && *route.CertID == id {
return errors.New(errCertificateDeleteReferenced)
}
if strings.TrimSpace(route.CertIDs) == "" {
continue
}
var certIDs []uint
if err := json.Unmarshal([]byte(route.CertIDs), &certIDs); err != nil {
return fmt.Errorf("proxy route %d cert_ids payload is invalid: %w", route.ID, err)
}
for _, certID := range certIDs {
if certID == id {
return errors.New(errCertificateDeleteReferenced)
}
}
domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, 0)
if err != nil {
return fmt.Errorf("proxy route %d domain_cert_ids payload is invalid: %w", route.ID, err)
}
for _, certID := range domainCertIDs {
if certID == id {
return errors.New(errCertificateDeleteReferenced)
}
}
if count > 0 {
return errors.New(errCertificateDeleteReferenced)
}
return nil
}
@@ -40,6 +40,8 @@ func setupTLSTestDB(t *testing.T) func() {
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.TLSCertificate{},
&model.Zone{},
&model.ZoneDomain{},
&model.ManagedDomain{},
&model.DNSAccount{},
&model.AcmeAccount{},
@@ -71,6 +73,22 @@ func setupTLSTestDB(t *testing.T) func() {
}
}
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)