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
+33 -242
View File
@@ -17,7 +17,6 @@ import (
"strings"
"unicode"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
@@ -254,63 +253,23 @@ func validateCertificateCoverage(certificate *model.TLSCertificate, domains []st
return nil
}
func loadTLSCertificates(ctx context.Context, certIDs []uint) ([]*model.TLSCertificate, error) {
certificates := make([]*model.TLSCertificate, 0, len(certIDs))
for _, certID := range certIDs {
certificate, err := lookupTLSCertificateByID(ctx, certID)
if err != nil {
return nil, err
}
certificates = append(certificates, certificate)
func loadProxyRouteZoneDomains(ctx context.Context, ids []uint) ([]model.ZoneDomain, error) {
if len(ids) == 0 {
return nil, errors.New(errProxyRouteZoneDomainsRequired)
}
return certificates, nil
}
func normalizeProxyRouteDomainValue(raw string) string {
return strings.ToLower(strings.TrimSpace(raw))
}
func mapRouteIdentityDomainError(err error) error {
if err == nil {
return nil
}
switch err.Error() {
case "domain is required":
return errors.New(errProxyRouteDomainRequired)
default:
if strings.Contains(err.Error(), " is invalid") {
return errors.New(errProxyRouteDomainInvalid)
seen := make(map[uint]struct{}, len(ids))
for _, id := range ids {
if id == 0 {
return nil, errors.New(errProxyRouteZoneDomainNotFound)
}
return err
}
}
func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string, rawDomains []string) ([]string, error) {
if len(rawDomains) > 0 {
domains, err := routeidentity.NormalizeDomains(rawDomains)
if err != nil {
return nil, mapRouteIdentityDomainError(err)
if _, ok := seen[id]; ok {
return nil, errors.New(errProxyRouteZoneDomainDuplicate)
}
domain := normalizeProxyRouteDomainValue(rawDomain)
if domain != "" && domain != domains[0] {
return nil, errors.New(errProxyRouteDomainMismatch)
}
return domains, nil
seen[id] = struct{}{}
}
if route != nil {
existingDomains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
if err == nil && len(existingDomains) > 0 {
domain := normalizeProxyRouteDomainValue(rawDomain)
if domain == "" || domain == existingDomains[0] {
return existingDomains, nil
}
}
}
domains, err := routeidentity.NormalizeDomains([]string{rawDomain})
domains, err := model.ListZoneDomainsByIDs(ctx, ids)
if err != nil {
return nil, mapRouteIdentityDomainError(err)
return nil, errors.New(errProxyRouteZoneDomainNotFound)
}
return domains, nil
}
@@ -322,7 +281,7 @@ func validateProxyRouteSiteName(siteName string) error {
return nil
}
func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string, domains []string) error {
func validateProxyRouteSiteNameUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string) error {
routes, err := model.ListProxyRoutes(ctx)
if err != nil {
return err
@@ -337,27 +296,33 @@ func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.Prox
if item == nil || item.ID == currentID {
continue
}
existingSiteName, existingDomains, err := routeidentity.ResolveFromRoute(item)
if err != nil {
return fmt.Errorf("existing route %d domains are invalid: %w", item.ID, err)
}
if existingSiteName == siteName {
if item.SiteName == siteName {
return errors.New(errProxyRouteSiteNameExists)
}
existingSet := make(map[string]struct{}, len(existingDomains))
for _, existingDomain := range existingDomains {
existingSet[existingDomain] = struct{}{}
}
for _, domain := range domains {
if _, ok := existingSet[domain]; ok {
return fmt.Errorf(errProxyRouteDomainExists, domain)
}
}
}
return nil
}
func validateProxyRouteZoneDomainCertificates(ctx context.Context, domains []model.ZoneDomain, enableHTTPS bool) error {
if !enableHTTPS {
return nil
}
for _, domain := range domains {
if domain.CertID == nil || *domain.CertID == 0 {
return errors.New(errProxyRouteCertRequired)
}
certificate, err := lookupTLSCertificateByID(ctx, *domain.CertID)
if err != nil {
return errors.New(errProxyRouteCertNotFound)
}
if err := validateCertificateCoverage(certificate, []string{domain.Domain}); err != nil {
return err
}
}
return nil
}
func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) {
if value < 0 {
return 0, fmt.Errorf("%s must be greater than or equal to 0", field)
@@ -365,150 +330,6 @@ func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) {
return value, nil
}
func normalizeProxyRouteCertificateIDs(ctx context.Context, enableHTTPS bool, certID *uint, certIDs []uint) ([]uint, error) {
if !enableHTTPS {
return []uint{}, nil
}
candidates := make([]uint, 0, len(certIDs)+1)
if certID != nil && *certID != 0 {
candidates = append(candidates, *certID)
}
candidates = append(candidates, certIDs...)
normalized := make([]uint, 0, len(candidates))
seen := make(map[uint]struct{}, len(candidates))
for _, item := range candidates {
if item == 0 {
continue
}
if _, ok := seen[item]; ok {
continue
}
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
return nil, errors.New(errProxyRouteCertNotFound)
}
seen[item] = struct{}{}
normalized = append(normalized, item)
}
if len(normalized) == 0 {
return nil, errors.New(errProxyRouteCertRequired)
}
return normalized, nil
}
func normalizeProxyRouteDomainCertificateIDs(
ctx context.Context,
domains []string,
enableHTTPS bool,
rawDomainCertIDs []uint,
certID *uint,
certIDs []uint,
) ([]uint, []uint, *uint, error) {
if !enableHTTPS {
return []uint{}, []uint{}, nil, nil
}
if len(rawDomainCertIDs) > 0 {
return normalizeExplicitDomainCertIDs(ctx, domains, rawDomainCertIDs)
}
normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(ctx, enableHTTPS, certID, certIDs)
if err != nil {
return nil, nil, nil, err
}
return normalizeDerivedDomainCertIDs(ctx, domains, normalizedCertIDs)
}
func validateProxyRouteDomainCertificateCoverage(ctx context.Context, domains []string, domainCertIDs []uint) error {
if len(domainCertIDs) == 0 {
return nil
}
domainsByCertID := make(map[uint][]string)
for index, certID := range domainCertIDs {
if certID == 0 {
continue
}
domainsByCertID[certID] = append(domainsByCertID[certID], domains[index])
}
for certID, assignedDomains := range domainsByCertID {
certificate, err := lookupTLSCertificateByID(ctx, certID)
if err != nil {
return errors.New(errProxyRouteCertNotFound)
}
if err := validateCertificateCoverage(certificate, assignedDomains); err != nil {
return err
}
}
return nil
}
func deriveDomainCertIDsFromCertificateSet(ctx context.Context, domains []string, certIDs []uint) ([]uint, error) {
certificates, err := loadTLSCertificates(ctx, certIDs)
if err != nil {
return nil, err
}
result := make([]uint, len(domains))
for domainIndex, domain := range domains {
if domainIndex < len(certificates) &&
certificates[domainIndex] != nil &&
validateCertificateCoverage(certificates[domainIndex], []string{domain}) == nil {
result[domainIndex] = certificates[domainIndex].ID
continue
}
assigned := uint(0)
for _, certificate := range certificates {
if certificate != nil &&
validateCertificateCoverage(certificate, []string{domain}) == nil {
assigned = certificate.ID
break
}
}
if assigned == 0 {
return nil, fmt.Errorf("certificate does not cover domain %s", domain)
}
result[domainIndex] = assigned
}
return result, nil
}
func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []uint{}, nil
}
var domainCertIDs []uint
if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil {
return nil, errors.New("domain_cert_ids payload is invalid")
}
if len(domainCertIDs) == 0 {
return []uint{}, nil
}
if domainCount > 0 && len(domainCertIDs) != domainCount {
return nil, errors.New("domain_cert_ids length does not match domains")
}
normalized := make([]uint, len(domainCertIDs))
copy(normalized, domainCertIDs)
return normalized, nil
}
func resolveProxyRouteDomainCertIDs(ctx context.Context, route *model.ProxyRoute, domains []string, certIDs []uint) ([]uint, error) {
domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, len(domains))
if err != nil {
return nil, err
}
if len(domainCertIDs) > 0 || len(certIDs) == 0 {
return domainCertIDs, nil
}
return deriveDomainCertIDsFromCertificateSet(ctx, domains, certIDs)
}
func normalizeProxyRouteLimitRate(raw string) (string, error) {
normalized := strings.ToLower(strings.TrimSpace(raw))
if normalized == "" || normalized == "0" {
@@ -727,36 +548,6 @@ func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, erro
return normalizeUpstreams(fallbackOriginURL, upstreams)
}
func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) {
text := strings.TrimSpace(raw)
if text == "" {
if fallbackCertID == nil || *fallbackCertID == 0 {
return []uint{}, nil
}
return []uint{*fallbackCertID}, nil
}
var certIDs []uint
if err := json.Unmarshal([]byte(text), &certIDs); err != nil {
return nil, errors.New("cert_ids payload is invalid")
}
normalized := make([]uint, 0, len(certIDs))
seen := make(map[uint]struct{}, len(certIDs))
for _, certID := range certIDs {
if certID == 0 {
continue
}
if _, ok := seen[certID]; ok {
continue
}
seen[certID] = struct{}{}
normalized = append(normalized, certID)
}
if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 {
return []uint{*fallbackCertID}, nil
}
return normalized, nil
}
func validateOriginURL(raw string) error {
if raw == "" {
return errors.New(errProxyRouteOriginEmpty)