From d0536fcdd5a4044c7a879577ed26eaeff26f5b60 Mon Sep 17 00:00:00 2001 From: ryan Date: Sun, 12 Jul 2026 14:52:19 +0800 Subject: [PATCH] refactor(proxy): bind routes through zone domains --- internal/apps/openflare/origin/logics.go | 10 +- .../openflare/proxy_route/build_helpers.go | 76 ++--- .../openflare/proxy_route/cert_helpers.go | 71 ----- internal/apps/openflare/proxy_route/errs.go | 93 +++--- .../apps/openflare/proxy_route/helpers.go | 275 +++--------------- internal/apps/openflare/proxy_route/logics.go | 177 ++++++----- .../apps/openflare/proxy_route/logics_test.go | 154 +++------- .../apps/openflare/proxy_route/routers.go | 3 +- internal/apps/openflare/tls/helpers.go | 17 -- internal/apps/openflare/tls/logics.go | 29 +- internal/apps/openflare/tls/logics_test.go | 18 ++ internal/model/openflare_proxy_route.go | 16 +- internal/model/openflare_zone.go | 31 ++ 13 files changed, 332 insertions(+), 638 deletions(-) delete mode 100644 internal/apps/openflare/proxy_route/cert_helpers.go diff --git a/internal/apps/openflare/origin/logics.go b/internal/apps/openflare/origin/logics.go index f4ea24b6..cfdf332f 100644 --- a/internal/apps/openflare/origin/logics.go +++ b/internal/apps/openflare/origin/logics.go @@ -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"), diff --git a/internal/apps/openflare/proxy_route/build_helpers.go b/internal/apps/openflare/proxy_route/build_helpers.go index d275227e..a667afc8 100644 --- a/internal/apps/openflare/proxy_route/build_helpers.go +++ b/internal/apps/openflare/proxy_route/build_helpers.go @@ -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 +} diff --git a/internal/apps/openflare/proxy_route/cert_helpers.go b/internal/apps/openflare/proxy_route/cert_helpers.go deleted file mode 100644 index de505514..00000000 --- a/internal/apps/openflare/proxy_route/cert_helpers.go +++ /dev/null @@ -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 - } -} diff --git a/internal/apps/openflare/proxy_route/errs.go b/internal/apps/openflare/proxy_route/errs.go index b72b9579..1b8270b3 100644 --- a/internal/apps/openflare/proxy_route/errs.go +++ b/internal/apps/openflare/proxy_route/errs.go @@ -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 = "源站路径不能包含协议" ) diff --git a/internal/apps/openflare/proxy_route/helpers.go b/internal/apps/openflare/proxy_route/helpers.go index f9169e84..a0095e37 100644 --- a/internal/apps/openflare/proxy_route/helpers.go +++ b/internal/apps/openflare/proxy_route/helpers.go @@ -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) diff --git a/internal/apps/openflare/proxy_route/logics.go b/internal/apps/openflare/proxy_route/logics.go index 561cfe14..dfd1c921 100644 --- a/internal/apps/openflare/proxy_route/logics.go +++ b/internal/apps/openflare/proxy_route/logics.go @@ -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, diff --git a/internal/apps/openflare/proxy_route/logics_test.go b/internal/apps/openflare/proxy_route/logics_test.go index 21b156f7..84bfd4ab 100644 --- a/internal/apps/openflare/proxy_route/logics_test.go +++ b/internal/apps/openflare/proxy_route/logics_test.go @@ -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) } diff --git a/internal/apps/openflare/proxy_route/routers.go b/internal/apps/openflare/proxy_route/routers.go index e9214a73..cef9f52a 100644 --- a/internal/apps/openflare/proxy_route/routers.go +++ b/internal/apps/openflare/proxy_route/routers.go @@ -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()) -} \ No newline at end of file +} diff --git a/internal/apps/openflare/tls/helpers.go b/internal/apps/openflare/tls/helpers.go index 839aa8c2..aa4b3bb8 100644 --- a/internal/apps/openflare/tls/helpers.go +++ b/internal/apps/openflare/tls/helpers.go @@ -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 -} diff --git a/internal/apps/openflare/tls/logics.go b/internal/apps/openflare/tls/logics.go index 3fe462c7..36afe517 100644 --- a/internal/apps/openflare/tls/logics.go +++ b/internal/apps/openflare/tls/logics.go @@ -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 } diff --git a/internal/apps/openflare/tls/logics_test.go b/internal/apps/openflare/tls/logics_test.go index 7b1ca131..79997c8e 100644 --- a/internal/apps/openflare/tls/logics_test.go +++ b/internal/apps/openflare/tls/logics_test.go @@ -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) diff --git a/internal/model/openflare_proxy_route.go b/internal/model/openflare_proxy_route.go index 7139081b..c2aa899c 100644 --- a/internal/model/openflare_proxy_route.go +++ b/internal/model/openflare_proxy_route.go @@ -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"` diff --git a/internal/model/openflare_zone.go b/internal/model/openflare_zone.go index a9e42f45..b203ff75 100644 --- a/internal/model/openflare_zone.go +++ b/internal/model/openflare_zone.go @@ -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)