mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 08:06:37 +08:00
refactor(proxy): bind routes through zone domains
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user