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)) items := make([]RouteSummary, 0, len(routes))
for _, route := range 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{ items = append(items, RouteSummary{
ID: route.ID, ID: route.ID,
Domain: route.Domain, Domain: domain,
OriginURL: route.OriginURL, OriginURL: route.OriginURL,
Enabled: route.Enabled, Enabled: route.Enabled,
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
@@ -11,15 +11,13 @@ import (
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
) )
type proxyRouteJSONFields struct { type proxyRouteJSONFields struct {
cacheRulesJSON string cacheRulesJSON string
upstreamsJSON string upstreamsJSON string
customHeadersJSON string customHeadersJSON string
certIDsJSON string
domainCertIDsJSON string
domainsJSON string
} }
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) { 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( func marshalProxyRouteJSONFields(
domains []string,
upstreams []string, upstreams []string,
cacheRules []string, cacheRules []string,
customHeaders []CustomHeaderInput, customHeaders []CustomHeaderInput,
certIDs []uint,
domainCertIDs []uint,
) (*proxyRouteJSONFields, error) { ) (*proxyRouteJSONFields, error) {
cacheRulesJSON, err := json.Marshal(cacheRules) cacheRulesJSON, err := json.Marshal(cacheRules)
if err != nil { if err != nil {
@@ -65,38 +60,13 @@ func marshalProxyRouteJSONFields(
if err != nil { if err != nil {
return nil, err 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{ return &proxyRouteJSONFields{
cacheRulesJSON: string(cacheRulesJSON), cacheRulesJSON: string(cacheRulesJSON),
upstreamsJSON: string(upstreamsJSON), upstreamsJSON: string(upstreamsJSON),
customHeadersJSON: string(customHeadersJSON), customHeadersJSON: string(customHeadersJSON),
certIDsJSON: string(certIDsJSON),
domainCertIDsJSON: string(domainCertIDsJSON),
domainsJSON: string(domainsJSON),
}, nil }, 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 { func normalizeProxyRouteBasicAuth(input *Input) error {
if !input.BasicAuthEnabled { if !input.BasicAuthEnabled {
input.BasicAuthUsername = "" input.BasicAuthUsername = ""
@@ -114,7 +84,7 @@ func normalizeProxyRouteBasicAuth(input *Input) error {
func populateProxyRouteFields( func populateProxyRouteFields(
route *model.ProxyRoute, route *model.ProxyRoute,
input Input, input Input,
siteName, domain string, siteName string,
jsonFields *proxyRouteJSONFields, jsonFields *proxyRouteJSONFields,
originID *uint, originID *uint,
upstreams []string, upstreams []string,
@@ -123,17 +93,12 @@ func populateProxyRouteFields(
limitRate, upstreamType string, limitRate, upstreamType string,
) { ) {
route.SiteName = siteName route.SiteName = siteName
route.Domain = domain
route.Domains = jsonFields.domainsJSON
route.OriginID = originID route.OriginID = originID
route.OriginURL = upstreams[0] route.OriginURL = upstreams[0]
route.OriginHost = originHost route.OriginHost = originHost
route.Upstreams = jsonFields.upstreamsJSON route.Upstreams = jsonFields.upstreamsJSON
route.Enabled = input.Enabled route.Enabled = input.Enabled
route.EnableHTTPS = input.EnableHTTPS route.EnableHTTPS = input.EnableHTTPS
route.CertID = input.CertID
route.CertIDs = jsonFields.certIDsJSON
route.DomainCertIDs = jsonFields.domainCertIDsJSON
route.RedirectHTTP = input.RedirectHTTP route.RedirectHTTP = input.RedirectHTTP
route.LimitConnPerServer = limitConnPerServer route.LimitConnPerServer = limitConnPerServer
route.LimitConnPerIP = limitConnPerIP route.LimitConnPerIP = limitConnPerIP
@@ -176,3 +141,40 @@ func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, u
} }
return nil 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 package proxy_route
const ( const (
errProxyRouteNotFound = "proxy route not found" errProxyRouteNotFound = "proxy route not found"
errProxyRouteIdentityExists = "proxy route identity already exists" errProxyRouteIdentityExists = "proxy route identity already exists"
errProxyRouteSiteNameExists = "site_name already exists" errProxyRouteSiteNameExists = "site_name already exists"
errProxyRouteDomainExists = "domain %s already exists" errProxyRouteDomainExists = "domain %s already exists"
errProxyRouteSiteNameEmpty = "site_name cannot be empty" errProxyRouteSiteNameEmpty = "site_name cannot be empty"
errProxyRouteDomainRequired = "at least one domain is required" errProxyRouteZoneDomainsRequired = "at least one zone domain is required"
errProxyRouteDomainInvalid = "domain format is invalid" errProxyRouteZoneDomainNotFound = "selected zone domain does not exist"
errProxyRouteDomainMismatch = "domain must match domains[0]" errProxyRouteZoneDomainDuplicate = "zone_domain_ids must not contain duplicates"
errProxyRouteOriginEmpty = "origin_url cannot be empty" errProxyRouteZoneDomainBound = "selected zone domain is already bound to another proxy route"
errProxyRouteOriginInvalid = "origin URL format is invalid" errProxyRouteOriginEmpty = "origin_url cannot be empty"
errProxyRouteOriginScheme = "origin URL must start with http:// or https://" errProxyRouteOriginInvalid = "origin URL format is invalid"
errProxyRouteOriginHostInvalid = "origin_host format is invalid" errProxyRouteOriginScheme = "origin URL must start with http:// or https://"
errProxyRouteUpstreamRequired = "at least one upstream is required" errProxyRouteOriginHostInvalid = "origin_host format is invalid"
errProxyRouteUpstreamScheme = "all upstreams must use the same scheme" errProxyRouteUpstreamRequired = "at least one upstream is required"
errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths" errProxyRouteUpstreamScheme = "all upstreams must use the same scheme"
errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings" errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths"
errProxyRouteOriginNotFound = "selected origin does not exist" errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings"
errProxyRouteCertNotFound = "selected certificate does not exist" errProxyRouteOriginNotFound = "selected origin does not exist"
errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled" errProxyRouteCertNotFound = "selected certificate does not exist"
errProxyRouteCertDomainLength = "domain_cert_ids must match domains length" errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled"
errProxyRouteRedirectHTTP = "redirect_http requires enable_https" errProxyRouteCertDomainLength = "domain_cert_ids must match domains length"
errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled" errProxyRouteRedirectHTTP = "redirect_http requires enable_https"
errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format" errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled"
errProxyRouteCachePolicy = "cache policy is not supported" errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format"
errProxyRouteCacheSuffix = "cache suffix format is invalid" errProxyRouteCachePolicy = "cache policy is not supported"
errProxyRouteCachePath = "cache path rule format is invalid" errProxyRouteCacheSuffix = "cache suffix format is invalid"
errProxyRouteCacheSuffixReq = "at least one suffix is required" errProxyRouteCachePath = "cache path rule format is invalid"
errProxyRouteCachePrefixReq = "at least one path prefix is required" errProxyRouteCacheSuffixReq = "at least one suffix is required"
errProxyRouteCacheExactReq = "at least one exact path is required" errProxyRouteCachePrefixReq = "at least one path prefix is required"
errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty" errProxyRouteCacheExactReq = "at least one exact path is required"
errProxyRouteHeaderKeyInvalid = "custom header key format is invalid" errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty"
errProxyRouteHeaderNewline = "custom headers cannot contain newlines" errProxyRouteHeaderKeyInvalid = "custom header key format is invalid"
errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream" errProxyRouteHeaderNewline = "custom headers cannot contain newlines"
errProxyRouteTunnelNodeMissing = "tunnel client node does not exist" errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream"
errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node" errProxyRouteTunnelNodeMissing = "tunnel client node does not exist"
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream" errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node"
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https" errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream" errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
errProxyRoutePagesNotFound = "pages 项目不存在" errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
errProxyRoutePagesDisabled = "pages 项目未启用" errProxyRoutePagesNotFound = "pages 项目不存在"
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署" errProxyRoutePagesDisabled = "pages 项目未启用"
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https" errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
errProxyRouteOriginPort = "端口格式不合法" errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
errProxyRouteOriginPortEmpty = "端口不能为空" errProxyRouteOriginPort = "端口格式不合法"
errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头" errProxyRouteOriginPortEmpty = "端口不能为空"
errProxyRouteOriginURIProto = "源站路径不能包含协议" errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头"
errProxyRouteOriginURIProto = "源站路径不能包含协议"
) )
+33 -242
View File
@@ -17,7 +17,6 @@ import (
"strings" "strings"
"unicode" "unicode"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -254,63 +253,23 @@ func validateCertificateCoverage(certificate *model.TLSCertificate, domains []st
return nil return nil
} }
func loadTLSCertificates(ctx context.Context, certIDs []uint) ([]*model.TLSCertificate, error) { func loadProxyRouteZoneDomains(ctx context.Context, ids []uint) ([]model.ZoneDomain, error) {
certificates := make([]*model.TLSCertificate, 0, len(certIDs)) if len(ids) == 0 {
for _, certID := range certIDs { return nil, errors.New(errProxyRouteZoneDomainsRequired)
certificate, err := lookupTLSCertificateByID(ctx, certID)
if err != nil {
return nil, err
}
certificates = append(certificates, certificate)
} }
return certificates, nil seen := make(map[uint]struct{}, len(ids))
} for _, id := range ids {
if id == 0 {
func normalizeProxyRouteDomainValue(raw string) string { return nil, errors.New(errProxyRouteZoneDomainNotFound)
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)
} }
return err if _, ok := seen[id]; ok {
} return nil, errors.New(errProxyRouteZoneDomainDuplicate)
}
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)
} }
domain := normalizeProxyRouteDomainValue(rawDomain) seen[id] = struct{}{}
if domain != "" && domain != domains[0] {
return nil, errors.New(errProxyRouteDomainMismatch)
}
return domains, nil
} }
domains, err := model.ListZoneDomainsByIDs(ctx, ids)
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})
if err != nil { if err != nil {
return nil, mapRouteIdentityDomainError(err) return nil, errors.New(errProxyRouteZoneDomainNotFound)
} }
return domains, nil return domains, nil
} }
@@ -322,7 +281,7 @@ func validateProxyRouteSiteName(siteName string) error {
return nil 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) routes, err := model.ListProxyRoutes(ctx)
if err != nil { if err != nil {
return err return err
@@ -337,27 +296,33 @@ func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.Prox
if item == nil || item.ID == currentID { if item == nil || item.ID == currentID {
continue continue
} }
existingSiteName, existingDomains, err := routeidentity.ResolveFromRoute(item) if item.SiteName == siteName {
if err != nil {
return fmt.Errorf("existing route %d domains are invalid: %w", item.ID, err)
}
if existingSiteName == siteName {
return errors.New(errProxyRouteSiteNameExists) 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 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) { func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) {
if value < 0 { if value < 0 {
return 0, fmt.Errorf("%s must be greater than or equal to 0", field) 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 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) { func normalizeProxyRouteLimitRate(raw string) (string, error) {
normalized := strings.ToLower(strings.TrimSpace(raw)) normalized := strings.ToLower(strings.TrimSpace(raw))
if normalized == "" || normalized == "0" { if normalized == "" || normalized == "0" {
@@ -727,36 +548,6 @@ func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, erro
return normalizeUpstreams(fallbackOriginURL, upstreams) 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 { func validateOriginURL(raw string) error {
if raw == "" { if raw == "" {
return errors.New(errProxyRouteOriginEmpty) return errors.New(errProxyRouteOriginEmpty)
+97 -80
View File
@@ -5,12 +5,14 @@ package proxy_route
import ( import (
"context" "context"
"encoding/json"
"errors" "errors"
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity" "github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
) )
// CustomHeaderInput 自定义响应头。 // CustomHeaderInput 自定义响应头。
@@ -22,8 +24,7 @@ type CustomHeaderInput struct {
// Input 代理规则创建/更新请求。 // Input 代理规则创建/更新请求。
type Input struct { type Input struct {
SiteName string `json:"site_name"` SiteName string `json:"site_name"`
Domain string `json:"domain"` ZoneDomainIDs []uint `json:"zone_domain_ids"`
Domains []string `json:"domains"`
OriginID *uint `json:"origin_id"` OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"` OriginURL string `json:"origin_url"`
OriginScheme string `json:"origin_scheme"` OriginScheme string `json:"origin_scheme"`
@@ -34,9 +35,6 @@ type Input struct {
Upstreams []string `json:"upstreams"` Upstreams []string `json:"upstreams"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"` 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"` RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server"` LimitConnPerServer int `json:"limit_conn_per_server"`
LimitConnPerIP int `json:"limit_conn_per_ip"` LimitConnPerIP int `json:"limit_conn_per_ip"`
@@ -61,10 +59,8 @@ type Input struct {
type View struct { type View struct {
ID uint `json:"id"` ID uint `json:"id"`
SiteName string `json:"site_name"` SiteName string `json:"site_name"`
Domain string `json:"domain"` ZoneDomainIDs []uint `json:"zone_domain_ids"`
Domains []string `json:"domains"` ZoneDomains []ZoneDomainView `json:"zone_domains"`
PrimaryDomain string `json:"primary_domain"`
DomainCount int `json:"domain_count"`
OriginID *uint `json:"origin_id"` OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"` OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host"` OriginHost string `json:"origin_host"`
@@ -72,9 +68,6 @@ type View struct {
UpstreamList []string `json:"upstream_list"` UpstreamList []string `json:"upstream_list"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"` 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"` RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server"` LimitConnPerServer int `json:"limit_conn_per_server"`
LimitConnPerIP int `json:"limit_conn_per_ip"` LimitConnPerIP int `json:"limit_conn_per_ip"`
@@ -99,6 +92,14 @@ type View struct {
UpdatedAt time.Time `json:"updated_at"` 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 列出全部代理规则。 // ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*View, error) { func ListProxyRoutes(ctx context.Context) ([]*View, error) {
routes, err := model.ListProxyRoutes(ctx) routes, err := model.ListProxyRoutes(ctx)
@@ -119,11 +120,16 @@ func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
// CreateProxyRoute 创建代理规则。 // CreateProxyRoute 创建代理规则。
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) { func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
route, err := buildProxyRoute(ctx, nil, input) route, _, err := buildProxyRoute(ctx, nil, input)
if err != nil { if err != nil {
return nil, err 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) { if isUniqueConstraintError(err) {
return nil, errors.New(errProxyRouteIdentityExists) return nil, errors.New(errProxyRouteIdentityExists)
} }
@@ -138,11 +144,16 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
if err != nil { if err != nil {
return nil, err return nil, err
} }
route, err = buildProxyRoute(ctx, route, input) route, _, err = buildProxyRoute(ctx, route, input)
if err != nil { if err != nil {
return nil, err 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) { if isUniqueConstraintError(err) {
return nil, errors.New(errProxyRouteIdentityExists) 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 { if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
return err 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) { func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) {
domains, err := normalizeProxyRouteDomainsInput(route, input.Domain, input.Domains) domains, err := loadProxyRouteZoneDomains(ctx, input.ZoneDomainIDs)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
domain := domains[0] siteName := strings.TrimSpace(input.SiteName)
siteName := routeidentity.ResolveSiteName(route, input.SiteName, domain)
upstreamType := normalizeUpstreamType(input.UpstreamType) upstreamType := normalizeUpstreamType(input.UpstreamType)
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input) _, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
originHost := strings.TrimSpace(input.OriginHost) originHost := strings.TrimSpace(input.OriginHost)
remark := strings.TrimSpace(input.Remark) remark := strings.TrimSpace(input.Remark)
cachePolicy := strings.TrimSpace(input.CachePolicy) cachePolicy := strings.TrimSpace(input.CachePolicy)
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules) cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders) customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server") limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip") limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate) limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
if err != nil { if err != nil {
return nil, err return nil, nil, err
} }
if err := validateProxyRouteZoneDomainCertificates(ctx, domains, input.EnableHTTPS); err != nil {
normalizeProxyRouteHTTPSInput(&input) return nil, nil, err
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs( }
ctx, jsonFields, err := marshalProxyRouteJSONFields(upstreams, cacheRules, customHeaders)
domains,
input.EnableHTTPS,
input.DomainCertIDs,
input.CertID,
input.CertIDs,
)
if err != nil { if err != nil {
return nil, err return nil, 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
} }
if err := validateProxyRouteSiteName(siteName); err != nil { if err := validateProxyRouteSiteName(siteName); err != nil {
return nil, err return nil, nil, err
} }
if err := validateProxyRouteIdentityUniqueness(ctx, route, siteName, domains); err != nil { if err := validateProxyRouteSiteNameUniqueness(ctx, route, siteName); err != nil {
return nil, err return nil, nil, err
} }
if err := validateOriginHost(originHost); err != nil { 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 { if input.RedirectHTTP && !input.EnableHTTPS {
return nil, errors.New(errProxyRouteRedirectHTTP) return nil, nil, errors.New(errProxyRouteRedirectHTTP)
} }
if err := normalizeProxyRouteBasicAuth(&input); err != nil { if err := normalizeProxyRouteBasicAuth(&input); err != nil {
return nil, err return nil, nil, err
} }
if route == nil { if route == nil {
@@ -243,7 +242,6 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
route, route,
input, input,
siteName, siteName,
domain,
jsonFields, jsonFields,
originID, originID,
upstreams, upstreams,
@@ -255,10 +253,41 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
limitRate, limitRate,
upstreamType, 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 { 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) { 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 { if route == nil {
return nil, errors.New("proxy route is 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 { if err != nil {
return nil, err return nil, err
} }
@@ -293,26 +322,17 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
if err != nil { if err != nil {
return nil, err return nil, err
} }
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) zoneDomainIDs := make([]uint, 0, len(domains))
if err != nil { zoneDomains := make([]ZoneDomainView, 0, len(domains))
return nil, err 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{ return &View{
ID: route.ID, ID: route.ID,
SiteName: routeidentity.ResolveSiteName(route, route.SiteName, primaryDomain), SiteName: route.SiteName,
Domain: primaryDomain, ZoneDomainIDs: zoneDomainIDs,
Domains: domains, ZoneDomains: zoneDomains,
PrimaryDomain: primaryDomain,
DomainCount: len(domains),
OriginID: route.OriginID, OriginID: route.OriginID,
OriginURL: route.OriginURL, OriginURL: route.OriginURL,
OriginHost: route.OriginHost, OriginHost: route.OriginHost,
@@ -320,9 +340,6 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
UpstreamList: upstreams, UpstreamList: upstreams,
Enabled: route.Enabled, Enabled: route.Enabled,
EnableHTTPS: route.EnableHTTPS, EnableHTTPS: route.EnableHTTPS,
CertID: certID,
CertIDs: certIDs,
DomainCertIDs: domainCertIDs,
RedirectHTTP: route.RedirectHTTP, RedirectHTTP: route.RedirectHTTP,
LimitConnPerServer: route.LimitConnPerServer, LimitConnPerServer: route.LimitConnPerServer,
LimitConnPerIP: route.LimitConnPerIP, LimitConnPerIP: route.LimitConnPerIP,
@@ -17,130 +17,66 @@ import (
func setupProxyRouteTestDB(t *testing.T) func() { func setupProxyRouteTestDB(t *testing.T) func() {
t.Helper() 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, 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) db.SetDB(sqliteDB)
return func() { return func() { db.SetDB(nil) }
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) cleanup := setupProxyRouteTestDB(t)
defer cleanup() defer cleanup()
ctx := context.Background() ctx := context.Background()
domainA := createZoneDomain(t, ctx, "api.example.com", nil)
domainB := createZoneDomain(t, ctx, "www.example.com", nil)
view, err := CreateProxyRoute(ctx, Input{ view, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domainA.ID, domainB.ID}, OriginURL: "http://origin.example.com:8080", Enabled: true})
SiteName: "example-site",
Domain: "example.com",
OriginURL: "http://origin.example.com:8080",
Enabled: true,
})
require.NoError(t, err) require.NoError(t, err)
assert.NotZero(t, view.ID) assert.Equal(t, []uint{domainA.ID, domainB.ID}, view.ZoneDomainIDs)
assert.Equal(t, "example-site", view.SiteName) require.Len(t, view.ZoneDomains, 2)
assert.Equal(t, "example.com", view.Domain) assert.Equal(t, "api.example.com", view.ZoneDomains[0].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)
_, err = CreateProxyRoute(ctx, Input{ func TestCreateProxyRouteRejectsInvalidZoneDomainBindings(t *testing.T) {
SiteName: "duplicate-site", cleanup := setupProxyRouteTestDB(t)
Domain: "example.com", defer cleanup()
OriginURL: "http://origin.example.com:8080", 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) 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) cleanup := setupProxyRouteTestDB(t)
defer cleanup() defer cleanup()
ctx := context.Background() ctx := context.Background()
domain := createZoneDomain(t, ctx, "api.example.com", nil)
first, err := CreateProxyRoute(ctx, Input{ _, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080", EnableHTTPS: true})
SiteName: "first-site", require.EqualError(t, err, errProxyRouteCertRequired)
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)
} }
@@ -11,7 +11,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func handleLogicError(c *gin.Context, err error) bool { func handleLogicError(c *gin.Context, err error) bool {
if err == nil { if err == nil {
return false return false
@@ -138,4 +137,4 @@ func DeleteProxyRouteHandler(c *gin.Context) {
return return
} }
c.JSON(http.StatusOK, response.OKNil()) c.JSON(http.StatusOK, response.OKNil())
} }
-17
View File
@@ -5,10 +5,8 @@ package tls
import ( import (
"crypto/x509" "crypto/x509"
"encoding/json"
"encoding/pem" "encoding/pem"
"errors" "errors"
"fmt"
"io" "io"
"mime/multipart" "mime/multipart"
"strings" "strings"
@@ -45,18 +43,3 @@ func isUniqueConstraintError(err error) bool {
} }
return strings.Contains(strings.ToLower(err.Error()), "unique") 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 { func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
routes, err := model.ListTLSProxyRouteRefs(ctx) count, err := model.CountZoneDomainsByCertificateID(ctx, id)
if err != nil { if err != nil {
return err return err
} }
for _, route := range routes { if count > 0 {
if route.CertID != nil && *route.CertID == id { return errors.New(errCertificateDeleteReferenced)
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)
}
}
} }
return nil return nil
} }
@@ -40,6 +40,8 @@ func setupTLSTestDB(t *testing.T) func() {
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate( require.NoError(t, sqliteDB.AutoMigrate(
&model.TLSCertificate{}, &model.TLSCertificate{},
&model.Zone{},
&model.ZoneDomain{},
&model.ManagedDomain{}, &model.ManagedDomain{},
&model.DNSAccount{}, &model.DNSAccount{},
&model.AcmeAccount{}, &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) { func generateTestCertificatePair(t *testing.T, dnsNames []string) (string, string) {
t.Helper() t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048) privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
+9 -7
View File
@@ -12,19 +12,21 @@ import (
// ProxyRoute OpenFlare 代理规则实体。 // ProxyRoute OpenFlare 代理规则实体。
type ProxyRoute struct { type ProxyRoute struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"` ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
SiteName string `json:"site_name" gorm:"size:255;not null;default:''"` SiteName string `json:"site_name" gorm:"size:255;not null;default:''"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"` // Legacy mirrors are maintained from ZoneDomain bindings until the staged
Domains string `json:"domains" gorm:"type:text;not null;default:'[]'"` // 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"` OriginID *uint `json:"origin_id" gorm:"index"`
OriginURL string `json:"origin_url" gorm:"size:2048;not null"` OriginURL string `json:"origin_url" gorm:"size:2048;not null"`
OriginHost string `json:"origin_host" gorm:"size:255"` OriginHost string `json:"origin_host" gorm:"size:255"`
Upstreams string `json:"upstreams" gorm:"type:text;not null;default:'[]'"` Upstreams string `json:"upstreams" gorm:"type:text;not null;default:'[]'"`
Enabled bool `json:"enabled" gorm:"not null;default:true"` Enabled bool `json:"enabled" gorm:"not null;default:true"`
EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"` EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"`
CertID *uint `json:"cert_id"` CertID *uint `json:"-"`
CertIDs string `json:"cert_ids" gorm:"type:text;not null;default:'[]'"` CertIDs string `json:"-" gorm:"type:text;not null;default:'[]'"`
DomainCertIDs string `json:"domain_cert_ids" 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"` RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"`
LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"` LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"`
LimitConnPerIP int `json:"limit_conn_per_ip" 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 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. // ReplaceZoneDomainRouteBindings replaces every ZoneDomain binding for a proxy route.
func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error { func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error {
conn := db.DB(ctx) conn := db.DB(ctx)