mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
refactor(proxy): bind routes through zone domains
This commit is contained in:
@@ -74,9 +74,17 @@ func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
|
||||
}
|
||||
items := make([]RouteSummary, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domain := ""
|
||||
if len(domains) > 0 {
|
||||
domain = domains[0].Domain
|
||||
}
|
||||
items = append(items, RouteSummary{
|
||||
ID: route.ID,
|
||||
Domain: route.Domain,
|
||||
Domain: domain,
|
||||
OriginURL: route.OriginURL,
|
||||
Enabled: route.Enabled,
|
||||
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
|
||||
@@ -11,15 +11,13 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type proxyRouteJSONFields struct {
|
||||
cacheRulesJSON string
|
||||
upstreamsJSON string
|
||||
customHeadersJSON string
|
||||
certIDsJSON string
|
||||
domainCertIDsJSON string
|
||||
domainsJSON string
|
||||
}
|
||||
|
||||
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) {
|
||||
@@ -46,12 +44,9 @@ func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input
|
||||
}
|
||||
|
||||
func marshalProxyRouteJSONFields(
|
||||
domains []string,
|
||||
upstreams []string,
|
||||
cacheRules []string,
|
||||
customHeaders []CustomHeaderInput,
|
||||
certIDs []uint,
|
||||
domainCertIDs []uint,
|
||||
) (*proxyRouteJSONFields, error) {
|
||||
cacheRulesJSON, err := json.Marshal(cacheRules)
|
||||
if err != nil {
|
||||
@@ -65,38 +60,13 @@ func marshalProxyRouteJSONFields(
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDsJSON, err := json.Marshal(certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainsJSON, err := json.Marshal(domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &proxyRouteJSONFields{
|
||||
cacheRulesJSON: string(cacheRulesJSON),
|
||||
upstreamsJSON: string(upstreamsJSON),
|
||||
customHeadersJSON: string(customHeadersJSON),
|
||||
certIDsJSON: string(certIDsJSON),
|
||||
domainCertIDsJSON: string(domainCertIDsJSON),
|
||||
domainsJSON: string(domainsJSON),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteHTTPSInput(input *Input) {
|
||||
if input.EnableHTTPS {
|
||||
return
|
||||
}
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteBasicAuth(input *Input) error {
|
||||
if !input.BasicAuthEnabled {
|
||||
input.BasicAuthUsername = ""
|
||||
@@ -114,7 +84,7 @@ func normalizeProxyRouteBasicAuth(input *Input) error {
|
||||
func populateProxyRouteFields(
|
||||
route *model.ProxyRoute,
|
||||
input Input,
|
||||
siteName, domain string,
|
||||
siteName string,
|
||||
jsonFields *proxyRouteJSONFields,
|
||||
originID *uint,
|
||||
upstreams []string,
|
||||
@@ -123,17 +93,12 @@ func populateProxyRouteFields(
|
||||
limitRate, upstreamType string,
|
||||
) {
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = jsonFields.domainsJSON
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = jsonFields.upstreamsJSON
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.CertIDs = jsonFields.certIDsJSON
|
||||
route.DomainCertIDs = jsonFields.domainCertIDsJSON
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.LimitConnPerServer = limitConnPerServer
|
||||
route.LimitConnPerIP = limitConnPerIP
|
||||
@@ -176,3 +141,40 @@ func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, u
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateProxyRouteRecord(tx *gorm.DB, route *model.ProxyRoute) error {
|
||||
return tx.Model(&model.ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
|
||||
"site_name": route.SiteName, "origin_id": route.OriginID, "origin_url": route.OriginURL,
|
||||
"domain": route.Domain, "domains": route.Domains, "cert_id": route.CertID,
|
||||
"cert_ids": route.CertIDs, "domain_cert_ids": route.DomainCertIDs,
|
||||
"origin_host": route.OriginHost, "upstreams": route.Upstreams, "enabled": route.Enabled,
|
||||
"enable_https": route.EnableHTTPS, "redirect_http": route.RedirectHTTP,
|
||||
"limit_conn_per_server": route.LimitConnPerServer, "limit_conn_per_ip": route.LimitConnPerIP,
|
||||
"limit_rate": route.LimitRate, "cache_enabled": route.CacheEnabled, "cache_policy": route.CachePolicy,
|
||||
"cache_rules": route.CacheRules, "custom_headers": route.CustomHeaders,
|
||||
"basic_auth_enabled": route.BasicAuthEnabled, "basic_auth_username": route.BasicAuthUsername,
|
||||
"basic_auth_password": route.BasicAuthPassword, "remark": route.Remark,
|
||||
"upstream_type": route.UpstreamType, "tunnel_node_id": route.TunnelNodeID,
|
||||
"tunnel_target_addr": route.TunnelTargetAddr, "tunnel_target_protocol": route.TunnelTargetProtocol,
|
||||
"pages_project_id": route.PagesProjectID,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func replaceZoneDomainRouteBindings(tx *gorm.DB, routeID uint, domainIDs []uint) error {
|
||||
var requested []model.ZoneDomain
|
||||
if err := tx.Where("id IN ?", domainIDs).Find(&requested).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if len(requested) != len(domainIDs) {
|
||||
return errors.New(errProxyRouteZoneDomainNotFound)
|
||||
}
|
||||
for _, domain := range requested {
|
||||
if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID {
|
||||
return errors.New(errProxyRouteZoneDomainBound)
|
||||
}
|
||||
}
|
||||
if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ? AND id NOT IN ?", routeID, domainIDs).Update("proxy_route_id", nil).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error
|
||||
}
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
func normalizeExplicitDomainCertIDs(ctx context.Context, domains []string, rawDomainCertIDs []uint) ([]uint, []uint, *uint, error) {
|
||||
if len(rawDomainCertIDs) != len(domains) {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
|
||||
}
|
||||
|
||||
normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs))
|
||||
uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs))
|
||||
seen := make(map[uint]struct{}, len(rawDomainCertIDs))
|
||||
hasAssignedCertificate := false
|
||||
for index, item := range rawDomainCertIDs {
|
||||
if item == 0 {
|
||||
continue
|
||||
}
|
||||
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
|
||||
}
|
||||
normalizedDomainCertIDs[index] = item
|
||||
hasAssignedCertificate = true
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
uniqueCertIDs = append(uniqueCertIDs, item)
|
||||
}
|
||||
if !hasAssignedCertificate {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
}
|
||||
|
||||
primaryCertID := &uniqueCertIDs[0]
|
||||
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
|
||||
}
|
||||
|
||||
func normalizeDerivedDomainCertIDs(
|
||||
ctx context.Context,
|
||||
domains []string,
|
||||
normalizedCertIDs []uint,
|
||||
) ([]uint, []uint, *uint, error) {
|
||||
switch {
|
||||
case len(normalizedCertIDs) == 0:
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
case len(normalizedCertIDs) == 1:
|
||||
domainCertIDs := make([]uint, len(domains))
|
||||
for index := range domainCertIDs {
|
||||
domainCertIDs[index] = normalizedCertIDs[0]
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
case len(normalizedCertIDs) == len(domains):
|
||||
domainCertIDs := make([]uint, len(normalizedCertIDs))
|
||||
copy(domainCertIDs, normalizedCertIDs)
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
default:
|
||||
domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(ctx, domains, normalizedCertIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
}
|
||||
}
|
||||
@@ -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 = "源站路径不能包含协议"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -5,12 +5,14 @@ package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// CustomHeaderInput 自定义响应头。
|
||||
@@ -22,8 +24,7 @@ type CustomHeaderInput struct {
|
||||
// Input 代理规则创建/更新请求。
|
||||
type Input struct {
|
||||
SiteName string `json:"site_name"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains"`
|
||||
ZoneDomainIDs []uint `json:"zone_domain_ids"`
|
||||
OriginID *uint `json:"origin_id"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginScheme string `json:"origin_scheme"`
|
||||
@@ -34,9 +35,6 @@ type Input struct {
|
||||
Upstreams []string `json:"upstreams"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs []uint `json:"cert_ids"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip"`
|
||||
@@ -61,10 +59,8 @@ type Input struct {
|
||||
type View struct {
|
||||
ID uint `json:"id"`
|
||||
SiteName string `json:"site_name"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains"`
|
||||
PrimaryDomain string `json:"primary_domain"`
|
||||
DomainCount int `json:"domain_count"`
|
||||
ZoneDomainIDs []uint `json:"zone_domain_ids"`
|
||||
ZoneDomains []ZoneDomainView `json:"zone_domains"`
|
||||
OriginID *uint `json:"origin_id"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginHost string `json:"origin_host"`
|
||||
@@ -72,9 +68,6 @@ type View struct {
|
||||
UpstreamList []string `json:"upstream_list"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs []uint `json:"cert_ids"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip"`
|
||||
@@ -99,6 +92,14 @@ type View struct {
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ZoneDomainView is the route-safe representation of a bound Zone domain.
|
||||
type ZoneDomainView struct {
|
||||
ID uint `json:"id"`
|
||||
ZoneID uint `json:"zone_id"`
|
||||
Domain string `json:"domain"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
}
|
||||
|
||||
// ListProxyRoutes 列出全部代理规则。
|
||||
func ListProxyRoutes(ctx context.Context) ([]*View, error) {
|
||||
routes, err := model.ListProxyRoutes(ctx)
|
||||
@@ -119,11 +120,16 @@ func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
|
||||
|
||||
// CreateProxyRoute 创建代理规则。
|
||||
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
|
||||
route, err := buildProxyRoute(ctx, nil, input)
|
||||
route, _, err := buildProxyRoute(ctx, nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.CreateProxyRouteRecord(ctx, route); err != nil {
|
||||
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Create(route).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs)
|
||||
}); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errProxyRouteIdentityExists)
|
||||
}
|
||||
@@ -138,11 +144,16 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
route, err = buildProxyRoute(ctx, route, input)
|
||||
route, _, err = buildProxyRoute(ctx, route, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.UpdateProxyRouteRecord(ctx, route); err != nil {
|
||||
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := updateProxyRouteRecord(tx, route); err != nil {
|
||||
return err
|
||||
}
|
||||
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs)
|
||||
}); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errProxyRouteIdentityExists)
|
||||
}
|
||||
@@ -156,84 +167,72 @@ func DeleteProxyRoute(ctx context.Context, id uint) error {
|
||||
if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return model.DeleteProxyRouteRecord(ctx, id)
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", id).Update("proxy_route_id", nil).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.ProxyRoute{}, id).Error
|
||||
})
|
||||
}
|
||||
|
||||
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, error) {
|
||||
domains, err := normalizeProxyRouteDomainsInput(route, input.Domain, input.Domains)
|
||||
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) {
|
||||
domains, err := loadProxyRouteZoneDomains(ctx, input.ZoneDomainIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
domain := domains[0]
|
||||
siteName := routeidentity.ResolveSiteName(route, input.SiteName, domain)
|
||||
siteName := strings.TrimSpace(input.SiteName)
|
||||
|
||||
upstreamType := normalizeUpstreamType(input.UpstreamType)
|
||||
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
originHost := strings.TrimSpace(input.OriginHost)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
cachePolicy := strings.TrimSpace(input.CachePolicy)
|
||||
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
normalizeProxyRouteHTTPSInput(&input)
|
||||
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
|
||||
ctx,
|
||||
domains,
|
||||
input.EnableHTTPS,
|
||||
input.DomainCertIDs,
|
||||
input.CertID,
|
||||
input.CertIDs,
|
||||
)
|
||||
if err := validateProxyRouteZoneDomainCertificates(ctx, domains, input.EnableHTTPS); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
jsonFields, err := marshalProxyRouteJSONFields(upstreams, cacheRules, customHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteDomainCertificateCoverage(ctx, domains, domainCertIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
jsonFields, err := marshalProxyRouteJSONFields(domains, upstreams, cacheRules, customHeaders, certIDs, domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
if err := validateProxyRouteSiteName(siteName); err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
if err := validateProxyRouteIdentityUniqueness(ctx, route, siteName, domains); err != nil {
|
||||
return nil, err
|
||||
if err := validateProxyRouteSiteNameUniqueness(ctx, route, siteName); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if err := validateOriginHost(originHost); err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
input.DomainCertIDs = domainCertIDs
|
||||
input.CertIDs = certIDs
|
||||
input.CertID = primaryCertID
|
||||
if input.RedirectHTTP && !input.EnableHTTPS {
|
||||
return nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
return nil, nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
}
|
||||
|
||||
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
if route == nil {
|
||||
@@ -243,7 +242,6 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
route,
|
||||
input,
|
||||
siteName,
|
||||
domain,
|
||||
jsonFields,
|
||||
originID,
|
||||
upstreams,
|
||||
@@ -255,10 +253,41 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
limitRate,
|
||||
upstreamType,
|
||||
)
|
||||
// Preserve legacy columns until the second-phase schema cleanup. They are
|
||||
// derived solely from ZoneDomain bindings and are not exposed by this API.
|
||||
populateLegacyZoneDomainFields(route, domains)
|
||||
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
return route, nil
|
||||
return route, domains, nil
|
||||
}
|
||||
|
||||
func mustMarshalProxyRouteLegacy(value any) string {
|
||||
encoded, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return string(encoded)
|
||||
}
|
||||
|
||||
func populateLegacyZoneDomainFields(route *model.ProxyRoute, domains []model.ZoneDomain) {
|
||||
legacyDomains := make([]string, 0, len(domains))
|
||||
legacyCertIDs := make([]uint, 0, len(domains))
|
||||
for _, domain := range domains {
|
||||
legacyDomains = append(legacyDomains, domain.Domain)
|
||||
if domain.CertID != nil {
|
||||
legacyCertIDs = append(legacyCertIDs, *domain.CertID)
|
||||
}
|
||||
}
|
||||
route.Domain = legacyDomains[0]
|
||||
route.Domains = mustMarshalProxyRouteLegacy(legacyDomains)
|
||||
route.CertIDs = mustMarshalProxyRouteLegacy(legacyCertIDs)
|
||||
if len(legacyCertIDs) > 0 {
|
||||
route.CertID = &legacyCertIDs[0]
|
||||
} else {
|
||||
route.CertID = nil
|
||||
}
|
||||
route.DomainCertIDs = route.CertIDs
|
||||
}
|
||||
|
||||
func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) {
|
||||
@@ -277,7 +306,7 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
|
||||
if route == nil {
|
||||
return nil, errors.New("proxy route is nil")
|
||||
}
|
||||
domains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
|
||||
domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -293,26 +322,17 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
zoneDomainIDs := make([]uint, 0, len(domains))
|
||||
zoneDomains := make([]ZoneDomainView, 0, len(domains))
|
||||
for _, domain := range domains {
|
||||
zoneDomainIDs = append(zoneDomainIDs, domain.ID)
|
||||
zoneDomains = append(zoneDomains, ZoneDomainView{ID: domain.ID, ZoneID: domain.ZoneID, Domain: domain.Domain, CertID: domain.CertID})
|
||||
}
|
||||
domainCertIDs, err := resolveProxyRouteDomainCertIDs(ctx, route, domains, certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var certID *uint
|
||||
if len(certIDs) > 0 {
|
||||
certID = &certIDs[0]
|
||||
}
|
||||
primaryDomain := domains[0]
|
||||
return &View{
|
||||
ID: route.ID,
|
||||
SiteName: routeidentity.ResolveSiteName(route, route.SiteName, primaryDomain),
|
||||
Domain: primaryDomain,
|
||||
Domains: domains,
|
||||
PrimaryDomain: primaryDomain,
|
||||
DomainCount: len(domains),
|
||||
SiteName: route.SiteName,
|
||||
ZoneDomainIDs: zoneDomainIDs,
|
||||
ZoneDomains: zoneDomains,
|
||||
OriginID: route.OriginID,
|
||||
OriginURL: route.OriginURL,
|
||||
OriginHost: route.OriginHost,
|
||||
@@ -320,9 +340,6 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
|
||||
UpstreamList: upstreams,
|
||||
Enabled: route.Enabled,
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: certID,
|
||||
CertIDs: certIDs,
|
||||
DomainCertIDs: domainCertIDs,
|
||||
RedirectHTTP: route.RedirectHTTP,
|
||||
LimitConnPerServer: route.LimitConnPerServer,
|
||||
LimitConnPerIP: route.LimitConnPerIP,
|
||||
|
||||
@@ -17,130 +17,66 @@ import (
|
||||
|
||||
func setupProxyRouteTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{}))
|
||||
|
||||
require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{}, &model.Zone{}, &model.ZoneDomain{}, &model.TLSCertificate{}))
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
return func() { db.SetDB(nil) }
|
||||
}
|
||||
|
||||
func createZoneDomain(t *testing.T, ctx context.Context, domain string, certID *uint) *model.ZoneDomain {
|
||||
t.Helper()
|
||||
zone := &model.Zone{Domain: "example.com"}
|
||||
var existing model.Zone
|
||||
if err := db.DB(ctx).Where("domain = ?", zone.Domain).First(&existing).Error; err == nil {
|
||||
zone = &existing
|
||||
} else {
|
||||
require.NoError(t, db.DB(ctx).Create(zone).Error)
|
||||
}
|
||||
item := &model.ZoneDomain{ZoneID: zone.ID, Domain: domain, CertID: certID}
|
||||
require.NoError(t, db.DB(ctx).Create(item).Error)
|
||||
return item
|
||||
}
|
||||
|
||||
func TestCreateProxyRoute(t *testing.T) {
|
||||
func TestCreateProxyRouteBindsZoneDomains(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
domainA := createZoneDomain(t, ctx, "api.example.com", nil)
|
||||
domainB := createZoneDomain(t, ctx, "www.example.com", nil)
|
||||
|
||||
view, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "example-site",
|
||||
Domain: "example.com",
|
||||
OriginURL: "http://origin.example.com:8080",
|
||||
Enabled: true,
|
||||
})
|
||||
view, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domainA.ID, domainB.ID}, OriginURL: "http://origin.example.com:8080", Enabled: true})
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, view.ID)
|
||||
assert.Equal(t, "example-site", view.SiteName)
|
||||
assert.Equal(t, "example.com", view.Domain)
|
||||
assert.Equal(t, []string{"example.com"}, view.Domains)
|
||||
assert.Equal(t, "http://origin.example.com:8080", view.OriginURL)
|
||||
assert.Equal(t, []string{"http://origin.example.com:8080"}, view.UpstreamList)
|
||||
assert.True(t, view.Enabled)
|
||||
assert.Equal(t, []uint{domainA.ID, domainB.ID}, view.ZoneDomainIDs)
|
||||
require.Len(t, view.ZoneDomains, 2)
|
||||
assert.Equal(t, "api.example.com", view.ZoneDomains[0].Domain)
|
||||
}
|
||||
|
||||
_, err = CreateProxyRoute(ctx, Input{
|
||||
SiteName: "duplicate-site",
|
||||
Domain: "example.com",
|
||||
OriginURL: "http://origin.example.com:8080",
|
||||
})
|
||||
func TestCreateProxyRouteRejectsInvalidZoneDomainBindings(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
domain := createZoneDomain(t, ctx, "api.example.com", nil)
|
||||
base := Input{SiteName: "api", OriginURL: "http://origin.example.com:8080"}
|
||||
|
||||
_, err := CreateProxyRoute(ctx, base)
|
||||
require.EqualError(t, err, errProxyRouteZoneDomainsRequired)
|
||||
base.ZoneDomainIDs = []uint{domain.ID, domain.ID}
|
||||
_, err = CreateProxyRoute(ctx, base)
|
||||
require.EqualError(t, err, errProxyRouteZoneDomainDuplicate)
|
||||
|
||||
first, err := CreateProxyRoute(ctx, Input{SiteName: "first", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080"})
|
||||
require.NoError(t, err)
|
||||
_, err = CreateProxyRoute(ctx, Input{SiteName: "second", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://other.example.com:8080"})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
require.NoError(t, DeleteProxyRoute(ctx, first.ID))
|
||||
}
|
||||
|
||||
func TestListProxyRoutes(t *testing.T) {
|
||||
func TestCreateProxyRouteHTTPSRequiresCoveringCertificate(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
first, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "first-site",
|
||||
Domain: "first.example.com",
|
||||
OriginURL: "http://origin-a.internal:80",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
second, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "second-site",
|
||||
Domain: "second.example.com",
|
||||
OriginURL: "http://origin-b.internal:80",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
routes, err := ListProxyRoutes(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, routes, 2)
|
||||
assert.Equal(t, second.ID, routes[0].ID)
|
||||
assert.Equal(t, first.ID, routes[1].ID)
|
||||
assert.Equal(t, "second.example.com", routes[0].Domain)
|
||||
assert.Equal(t, "first.example.com", routes[1].Domain)
|
||||
}
|
||||
|
||||
func TestValidateProxyRouteIdentityUniquenessUsesDecodedPrimaryDomain(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
existing := &model.ProxyRoute{
|
||||
SiteName: "",
|
||||
Domain: "legacy.example.com",
|
||||
Domains: `["primary.example.com"]`,
|
||||
OriginURL: "http://origin.example.com:8080",
|
||||
Upstreams: `["http://origin.example.com:8080"]`,
|
||||
Enabled: true,
|
||||
UpstreamType: "direct",
|
||||
}
|
||||
require.NoError(t, model.CreateProxyRouteRecord(ctx, existing))
|
||||
|
||||
view, err := GetProxyRoute(ctx, existing.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "primary.example.com", view.SiteName)
|
||||
|
||||
_, err = CreateProxyRoute(ctx, Input{
|
||||
SiteName: "primary.example.com",
|
||||
Domain: "other.example.com",
|
||||
OriginURL: "http://origin-b.example.com:8080",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "site_name already exists")
|
||||
}
|
||||
|
||||
func TestUpdateProxyRouteAuthConfig(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "auth-site",
|
||||
Domain: "auth.example.com",
|
||||
OriginURL: "http://origin.example.com:8080",
|
||||
Enabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
updated, err := UpdateProxyRoute(ctx, created.ID, Input{
|
||||
SiteName: created.SiteName,
|
||||
Domain: created.Domain,
|
||||
Domains: created.Domains,
|
||||
OriginURL: created.OriginURL,
|
||||
Enabled: created.Enabled,
|
||||
BasicAuthEnabled: true,
|
||||
BasicAuthUsername: "admin",
|
||||
BasicAuthPassword: "secret",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, updated.BasicAuthEnabled)
|
||||
assert.Equal(t, "admin", updated.BasicAuthUsername)
|
||||
assert.Equal(t, "secret", updated.BasicAuthPassword)
|
||||
domain := createZoneDomain(t, ctx, "api.example.com", nil)
|
||||
_, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080", EnableHTTPS: true})
|
||||
require.EqualError(t, err, errProxyRouteCertRequired)
|
||||
}
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
@@ -138,4 +137,4 @@ func DeleteProxyRouteHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,10 +5,8 @@ package tls
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"strings"
|
||||
@@ -45,18 +43,3 @@ func isUniqueConstraintError(err error) bool {
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return nil, nil
|
||||
}
|
||||
var domainCertIDs []uint
|
||||
if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if domainCount > 0 && len(domainCertIDs) != domainCount {
|
||||
return nil, fmt.Errorf("domain_cert_ids length mismatch")
|
||||
}
|
||||
return domainCertIDs, nil
|
||||
}
|
||||
|
||||
@@ -430,35 +430,12 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
|
||||
}
|
||||
|
||||
func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
|
||||
routes, err := model.ListTLSProxyRouteRefs(ctx)
|
||||
count, err := model.CountZoneDomainsByCertificateID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, route := range routes {
|
||||
if route.CertID != nil && *route.CertID == id {
|
||||
return errors.New(errCertificateDeleteReferenced)
|
||||
}
|
||||
if strings.TrimSpace(route.CertIDs) == "" {
|
||||
continue
|
||||
}
|
||||
var certIDs []uint
|
||||
if err := json.Unmarshal([]byte(route.CertIDs), &certIDs); err != nil {
|
||||
return fmt.Errorf("proxy route %d cert_ids payload is invalid: %w", route.ID, err)
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if certID == id {
|
||||
return errors.New(errCertificateDeleteReferenced)
|
||||
}
|
||||
}
|
||||
domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, 0)
|
||||
if err != nil {
|
||||
return fmt.Errorf("proxy route %d domain_cert_ids payload is invalid: %w", route.ID, err)
|
||||
}
|
||||
for _, certID := range domainCertIDs {
|
||||
if certID == id {
|
||||
return errors.New(errCertificateDeleteReferenced)
|
||||
}
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(errCertificateDeleteReferenced)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -40,6 +40,8 @@ func setupTLSTestDB(t *testing.T) func() {
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.TLSCertificate{},
|
||||
&model.Zone{},
|
||||
&model.ZoneDomain{},
|
||||
&model.ManagedDomain{},
|
||||
&model.DNSAccount{},
|
||||
&model.AcmeAccount{},
|
||||
@@ -71,6 +73,22 @@ func setupTLSTestDB(t *testing.T) func() {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteCertificateRejectsZoneDomainReference(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
certPEM, keyPEM := generateTestCertificatePair(t, []string{"api.example.com"})
|
||||
certificate, err := CreateCertificate(ctx, CertificateInput{Name: "api-cert", CertPEM: certPEM, KeyPEM: keyPEM})
|
||||
require.NoError(t, err)
|
||||
zone := &model.Zone{Domain: "example.com"}
|
||||
require.NoError(t, db.DB(ctx).Create(zone).Error)
|
||||
require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, Domain: "api.example.com", CertID: &certificate.ID}).Error)
|
||||
|
||||
err = DeleteCertificate(ctx, certificate.ID)
|
||||
require.EqualError(t, err, errCertificateDeleteReferenced)
|
||||
}
|
||||
|
||||
func generateTestCertificatePair(t *testing.T, dnsNames []string) (string, string) {
|
||||
t.Helper()
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user