mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 08:36:37 +08:00
refactor(proxy): bind routes through zone domains
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user