mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 07:26:36 +08:00
[功能] 添加域名证书绑定支持,允许为每个域名单独选择证书并优化相关逻辑
This commit is contained in:
@@ -74,6 +74,7 @@ type snapshotRoute struct {
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id,omitempty"`
|
||||
CertIDs []uint `json:"cert_ids,omitempty"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids,omitempty"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server,omitempty"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"`
|
||||
@@ -472,6 +473,7 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: route.CertID,
|
||||
CertIDs: mustDecodeSnapshotCertIDs(route),
|
||||
DomainCertIDs: mustDecodeSnapshotDomainCertIDs(route, domains),
|
||||
RedirectHTTP: route.RedirectHTTP,
|
||||
LimitConnPerServer: route.LimitConnPerServer,
|
||||
LimitConnPerIP: route.LimitConnPerIP,
|
||||
@@ -497,6 +499,24 @@ func mustDecodeSnapshotCertIDs(route *model.ProxyRoute) []uint {
|
||||
return certIDs
|
||||
}
|
||||
|
||||
func mustDecodeSnapshotDomainCertIDs(
|
||||
route *model.ProxyRoute,
|
||||
domains []string,
|
||||
) []uint {
|
||||
if route == nil {
|
||||
return []uint{}
|
||||
}
|
||||
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
|
||||
if err != nil {
|
||||
return []uint{}
|
||||
}
|
||||
domainCertIDs, err := resolveProxyRouteDomainCertIDs(route, domains, certIDs)
|
||||
if err != nil {
|
||||
return []uint{}
|
||||
}
|
||||
return domainCertIDs
|
||||
}
|
||||
|
||||
func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
@@ -544,6 +564,14 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
|
||||
routes[index].CertID = primaryCertID
|
||||
routes[index].CertIDs = normalizedCertIDs
|
||||
}
|
||||
normalizedDomainCertIDs, err := normalizeSnapshotDomainCertificateIDs(
|
||||
routes[index].Domains,
|
||||
routes[index].CertIDs,
|
||||
routes[index].DomainCertIDs,
|
||||
)
|
||||
if err == nil {
|
||||
routes[index].DomainCertIDs = normalizedDomainCertIDs
|
||||
}
|
||||
normalizedUpstreams, err := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams)
|
||||
if err == nil {
|
||||
routes[index].OriginURL = normalizedUpstreams[0]
|
||||
@@ -583,7 +611,7 @@ func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRo
|
||||
}
|
||||
|
||||
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || !uintSliceEqual(left.CertIDs, right.CertIDs) {
|
||||
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || !uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) {
|
||||
return false
|
||||
}
|
||||
if len(left.Domains) != len(right.Domains) {
|
||||
@@ -814,48 +842,79 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("route %s cert_ids are invalid: %w", route.Domain, err)
|
||||
}
|
||||
if len(certIDs) > 0 {
|
||||
certificates, err := loadTLSCertificates(certIDs)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("route %s certificate lookup failed: %w", route.Domain, err)
|
||||
}
|
||||
if err := validateCertificateCoverageSet(certificates, domains); err != nil {
|
||||
return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
}
|
||||
for _, certificate := range certificates {
|
||||
supportFiles = append(supportFiles,
|
||||
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
|
||||
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)},
|
||||
)
|
||||
}
|
||||
if route.RedirectHTTP {
|
||||
builder.WriteString(renderHTTPRedirectServer(serverNames))
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg))
|
||||
}
|
||||
builder.WriteString(renderHTTPSServerWithCertificates(serverNames, route.OriginURL, route.OriginHost, certIDs, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg))
|
||||
continue
|
||||
domainCertIDs, err := resolveProxyRouteDomainCertIDs(route, domains, certIDs)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("route %s domain_cert_ids are invalid: %w", route.Domain, err)
|
||||
}
|
||||
if route.CertID == nil || *route.CertID == 0 {
|
||||
return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
certificate, err := model.GetTLSCertificateByID(*route.CertID)
|
||||
if len(certIDs) == 0 {
|
||||
return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
certificates, err := loadTLSCertificates(certIDs)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 关联证书不存在", route.Domain)
|
||||
return "", nil, fmt.Errorf("route %s certificate lookup failed: %w", route.Domain, err)
|
||||
}
|
||||
if err := validateCertificateCoverage(certificate, domains); err != nil {
|
||||
return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
certificateByID := make(map[uint]*model.TLSCertificate, len(certificates))
|
||||
for _, certificate := range certificates {
|
||||
if certificate == nil {
|
||||
continue
|
||||
}
|
||||
certificateByID[certificate.ID] = certificate
|
||||
supportFiles = append(supportFiles,
|
||||
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
|
||||
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)},
|
||||
)
|
||||
}
|
||||
supportFiles = append(supportFiles,
|
||||
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
|
||||
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)},
|
||||
)
|
||||
|
||||
httpOnlyDomains := make([]string, 0, len(domains))
|
||||
domainsByCertID := make(map[uint][]string, len(certIDs))
|
||||
for index, domain := range domains {
|
||||
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
|
||||
httpOnlyDomains = append(httpOnlyDomains, domain)
|
||||
continue
|
||||
}
|
||||
domainsByCertID[domainCertIDs[index]] = append(
|
||||
domainsByCertID[domainCertIDs[index]],
|
||||
domain,
|
||||
)
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
certificate := certificateByID[certID]
|
||||
if certificate == nil {
|
||||
return "", nil, fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
|
||||
}
|
||||
if err := validateCertificateCoverage(certificate, assignedDomains); err != nil {
|
||||
return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
}
|
||||
}
|
||||
|
||||
if route.RedirectHTTP {
|
||||
builder.WriteString(renderHTTPRedirectServer(serverNames))
|
||||
if len(httpOnlyDomains) > 0 {
|
||||
builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
|
||||
}
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg))
|
||||
}
|
||||
builder.WriteString(renderHTTPSServer(serverNames, route.OriginURL, route.OriginHost, certificate.ID, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg))
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), route.OriginURL, route.OriginHost, certID, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg))
|
||||
}
|
||||
}
|
||||
return builder.String(), dedupeSupportFiles(supportFiles), nil
|
||||
}
|
||||
@@ -988,6 +1047,37 @@ func normalizeSnapshotCertificateIDs(primaryCertID *uint, certIDs []uint) ([]uin
|
||||
return normalized, normalizedPrimary, nil
|
||||
}
|
||||
|
||||
func normalizeSnapshotDomainCertificateIDs(
|
||||
domains []string,
|
||||
certIDs []uint,
|
||||
domainCertIDs []uint,
|
||||
) ([]uint, error) {
|
||||
if len(domainCertIDs) > 0 {
|
||||
if len(domains) > 0 && len(domainCertIDs) != len(domains) {
|
||||
return nil, errors.New("snapshot domain_cert_ids length is invalid")
|
||||
}
|
||||
normalized := make([]uint, len(domainCertIDs))
|
||||
copy(normalized, domainCertIDs)
|
||||
return normalized, nil
|
||||
}
|
||||
if len(certIDs) == 0 {
|
||||
return []uint{}, nil
|
||||
}
|
||||
if len(certIDs) == 1 {
|
||||
normalized := make([]uint, len(domains))
|
||||
for index := range normalized {
|
||||
normalized[index] = certIDs[0]
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
if len(certIDs) == len(domains) {
|
||||
normalized := make([]uint, len(certIDs))
|
||||
copy(normalized, certIDs)
|
||||
return normalized, nil
|
||||
}
|
||||
return []uint{}, nil
|
||||
}
|
||||
|
||||
func uintPointerEqual(left *uint, right *uint) bool {
|
||||
if left == nil || right == nil {
|
||||
return left == nil && right == nil
|
||||
|
||||
@@ -403,7 +403,7 @@ func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
EnableHTTPS: true,
|
||||
CertIDs: []uint{appCertificate.ID, wwwCertificate.ID},
|
||||
DomainCertIDs: []uint{appCertificate.ID, wwwCertificate.ID},
|
||||
RedirectHTTP: true,
|
||||
CacheEnabled: true,
|
||||
CachePolicy: proxyRouteCachePolicyPathPrefix,
|
||||
@@ -419,6 +419,9 @@ func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t
|
||||
if len(route.CertIDs) != 2 || route.CertIDs[0] != appCertificate.ID || route.CertIDs[1] != wwwCertificate.ID {
|
||||
t.Fatalf("expected cert_ids to persist in order, got %#v", route.CertIDs)
|
||||
}
|
||||
if len(route.DomainCertIDs) != 2 || route.DomainCertIDs[0] != appCertificate.ID || route.DomainCertIDs[1] != wwwCertificate.ID {
|
||||
t.Fatalf("expected domain_cert_ids to persist per domain, got %#v", route.DomainCertIDs)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root")
|
||||
if err != nil {
|
||||
@@ -439,6 +442,59 @@ func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"cert_ids":[`) {
|
||||
t.Fatal("expected snapshot to include cert_ids")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"domain_cert_ids":[`) {
|
||||
t.Fatal("expected snapshot to include domain_cert_ids")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionSkipsHTTPSForDomainsWithoutCertificate(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
appCertPEM, appKeyPEM := generateCertificatePair(t, []string{"app.example.com"})
|
||||
appCertificate, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "app-only",
|
||||
CertPEM: appCertPEM,
|
||||
KeyPEM: appKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate app-only failed: %v", err)
|
||||
}
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "partial-https-site",
|
||||
Domains: []string{"app.example.com", "www.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
EnableHTTPS: true,
|
||||
DomainCertIDs: []uint{appCertificate.ID, 0},
|
||||
RedirectHTTP: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if len(route.CertIDs) != 1 || route.CertIDs[0] != appCertificate.ID {
|
||||
t.Fatalf("expected website cert_ids to keep used certificates only, got %#v", route.CertIDs)
|
||||
}
|
||||
if len(route.DomainCertIDs) != 2 || route.DomainCertIDs[0] != appCertificate.ID || route.DomainCertIDs[1] != 0 {
|
||||
t.Fatalf("expected domain_cert_ids to preserve unassigned domains, got %#v", route.DomainCertIDs)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root")
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if strings.Contains(result.Version.RenderedConfig, "listen 443 ssl;\n http2 on;\n server_name app.example.com www.example.com;") {
|
||||
t.Fatal("expected https server block to exclude domains without certificate")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "listen 443 ssl;\n http2 on;\n server_name app.example.com;") {
|
||||
t.Fatal("expected https server block to contain only the certified domain")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "listen 80;\n server_name app.example.com;\n\n return 301 https://$host$request_uri;") {
|
||||
t.Fatal("expected certified domain to keep http redirect")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "listen 80;\n server_name www.example.com;") {
|
||||
t.Fatal("expected non-certified domain to stay on plain http")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffConfigVersionTracksAddedDomainWithinWebsite(t *testing.T) {
|
||||
|
||||
@@ -44,6 +44,7 @@ type ProxyRouteInput struct {
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs []uint `json:"cert_ids"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip"`
|
||||
@@ -71,6 +72,7 @@ type ProxyRouteView struct {
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs []uint `json:"cert_ids"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip"`
|
||||
@@ -194,14 +196,33 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDs, err := normalizeProxyRouteCertificateIDs(input.EnableHTTPS, input.CertID, input.CertIDs)
|
||||
if !input.EnableHTTPS {
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
|
||||
domains,
|
||||
input.EnableHTTPS,
|
||||
input.DomainCertIDs,
|
||||
input.CertID,
|
||||
input.CertIDs,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteDomainCertificateCoverage(domains, domainCertIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDsJSON, err := json.Marshal(certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainsJSON, err := json.Marshal(domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -216,15 +237,9 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
if err := validateOriginHost(originHost); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !input.EnableHTTPS {
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
}
|
||||
input.DomainCertIDs = domainCertIDs
|
||||
input.CertIDs = certIDs
|
||||
if len(certIDs) > 0 {
|
||||
input.CertID = &certIDs[0]
|
||||
}
|
||||
input.CertID = primaryCertID
|
||||
if input.RedirectHTTP && !input.EnableHTTPS {
|
||||
return nil, errors.New("redirect_http requires enable_https")
|
||||
}
|
||||
@@ -243,6 +258,7 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.CertIDs = string(certIDsJSON)
|
||||
route.DomainCertIDs = string(domainCertIDsJSON)
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.LimitConnPerServer = limitConnPerServer
|
||||
route.LimitConnPerIP = limitConnPerIP
|
||||
@@ -291,6 +307,10 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDs, err := resolveProxyRouteDomainCertIDs(route, domains, certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var certID *uint
|
||||
if len(certIDs) > 0 {
|
||||
certID = &certIDs[0]
|
||||
@@ -312,6 +332,7 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) {
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: certID,
|
||||
CertIDs: certIDs,
|
||||
DomainCertIDs: domainCertIDs,
|
||||
RedirectHTTP: route.RedirectHTTP,
|
||||
LimitConnPerServer: route.LimitConnPerServer,
|
||||
LimitConnPerIP: route.LimitConnPerIP,
|
||||
@@ -476,6 +497,185 @@ func normalizeProxyRouteCertificateIDs(enableHTTPS bool, certID *uint, certIDs [
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteDomainCertificateIDs(
|
||||
domains []string,
|
||||
enableHTTPS bool,
|
||||
rawDomainCertIDs []uint,
|
||||
certID *uint,
|
||||
certIDs []uint,
|
||||
) ([]uint, []uint, *uint, error) {
|
||||
if !enableHTTPS {
|
||||
return []uint{}, []uint{}, nil, nil
|
||||
}
|
||||
|
||||
if len(rawDomainCertIDs) > 0 {
|
||||
if len(rawDomainCertIDs) != len(domains) {
|
||||
return nil, nil, nil, errors.New("domain_cert_ids must match domains length")
|
||||
}
|
||||
|
||||
normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs))
|
||||
uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs))
|
||||
seen := make(map[uint]struct{}, len(rawDomainCertIDs))
|
||||
hasAssignedCertificate := false
|
||||
for index, item := range rawDomainCertIDs {
|
||||
if item == 0 {
|
||||
continue
|
||||
}
|
||||
if _, err := model.GetTLSCertificateByID(item); err != nil {
|
||||
return nil, nil, nil, errors.New("selected certificate does not exist")
|
||||
}
|
||||
normalizedDomainCertIDs[index] = item
|
||||
hasAssignedCertificate = true
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
uniqueCertIDs = append(uniqueCertIDs, item)
|
||||
}
|
||||
if !hasAssignedCertificate {
|
||||
return nil, nil, nil, errors.New("must select a certificate when HTTPS is enabled")
|
||||
}
|
||||
|
||||
primaryCertID := &uniqueCertIDs[0]
|
||||
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
|
||||
}
|
||||
|
||||
normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(
|
||||
enableHTTPS,
|
||||
certID,
|
||||
certIDs,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
switch {
|
||||
case len(normalizedCertIDs) == 0:
|
||||
return nil, nil, nil, errors.New("must select a certificate when HTTPS is enabled")
|
||||
case len(normalizedCertIDs) == 1:
|
||||
domainCertIDs := make([]uint, len(domains))
|
||||
for index := range domainCertIDs {
|
||||
domainCertIDs[index] = normalizedCertIDs[0]
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
case len(normalizedCertIDs) == len(domains):
|
||||
domainCertIDs := make([]uint, len(normalizedCertIDs))
|
||||
copy(domainCertIDs, normalizedCertIDs)
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
default:
|
||||
domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(
|
||||
domains,
|
||||
normalizedCertIDs,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
}
|
||||
}
|
||||
|
||||
func validateProxyRouteDomainCertificateCoverage(
|
||||
domains []string,
|
||||
domainCertIDs []uint,
|
||||
) error {
|
||||
if len(domainCertIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
domainsByCertID := make(map[uint][]string)
|
||||
for index, certID := range domainCertIDs {
|
||||
if certID == 0 {
|
||||
continue
|
||||
}
|
||||
domainsByCertID[certID] = append(domainsByCertID[certID], domains[index])
|
||||
}
|
||||
|
||||
for certID, assignedDomains := range domainsByCertID {
|
||||
certificate, err := model.GetTLSCertificateByID(certID)
|
||||
if err != nil {
|
||||
return errors.New("selected certificate does not exist")
|
||||
}
|
||||
if err := validateCertificateCoverage(certificate, assignedDomains); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deriveDomainCertIDsFromCertificateSet(
|
||||
domains []string,
|
||||
certIDs []uint,
|
||||
) ([]uint, error) {
|
||||
certificates, err := loadTLSCertificates(certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
result := make([]uint, len(domains))
|
||||
for domainIndex, domain := range domains {
|
||||
if domainIndex < len(certificates) &&
|
||||
certificates[domainIndex] != nil &&
|
||||
validateCertificateCoverage(certificates[domainIndex], []string{domain}) == nil {
|
||||
result[domainIndex] = certificates[domainIndex].ID
|
||||
continue
|
||||
}
|
||||
|
||||
assigned := uint(0)
|
||||
for _, certificate := range certificates {
|
||||
if certificate != nil &&
|
||||
validateCertificateCoverage(certificate, []string{domain}) == nil {
|
||||
assigned = certificate.ID
|
||||
break
|
||||
}
|
||||
}
|
||||
if assigned == 0 {
|
||||
return nil, fmt.Errorf("certificate does not cover domain %s", domain)
|
||||
}
|
||||
result[domainIndex] = assigned
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []uint{}, nil
|
||||
}
|
||||
|
||||
var domainCertIDs []uint
|
||||
if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil {
|
||||
return nil, errors.New("domain_cert_ids payload is invalid")
|
||||
}
|
||||
if len(domainCertIDs) == 0 {
|
||||
return []uint{}, nil
|
||||
}
|
||||
if domainCount > 0 && len(domainCertIDs) != domainCount {
|
||||
return nil, errors.New("domain_cert_ids length does not match domains")
|
||||
}
|
||||
|
||||
normalized := make([]uint, len(domainCertIDs))
|
||||
copy(normalized, domainCertIDs)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func resolveProxyRouteDomainCertIDs(
|
||||
route *model.ProxyRoute,
|
||||
domains []string,
|
||||
certIDs []uint,
|
||||
) ([]uint, error) {
|
||||
domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, len(domains))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(domainCertIDs) > 0 || len(certIDs) == 0 {
|
||||
return domainCertIDs, nil
|
||||
}
|
||||
return deriveDomainCertIDsFromCertificateSet(domains, certIDs)
|
||||
}
|
||||
|
||||
func normalizeProxyRouteLimitRate(raw string) (string, error) {
|
||||
normalized := strings.ToLower(strings.TrimSpace(raw))
|
||||
if normalized == "" || normalized == "0" {
|
||||
|
||||
@@ -125,6 +125,15 @@ func DeleteTLSCertificate(id uint) error {
|
||||
return errors.New("certificate is still referenced by proxy routes")
|
||||
}
|
||||
}
|
||||
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("certificate is still referenced by proxy routes")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
certificate, err := model.GetTLSCertificateByID(id)
|
||||
|
||||
Reference in New Issue
Block a user