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
+9 -1
View File
@@ -74,9 +74,17 @@ func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
}
items := make([]RouteSummary, 0, len(routes))
for _, route := range routes {
domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
domain := ""
if len(domains) > 0 {
domain = domains[0].Domain
}
items = append(items, RouteSummary{
ID: route.ID,
Domain: route.Domain,
Domain: domain,
OriginURL: route.OriginURL,
Enabled: route.Enabled,
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
@@ -11,15 +11,13 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
type proxyRouteJSONFields struct {
cacheRulesJSON string
upstreamsJSON string
customHeadersJSON string
certIDsJSON string
domainCertIDsJSON string
domainsJSON string
}
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) {
@@ -46,12 +44,9 @@ func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input
}
func marshalProxyRouteJSONFields(
domains []string,
upstreams []string,
cacheRules []string,
customHeaders []CustomHeaderInput,
certIDs []uint,
domainCertIDs []uint,
) (*proxyRouteJSONFields, error) {
cacheRulesJSON, err := json.Marshal(cacheRules)
if err != nil {
@@ -65,38 +60,13 @@ func marshalProxyRouteJSONFields(
if 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
}
return &proxyRouteJSONFields{
cacheRulesJSON: string(cacheRulesJSON),
upstreamsJSON: string(upstreamsJSON),
customHeadersJSON: string(customHeadersJSON),
certIDsJSON: string(certIDsJSON),
domainCertIDsJSON: string(domainCertIDsJSON),
domainsJSON: string(domainsJSON),
}, nil
}
func normalizeProxyRouteHTTPSInput(input *Input) {
if input.EnableHTTPS {
return
}
input.RedirectHTTP = false
input.CertID = nil
input.CertIDs = nil
input.DomainCertIDs = nil
}
func normalizeProxyRouteBasicAuth(input *Input) error {
if !input.BasicAuthEnabled {
input.BasicAuthUsername = ""
@@ -114,7 +84,7 @@ func normalizeProxyRouteBasicAuth(input *Input) error {
func populateProxyRouteFields(
route *model.ProxyRoute,
input Input,
siteName, domain string,
siteName string,
jsonFields *proxyRouteJSONFields,
originID *uint,
upstreams []string,
@@ -123,17 +93,12 @@ func populateProxyRouteFields(
limitRate, upstreamType string,
) {
route.SiteName = siteName
route.Domain = domain
route.Domains = jsonFields.domainsJSON
route.OriginID = originID
route.OriginURL = upstreams[0]
route.OriginHost = originHost
route.Upstreams = jsonFields.upstreamsJSON
route.Enabled = input.Enabled
route.EnableHTTPS = input.EnableHTTPS
route.CertID = input.CertID
route.CertIDs = jsonFields.certIDsJSON
route.DomainCertIDs = jsonFields.domainCertIDsJSON
route.RedirectHTTP = input.RedirectHTTP
route.LimitConnPerServer = limitConnPerServer
route.LimitConnPerIP = limitConnPerIP
@@ -176,3 +141,40 @@ func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, u
}
return nil
}
func updateProxyRouteRecord(tx *gorm.DB, route *model.ProxyRoute) error {
return tx.Model(&model.ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
"site_name": route.SiteName, "origin_id": route.OriginID, "origin_url": route.OriginURL,
"domain": route.Domain, "domains": route.Domains, "cert_id": route.CertID,
"cert_ids": route.CertIDs, "domain_cert_ids": route.DomainCertIDs,
"origin_host": route.OriginHost, "upstreams": route.Upstreams, "enabled": route.Enabled,
"enable_https": route.EnableHTTPS, "redirect_http": route.RedirectHTTP,
"limit_conn_per_server": route.LimitConnPerServer, "limit_conn_per_ip": route.LimitConnPerIP,
"limit_rate": route.LimitRate, "cache_enabled": route.CacheEnabled, "cache_policy": route.CachePolicy,
"cache_rules": route.CacheRules, "custom_headers": route.CustomHeaders,
"basic_auth_enabled": route.BasicAuthEnabled, "basic_auth_username": route.BasicAuthUsername,
"basic_auth_password": route.BasicAuthPassword, "remark": route.Remark,
"upstream_type": route.UpstreamType, "tunnel_node_id": route.TunnelNodeID,
"tunnel_target_addr": route.TunnelTargetAddr, "tunnel_target_protocol": route.TunnelTargetProtocol,
"pages_project_id": route.PagesProjectID,
}).Error
}
func replaceZoneDomainRouteBindings(tx *gorm.DB, routeID uint, domainIDs []uint) error {
var requested []model.ZoneDomain
if err := tx.Where("id IN ?", domainIDs).Find(&requested).Error; err != nil {
return err
}
if len(requested) != len(domainIDs) {
return errors.New(errProxyRouteZoneDomainNotFound)
}
for _, domain := range requested {
if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID {
return errors.New(errProxyRouteZoneDomainBound)
}
}
if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ? AND id NOT IN ?", routeID, domainIDs).Update("proxy_route_id", nil).Error; err != nil {
return err
}
return tx.Model(&model.ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error
}
@@ -1,71 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"errors"
)
func normalizeExplicitDomainCertIDs(ctx context.Context, domains []string, rawDomainCertIDs []uint) ([]uint, []uint, *uint, error) {
if len(rawDomainCertIDs) != len(domains) {
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
}
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 := lookupTLSCertificateByID(ctx, item); err != nil {
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
}
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(errProxyRouteCertRequired)
}
primaryCertID := &uniqueCertIDs[0]
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
}
func normalizeDerivedDomainCertIDs(
ctx context.Context,
domains []string,
normalizedCertIDs []uint,
) ([]uint, []uint, *uint, error) {
switch {
case len(normalizedCertIDs) == 0:
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
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(ctx, domains, normalizedCertIDs)
if err != nil {
return nil, nil, nil, err
}
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
}
}
+47 -46
View File
@@ -4,50 +4,51 @@
package proxy_route
const (
errProxyRouteNotFound = "proxy route not found"
errProxyRouteIdentityExists = "proxy route identity already exists"
errProxyRouteSiteNameExists = "site_name already exists"
errProxyRouteDomainExists = "domain %s already exists"
errProxyRouteSiteNameEmpty = "site_name cannot be empty"
errProxyRouteDomainRequired = "at least one domain is required"
errProxyRouteDomainInvalid = "domain format is invalid"
errProxyRouteDomainMismatch = "domain must match domains[0]"
errProxyRouteOriginEmpty = "origin_url cannot be empty"
errProxyRouteOriginInvalid = "origin URL format is invalid"
errProxyRouteOriginScheme = "origin URL must start with http:// or https://"
errProxyRouteOriginHostInvalid = "origin_host format is invalid"
errProxyRouteUpstreamRequired = "at least one upstream is required"
errProxyRouteUpstreamScheme = "all upstreams must use the same scheme"
errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths"
errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings"
errProxyRouteOriginNotFound = "selected origin does not exist"
errProxyRouteCertNotFound = "selected certificate does not exist"
errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled"
errProxyRouteCertDomainLength = "domain_cert_ids must match domains length"
errProxyRouteRedirectHTTP = "redirect_http requires enable_https"
errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled"
errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format"
errProxyRouteCachePolicy = "cache policy is not supported"
errProxyRouteCacheSuffix = "cache suffix format is invalid"
errProxyRouteCachePath = "cache path rule format is invalid"
errProxyRouteCacheSuffixReq = "at least one suffix is required"
errProxyRouteCachePrefixReq = "at least one path prefix is required"
errProxyRouteCacheExactReq = "at least one exact path is required"
errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty"
errProxyRouteHeaderKeyInvalid = "custom header key format is invalid"
errProxyRouteHeaderNewline = "custom headers cannot contain newlines"
errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream"
errProxyRouteTunnelNodeMissing = "tunnel client node does not exist"
errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node"
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
errProxyRoutePagesNotFound = "pages 项目不存在"
errProxyRoutePagesDisabled = "pages 项目未启用"
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
errProxyRouteOriginPort = "端口格式不合法"
errProxyRouteOriginPortEmpty = "端口不能为空"
errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头"
errProxyRouteOriginURIProto = "源站路径不能包含协议"
errProxyRouteNotFound = "proxy route not found"
errProxyRouteIdentityExists = "proxy route identity already exists"
errProxyRouteSiteNameExists = "site_name already exists"
errProxyRouteDomainExists = "domain %s already exists"
errProxyRouteSiteNameEmpty = "site_name cannot be empty"
errProxyRouteZoneDomainsRequired = "at least one zone domain is required"
errProxyRouteZoneDomainNotFound = "selected zone domain does not exist"
errProxyRouteZoneDomainDuplicate = "zone_domain_ids must not contain duplicates"
errProxyRouteZoneDomainBound = "selected zone domain is already bound to another proxy route"
errProxyRouteOriginEmpty = "origin_url cannot be empty"
errProxyRouteOriginInvalid = "origin URL format is invalid"
errProxyRouteOriginScheme = "origin URL must start with http:// or https://"
errProxyRouteOriginHostInvalid = "origin_host format is invalid"
errProxyRouteUpstreamRequired = "at least one upstream is required"
errProxyRouteUpstreamScheme = "all upstreams must use the same scheme"
errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths"
errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings"
errProxyRouteOriginNotFound = "selected origin does not exist"
errProxyRouteCertNotFound = "selected certificate does not exist"
errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled"
errProxyRouteCertDomainLength = "domain_cert_ids must match domains length"
errProxyRouteRedirectHTTP = "redirect_http requires enable_https"
errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled"
errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format"
errProxyRouteCachePolicy = "cache policy is not supported"
errProxyRouteCacheSuffix = "cache suffix format is invalid"
errProxyRouteCachePath = "cache path rule format is invalid"
errProxyRouteCacheSuffixReq = "at least one suffix is required"
errProxyRouteCachePrefixReq = "at least one path prefix is required"
errProxyRouteCacheExactReq = "at least one exact path is required"
errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty"
errProxyRouteHeaderKeyInvalid = "custom header key format is invalid"
errProxyRouteHeaderNewline = "custom headers cannot contain newlines"
errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream"
errProxyRouteTunnelNodeMissing = "tunnel client node does not exist"
errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node"
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
errProxyRoutePagesNotFound = "pages 项目不存在"
errProxyRoutePagesDisabled = "pages 项目未启用"
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
errProxyRouteOriginPort = "端口格式不合法"
errProxyRouteOriginPortEmpty = "端口不能为空"
errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头"
errProxyRouteOriginURIProto = "源站路径不能包含协议"
)
+33 -242
View File
@@ -17,7 +17,6 @@ import (
"strings"
"unicode"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
@@ -254,63 +253,23 @@ func validateCertificateCoverage(certificate *model.TLSCertificate, domains []st
return nil
}
func loadTLSCertificates(ctx context.Context, certIDs []uint) ([]*model.TLSCertificate, error) {
certificates := make([]*model.TLSCertificate, 0, len(certIDs))
for _, certID := range certIDs {
certificate, err := lookupTLSCertificateByID(ctx, certID)
if err != nil {
return nil, err
}
certificates = append(certificates, certificate)
func loadProxyRouteZoneDomains(ctx context.Context, ids []uint) ([]model.ZoneDomain, error) {
if len(ids) == 0 {
return nil, errors.New(errProxyRouteZoneDomainsRequired)
}
return certificates, nil
}
func normalizeProxyRouteDomainValue(raw string) string {
return strings.ToLower(strings.TrimSpace(raw))
}
func mapRouteIdentityDomainError(err error) error {
if err == nil {
return nil
}
switch err.Error() {
case "domain is required":
return errors.New(errProxyRouteDomainRequired)
default:
if strings.Contains(err.Error(), " is invalid") {
return errors.New(errProxyRouteDomainInvalid)
seen := make(map[uint]struct{}, len(ids))
for _, id := range ids {
if id == 0 {
return nil, errors.New(errProxyRouteZoneDomainNotFound)
}
return err
}
}
func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string, rawDomains []string) ([]string, error) {
if len(rawDomains) > 0 {
domains, err := routeidentity.NormalizeDomains(rawDomains)
if err != nil {
return nil, mapRouteIdentityDomainError(err)
if _, ok := seen[id]; ok {
return nil, errors.New(errProxyRouteZoneDomainDuplicate)
}
domain := normalizeProxyRouteDomainValue(rawDomain)
if domain != "" && domain != domains[0] {
return nil, errors.New(errProxyRouteDomainMismatch)
}
return domains, nil
seen[id] = struct{}{}
}
if route != nil {
existingDomains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
if err == nil && len(existingDomains) > 0 {
domain := normalizeProxyRouteDomainValue(rawDomain)
if domain == "" || domain == existingDomains[0] {
return existingDomains, nil
}
}
}
domains, err := routeidentity.NormalizeDomains([]string{rawDomain})
domains, err := model.ListZoneDomainsByIDs(ctx, ids)
if err != nil {
return nil, mapRouteIdentityDomainError(err)
return nil, errors.New(errProxyRouteZoneDomainNotFound)
}
return domains, nil
}
@@ -322,7 +281,7 @@ func validateProxyRouteSiteName(siteName string) error {
return nil
}
func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string, domains []string) error {
func validateProxyRouteSiteNameUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string) error {
routes, err := model.ListProxyRoutes(ctx)
if err != nil {
return err
@@ -337,27 +296,33 @@ func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.Prox
if item == nil || item.ID == currentID {
continue
}
existingSiteName, existingDomains, err := routeidentity.ResolveFromRoute(item)
if err != nil {
return fmt.Errorf("existing route %d domains are invalid: %w", item.ID, err)
}
if existingSiteName == siteName {
if item.SiteName == siteName {
return errors.New(errProxyRouteSiteNameExists)
}
existingSet := make(map[string]struct{}, len(existingDomains))
for _, existingDomain := range existingDomains {
existingSet[existingDomain] = struct{}{}
}
for _, domain := range domains {
if _, ok := existingSet[domain]; ok {
return fmt.Errorf(errProxyRouteDomainExists, domain)
}
}
}
return nil
}
func validateProxyRouteZoneDomainCertificates(ctx context.Context, domains []model.ZoneDomain, enableHTTPS bool) error {
if !enableHTTPS {
return nil
}
for _, domain := range domains {
if domain.CertID == nil || *domain.CertID == 0 {
return errors.New(errProxyRouteCertRequired)
}
certificate, err := lookupTLSCertificateByID(ctx, *domain.CertID)
if err != nil {
return errors.New(errProxyRouteCertNotFound)
}
if err := validateCertificateCoverage(certificate, []string{domain.Domain}); err != nil {
return err
}
}
return nil
}
func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) {
if value < 0 {
return 0, fmt.Errorf("%s must be greater than or equal to 0", field)
@@ -365,150 +330,6 @@ func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) {
return value, nil
}
func normalizeProxyRouteCertificateIDs(ctx context.Context, enableHTTPS bool, certID *uint, certIDs []uint) ([]uint, error) {
if !enableHTTPS {
return []uint{}, nil
}
candidates := make([]uint, 0, len(certIDs)+1)
if certID != nil && *certID != 0 {
candidates = append(candidates, *certID)
}
candidates = append(candidates, certIDs...)
normalized := make([]uint, 0, len(candidates))
seen := make(map[uint]struct{}, len(candidates))
for _, item := range candidates {
if item == 0 {
continue
}
if _, ok := seen[item]; ok {
continue
}
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
return nil, errors.New(errProxyRouteCertNotFound)
}
seen[item] = struct{}{}
normalized = append(normalized, item)
}
if len(normalized) == 0 {
return nil, errors.New(errProxyRouteCertRequired)
}
return normalized, nil
}
func normalizeProxyRouteDomainCertificateIDs(
ctx context.Context,
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 {
return normalizeExplicitDomainCertIDs(ctx, domains, rawDomainCertIDs)
}
normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(ctx, enableHTTPS, certID, certIDs)
if err != nil {
return nil, nil, nil, err
}
return normalizeDerivedDomainCertIDs(ctx, domains, normalizedCertIDs)
}
func validateProxyRouteDomainCertificateCoverage(ctx context.Context, 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 := lookupTLSCertificateByID(ctx, certID)
if err != nil {
return errors.New(errProxyRouteCertNotFound)
}
if err := validateCertificateCoverage(certificate, assignedDomains); err != nil {
return err
}
}
return nil
}
func deriveDomainCertIDsFromCertificateSet(ctx context.Context, domains []string, certIDs []uint) ([]uint, error) {
certificates, err := loadTLSCertificates(ctx, 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(ctx context.Context, 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(ctx, domains, certIDs)
}
func normalizeProxyRouteLimitRate(raw string) (string, error) {
normalized := strings.ToLower(strings.TrimSpace(raw))
if normalized == "" || normalized == "0" {
@@ -727,36 +548,6 @@ func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, erro
return normalizeUpstreams(fallbackOriginURL, upstreams)
}
func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) {
text := strings.TrimSpace(raw)
if text == "" {
if fallbackCertID == nil || *fallbackCertID == 0 {
return []uint{}, nil
}
return []uint{*fallbackCertID}, nil
}
var certIDs []uint
if err := json.Unmarshal([]byte(text), &certIDs); err != nil {
return nil, errors.New("cert_ids payload is invalid")
}
normalized := make([]uint, 0, len(certIDs))
seen := make(map[uint]struct{}, len(certIDs))
for _, certID := range certIDs {
if certID == 0 {
continue
}
if _, ok := seen[certID]; ok {
continue
}
seen[certID] = struct{}{}
normalized = append(normalized, certID)
}
if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 {
return []uint{*fallbackCertID}, nil
}
return normalized, nil
}
func validateOriginURL(raw string) error {
if raw == "" {
return errors.New(errProxyRouteOriginEmpty)
+97 -80
View File
@@ -5,12 +5,14 @@ package proxy_route
import (
"context"
"encoding/json"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// CustomHeaderInput 自定义响应头。
@@ -22,8 +24,7 @@ type CustomHeaderInput struct {
// Input 代理规则创建/更新请求。
type Input struct {
SiteName string `json:"site_name"`
Domain string `json:"domain"`
Domains []string `json:"domains"`
ZoneDomainIDs []uint `json:"zone_domain_ids"`
OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"`
OriginScheme string `json:"origin_scheme"`
@@ -34,9 +35,6 @@ type Input struct {
Upstreams []string `json:"upstreams"`
Enabled bool `json:"enabled"`
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"`
@@ -61,10 +59,8 @@ type Input struct {
type View struct {
ID uint `json:"id"`
SiteName string `json:"site_name"`
Domain string `json:"domain"`
Domains []string `json:"domains"`
PrimaryDomain string `json:"primary_domain"`
DomainCount int `json:"domain_count"`
ZoneDomainIDs []uint `json:"zone_domain_ids"`
ZoneDomains []ZoneDomainView `json:"zone_domains"`
OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host"`
@@ -72,9 +68,6 @@ type View struct {
UpstreamList []string `json:"upstream_list"`
Enabled bool `json:"enabled"`
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"`
@@ -99,6 +92,14 @@ type View struct {
UpdatedAt time.Time `json:"updated_at"`
}
// ZoneDomainView is the route-safe representation of a bound Zone domain.
type ZoneDomainView struct {
ID uint `json:"id"`
ZoneID uint `json:"zone_id"`
Domain string `json:"domain"`
CertID *uint `json:"cert_id"`
}
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*View, error) {
routes, err := model.ListProxyRoutes(ctx)
@@ -119,11 +120,16 @@ func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
// CreateProxyRoute 创建代理规则。
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
route, err := buildProxyRoute(ctx, nil, input)
route, _, err := buildProxyRoute(ctx, nil, input)
if err != nil {
return nil, err
}
if err = model.CreateProxyRouteRecord(ctx, route); err != nil {
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Create(route).Error; err != nil {
return err
}
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs)
}); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errProxyRouteIdentityExists)
}
@@ -138,11 +144,16 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
if err != nil {
return nil, err
}
route, err = buildProxyRoute(ctx, route, input)
route, _, err = buildProxyRoute(ctx, route, input)
if err != nil {
return nil, err
}
if err = model.UpdateProxyRouteRecord(ctx, route); err != nil {
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := updateProxyRouteRecord(tx, route); err != nil {
return err
}
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs)
}); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errProxyRouteIdentityExists)
}
@@ -156,84 +167,72 @@ func DeleteProxyRoute(ctx context.Context, id uint) error {
if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
return err
}
return model.DeleteProxyRouteRecord(ctx, id)
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", id).Update("proxy_route_id", nil).Error; err != nil {
return err
}
return tx.Delete(&model.ProxyRoute{}, id).Error
})
}
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, error) {
domains, err := normalizeProxyRouteDomainsInput(route, input.Domain, input.Domains)
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) {
domains, err := loadProxyRouteZoneDomains(ctx, input.ZoneDomainIDs)
if err != nil {
return nil, err
return nil, nil, err
}
domain := domains[0]
siteName := routeidentity.ResolveSiteName(route, input.SiteName, domain)
siteName := strings.TrimSpace(input.SiteName)
upstreamType := normalizeUpstreamType(input.UpstreamType)
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
if err != nil {
return nil, err
return nil, nil, err
}
originHost := strings.TrimSpace(input.OriginHost)
remark := strings.TrimSpace(input.Remark)
cachePolicy := strings.TrimSpace(input.CachePolicy)
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
if err != nil {
return nil, err
return nil, nil, err
}
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
if err != nil {
return nil, err
return nil, nil, err
}
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
if err != nil {
return nil, err
return nil, nil, err
}
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
if err != nil {
return nil, err
return nil, nil, err
}
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
if err != nil {
return nil, err
return nil, nil, err
}
normalizeProxyRouteHTTPSInput(&input)
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
ctx,
domains,
input.EnableHTTPS,
input.DomainCertIDs,
input.CertID,
input.CertIDs,
)
if err := validateProxyRouteZoneDomainCertificates(ctx, domains, input.EnableHTTPS); err != nil {
return nil, nil, err
}
jsonFields, err := marshalProxyRouteJSONFields(upstreams, cacheRules, customHeaders)
if err != nil {
return nil, err
}
if err := validateProxyRouteDomainCertificateCoverage(ctx, domains, domainCertIDs); err != nil {
return nil, err
}
jsonFields, err := marshalProxyRouteJSONFields(domains, upstreams, cacheRules, customHeaders, certIDs, domainCertIDs)
if err != nil {
return nil, err
return nil, nil, err
}
if err := validateProxyRouteSiteName(siteName); err != nil {
return nil, err
return nil, nil, err
}
if err := validateProxyRouteIdentityUniqueness(ctx, route, siteName, domains); err != nil {
return nil, err
if err := validateProxyRouteSiteNameUniqueness(ctx, route, siteName); err != nil {
return nil, nil, err
}
if err := validateOriginHost(originHost); err != nil {
return nil, err
return nil, nil, err
}
input.DomainCertIDs = domainCertIDs
input.CertIDs = certIDs
input.CertID = primaryCertID
if input.RedirectHTTP && !input.EnableHTTPS {
return nil, errors.New(errProxyRouteRedirectHTTP)
return nil, nil, errors.New(errProxyRouteRedirectHTTP)
}
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
return nil, err
return nil, nil, err
}
if route == nil {
@@ -243,7 +242,6 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
route,
input,
siteName,
domain,
jsonFields,
originID,
upstreams,
@@ -255,10 +253,41 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
limitRate,
upstreamType,
)
// Preserve legacy columns until the second-phase schema cleanup. They are
// derived solely from ZoneDomain bindings and are not exposed by this API.
populateLegacyZoneDomainFields(route, domains)
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
return nil, err
return nil, nil, err
}
return route, nil
return route, domains, nil
}
func mustMarshalProxyRouteLegacy(value any) string {
encoded, err := json.Marshal(value)
if err != nil {
panic(err)
}
return string(encoded)
}
func populateLegacyZoneDomainFields(route *model.ProxyRoute, domains []model.ZoneDomain) {
legacyDomains := make([]string, 0, len(domains))
legacyCertIDs := make([]uint, 0, len(domains))
for _, domain := range domains {
legacyDomains = append(legacyDomains, domain.Domain)
if domain.CertID != nil {
legacyCertIDs = append(legacyCertIDs, *domain.CertID)
}
}
route.Domain = legacyDomains[0]
route.Domains = mustMarshalProxyRouteLegacy(legacyDomains)
route.CertIDs = mustMarshalProxyRouteLegacy(legacyCertIDs)
if len(legacyCertIDs) > 0 {
route.CertID = &legacyCertIDs[0]
} else {
route.CertID = nil
}
route.DomainCertIDs = route.CertIDs
}
func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) {
@@ -277,7 +306,7 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
if route == nil {
return nil, errors.New("proxy route is nil")
}
domains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
@@ -293,26 +322,17 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
if err != nil {
return nil, err
}
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
if err != nil {
return nil, err
zoneDomainIDs := make([]uint, 0, len(domains))
zoneDomains := make([]ZoneDomainView, 0, len(domains))
for _, domain := range domains {
zoneDomainIDs = append(zoneDomainIDs, domain.ID)
zoneDomains = append(zoneDomains, ZoneDomainView{ID: domain.ID, ZoneID: domain.ZoneID, Domain: domain.Domain, CertID: domain.CertID})
}
domainCertIDs, err := resolveProxyRouteDomainCertIDs(ctx, route, domains, certIDs)
if err != nil {
return nil, err
}
var certID *uint
if len(certIDs) > 0 {
certID = &certIDs[0]
}
primaryDomain := domains[0]
return &View{
ID: route.ID,
SiteName: routeidentity.ResolveSiteName(route, route.SiteName, primaryDomain),
Domain: primaryDomain,
Domains: domains,
PrimaryDomain: primaryDomain,
DomainCount: len(domains),
SiteName: route.SiteName,
ZoneDomainIDs: zoneDomainIDs,
ZoneDomains: zoneDomains,
OriginID: route.OriginID,
OriginURL: route.OriginURL,
OriginHost: route.OriginHost,
@@ -320,9 +340,6 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
UpstreamList: upstreams,
Enabled: route.Enabled,
EnableHTTPS: route.EnableHTTPS,
CertID: certID,
CertIDs: certIDs,
DomainCertIDs: domainCertIDs,
RedirectHTTP: route.RedirectHTTP,
LimitConnPerServer: route.LimitConnPerServer,
LimitConnPerIP: route.LimitConnPerIP,
@@ -17,130 +17,66 @@ import (
func setupProxyRouteTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{}))
require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{}, &model.Zone{}, &model.ZoneDomain{}, &model.TLSCertificate{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
return func() { db.SetDB(nil) }
}
func createZoneDomain(t *testing.T, ctx context.Context, domain string, certID *uint) *model.ZoneDomain {
t.Helper()
zone := &model.Zone{Domain: "example.com"}
var existing model.Zone
if err := db.DB(ctx).Where("domain = ?", zone.Domain).First(&existing).Error; err == nil {
zone = &existing
} else {
require.NoError(t, db.DB(ctx).Create(zone).Error)
}
item := &model.ZoneDomain{ZoneID: zone.ID, Domain: domain, CertID: certID}
require.NoError(t, db.DB(ctx).Create(item).Error)
return item
}
func TestCreateProxyRoute(t *testing.T) {
func TestCreateProxyRouteBindsZoneDomains(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domainA := createZoneDomain(t, ctx, "api.example.com", nil)
domainB := createZoneDomain(t, ctx, "www.example.com", nil)
view, err := CreateProxyRoute(ctx, Input{
SiteName: "example-site",
Domain: "example.com",
OriginURL: "http://origin.example.com:8080",
Enabled: true,
})
view, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domainA.ID, domainB.ID}, OriginURL: "http://origin.example.com:8080", Enabled: true})
require.NoError(t, err)
assert.NotZero(t, view.ID)
assert.Equal(t, "example-site", view.SiteName)
assert.Equal(t, "example.com", view.Domain)
assert.Equal(t, []string{"example.com"}, view.Domains)
assert.Equal(t, "http://origin.example.com:8080", view.OriginURL)
assert.Equal(t, []string{"http://origin.example.com:8080"}, view.UpstreamList)
assert.True(t, view.Enabled)
assert.Equal(t, []uint{domainA.ID, domainB.ID}, view.ZoneDomainIDs)
require.Len(t, view.ZoneDomains, 2)
assert.Equal(t, "api.example.com", view.ZoneDomains[0].Domain)
}
_, err = CreateProxyRoute(ctx, Input{
SiteName: "duplicate-site",
Domain: "example.com",
OriginURL: "http://origin.example.com:8080",
})
func TestCreateProxyRouteRejectsInvalidZoneDomainBindings(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domain := createZoneDomain(t, ctx, "api.example.com", nil)
base := Input{SiteName: "api", OriginURL: "http://origin.example.com:8080"}
_, err := CreateProxyRoute(ctx, base)
require.EqualError(t, err, errProxyRouteZoneDomainsRequired)
base.ZoneDomainIDs = []uint{domain.ID, domain.ID}
_, err = CreateProxyRoute(ctx, base)
require.EqualError(t, err, errProxyRouteZoneDomainDuplicate)
first, err := CreateProxyRoute(ctx, Input{SiteName: "first", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080"})
require.NoError(t, err)
_, err = CreateProxyRoute(ctx, Input{SiteName: "second", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://other.example.com:8080"})
require.Error(t, err)
assert.Contains(t, err.Error(), "already exists")
require.NoError(t, DeleteProxyRoute(ctx, first.ID))
}
func TestListProxyRoutes(t *testing.T) {
func TestCreateProxyRouteHTTPSRequiresCoveringCertificate(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
first, err := CreateProxyRoute(ctx, Input{
SiteName: "first-site",
Domain: "first.example.com",
OriginURL: "http://origin-a.internal:80",
})
require.NoError(t, err)
second, err := CreateProxyRoute(ctx, Input{
SiteName: "second-site",
Domain: "second.example.com",
OriginURL: "http://origin-b.internal:80",
})
require.NoError(t, err)
routes, err := ListProxyRoutes(ctx)
require.NoError(t, err)
require.Len(t, routes, 2)
assert.Equal(t, second.ID, routes[0].ID)
assert.Equal(t, first.ID, routes[1].ID)
assert.Equal(t, "second.example.com", routes[0].Domain)
assert.Equal(t, "first.example.com", routes[1].Domain)
}
func TestValidateProxyRouteIdentityUniquenessUsesDecodedPrimaryDomain(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
existing := &model.ProxyRoute{
SiteName: "",
Domain: "legacy.example.com",
Domains: `["primary.example.com"]`,
OriginURL: "http://origin.example.com:8080",
Upstreams: `["http://origin.example.com:8080"]`,
Enabled: true,
UpstreamType: "direct",
}
require.NoError(t, model.CreateProxyRouteRecord(ctx, existing))
view, err := GetProxyRoute(ctx, existing.ID)
require.NoError(t, err)
assert.Equal(t, "primary.example.com", view.SiteName)
_, err = CreateProxyRoute(ctx, Input{
SiteName: "primary.example.com",
Domain: "other.example.com",
OriginURL: "http://origin-b.example.com:8080",
})
require.Error(t, err)
assert.Contains(t, err.Error(), "site_name already exists")
}
func TestUpdateProxyRouteAuthConfig(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateProxyRoute(ctx, Input{
SiteName: "auth-site",
Domain: "auth.example.com",
OriginURL: "http://origin.example.com:8080",
Enabled: true,
})
require.NoError(t, err)
updated, err := UpdateProxyRoute(ctx, created.ID, Input{
SiteName: created.SiteName,
Domain: created.Domain,
Domains: created.Domains,
OriginURL: created.OriginURL,
Enabled: created.Enabled,
BasicAuthEnabled: true,
BasicAuthUsername: "admin",
BasicAuthPassword: "secret",
})
require.NoError(t, err)
assert.True(t, updated.BasicAuthEnabled)
assert.Equal(t, "admin", updated.BasicAuthUsername)
assert.Equal(t, "secret", updated.BasicAuthPassword)
domain := createZoneDomain(t, ctx, "api.example.com", nil)
_, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080", EnableHTTPS: true})
require.EqualError(t, err, errProxyRouteCertRequired)
}
@@ -11,7 +11,6 @@ import (
"github.com/gin-gonic/gin"
)
func handleLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
@@ -138,4 +137,4 @@ func DeleteProxyRouteHandler(c *gin.Context) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
}
-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)
+9 -7
View File
@@ -12,19 +12,21 @@ import (
// ProxyRoute OpenFlare 代理规则实体。
type ProxyRoute struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
SiteName string `json:"site_name" gorm:"size:255;not null;default:''"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
Domains string `json:"domains" gorm:"type:text;not null;default:'[]'"`
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
SiteName string `json:"site_name" gorm:"size:255;not null;default:''"`
// Legacy mirrors are maintained from ZoneDomain bindings until the staged
// schema cleanup. They are not route API fields.
Domain string `json:"-" gorm:"uniqueIndex;size:255;not null"`
Domains string `json:"-" gorm:"type:text;not null;default:'[]'"`
OriginID *uint `json:"origin_id" gorm:"index"`
OriginURL string `json:"origin_url" gorm:"size:2048;not null"`
OriginHost string `json:"origin_host" gorm:"size:255"`
Upstreams string `json:"upstreams" gorm:"type:text;not null;default:'[]'"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"`
CertID *uint `json:"cert_id"`
CertIDs string `json:"cert_ids" gorm:"type:text;not null;default:'[]'"`
DomainCertIDs string `json:"domain_cert_ids" gorm:"type:text;not null;default:'[]'"`
CertID *uint `json:"-"`
CertIDs string `json:"-" gorm:"type:text;not null;default:'[]'"`
DomainCertIDs string `json:"-" gorm:"type:text;not null;default:'[]'"`
RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"`
LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"`
LimitConnPerIP int `json:"limit_conn_per_ip" gorm:"not null;default:0"`
+31
View File
@@ -61,6 +61,37 @@ func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]ZoneDomain,
return domains, nil
}
// ListZoneDomainsByIDs returns explicit domains in the requested ID order.
func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]ZoneDomain, error) {
if len(domainIDs) == 0 {
return []ZoneDomain{}, nil
}
var domains []ZoneDomain
if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil {
return nil, err
}
byID := make(map[uint]ZoneDomain, len(domains))
for _, domain := range domains {
byID[domain.ID] = domain
}
ordered := make([]ZoneDomain, 0, len(domainIDs))
for _, id := range domainIDs {
domain, ok := byID[id]
if !ok {
return nil, fmt.Errorf("one or more zone domains do not exist")
}
ordered = append(ordered, domain)
}
return ordered, nil
}
// CountZoneDomainsByCertificateID reports whether a certificate is assigned to a Zone domain.
func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) {
var count int64
err := db.DB(ctx).Model(&ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error
return count, err
}
// ReplaceZoneDomainRouteBindings replaces every ZoneDomain binding for a proxy route.
func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error {
conn := db.DB(ctx)