mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-10 17:26:38 +08:00
[功能] 添加站点名称和多域名支持到代理路由,更新相关逻辑和测试
This commit is contained in:
@@ -58,7 +58,9 @@ type ConfigOptionDiffItem struct {
|
||||
}
|
||||
|
||||
type snapshotRoute struct {
|
||||
SiteName string `json:"site_name,omitempty"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains,omitempty"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginHost string `json:"origin_host,omitempty"`
|
||||
Upstreams []string `json:"upstreams,omitempty"`
|
||||
@@ -224,7 +226,7 @@ func DiffConfigVersion() (*ConfigDiffResult, error) {
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
for _, route := range bundle.SnapshotRoutes {
|
||||
result.AddedDomains = append(result.AddedDomains, route.Domain)
|
||||
result.AddedDomains = append(result.AddedDomains, route.Domains...)
|
||||
}
|
||||
result.MainConfigChanged = true
|
||||
result.ChangedOptionKeys = openRestyOptionKeys()
|
||||
@@ -238,14 +240,8 @@ func DiffConfigVersion() (*ConfigDiffResult, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
currentMap := make(map[string]snapshotRoute, len(bundle.SnapshotRoutes))
|
||||
for _, route := range bundle.SnapshotRoutes {
|
||||
currentMap[route.Domain] = route
|
||||
}
|
||||
activeMap := make(map[string]snapshotRoute, len(activeSnapshot.Routes))
|
||||
for _, route := range activeSnapshot.Routes {
|
||||
activeMap[route.Domain] = route
|
||||
}
|
||||
currentMap := flattenSnapshotRoutesByDomain(bundle.SnapshotRoutes)
|
||||
activeMap := flattenSnapshotRoutesByDomain(activeSnapshot.Routes)
|
||||
for domain, currentRoute := range currentMap {
|
||||
activeRoute, ok := activeMap[domain]
|
||||
if !ok {
|
||||
@@ -403,6 +399,10 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) {
|
||||
func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
items := make([]snapshotRoute, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("route %s domains are invalid", route.Domain)
|
||||
}
|
||||
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
@@ -416,7 +416,9 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
return nil, fmt.Errorf("路由 %s 缓存规则无效", route.Domain)
|
||||
}
|
||||
items = append(items, snapshotRoute{
|
||||
Domain: route.Domain,
|
||||
SiteName: normalizeProxyRouteSiteNameInput(route, route.SiteName, domains[0]),
|
||||
Domain: domains[0],
|
||||
Domains: domains,
|
||||
OriginURL: route.OriginURL,
|
||||
OriginHost: route.OriginHost,
|
||||
Upstreams: upstreams,
|
||||
@@ -459,6 +461,19 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
|
||||
return []snapshotRoute{}
|
||||
}
|
||||
for index := range routes {
|
||||
normalizedDomains, err := decodeStoredDomains("", routes[index].Domain)
|
||||
if len(routes[index].Domains) > 0 {
|
||||
normalizedDomains, err = normalizeProxyRouteDomains(routes[index].Domains)
|
||||
}
|
||||
if err == nil && len(normalizedDomains) > 0 {
|
||||
routes[index].Domains = normalizedDomains
|
||||
routes[index].Domain = normalizedDomains[0]
|
||||
routes[index].SiteName = normalizeProxyRouteSiteNameInput(
|
||||
&model.ProxyRoute{SiteName: routes[index].SiteName},
|
||||
routes[index].SiteName,
|
||||
normalizedDomains[0],
|
||||
)
|
||||
}
|
||||
normalizedHeaders, err := normalizeCustomHeaders(routes[index].CustomHeaders)
|
||||
if err == nil {
|
||||
routes[index].CustomHeaders = normalizedHeaders
|
||||
@@ -477,10 +492,30 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
|
||||
return routes
|
||||
}
|
||||
|
||||
func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRoute {
|
||||
domainMap := make(map[string]snapshotRoute)
|
||||
for _, route := range normalizeSnapshotRoutes(routes) {
|
||||
for _, domain := range route.Domains {
|
||||
item := route
|
||||
item.Domain = domain
|
||||
domainMap[domain] = item
|
||||
}
|
||||
}
|
||||
return domainMap
|
||||
}
|
||||
|
||||
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
if left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || !uintPointerEqual(left.CertID, right.CertID) {
|
||||
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || !uintPointerEqual(left.CertID, right.CertID) {
|
||||
return false
|
||||
}
|
||||
if len(left.Domains) != len(right.Domains) {
|
||||
return false
|
||||
}
|
||||
for index := range left.Domains {
|
||||
if left.Domains[index] != right.Domains[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.Upstreams) != len(right.Upstreams) {
|
||||
return false
|
||||
}
|
||||
@@ -658,6 +693,15 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
builder.WriteString("# This file is generated by OpenFlare. Do not edit manually.\n")
|
||||
supportFiles := make([]SupportFile, 0)
|
||||
for _, route := range routes {
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("route %s domains are invalid", route.Domain)
|
||||
}
|
||||
serverNames := renderServerNames(domains)
|
||||
displayName := route.SiteName
|
||||
if strings.TrimSpace(displayName) == "" {
|
||||
displayName = domains[0]
|
||||
}
|
||||
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
@@ -680,7 +724,7 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPProxyServer(route.Domain, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, upstreamConfig, cfg))
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, upstreamConfig, cfg))
|
||||
continue
|
||||
}
|
||||
if route.CertID == nil || *route.CertID == 0 {
|
||||
@@ -690,16 +734,19 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 关联证书不存在", route.Domain)
|
||||
}
|
||||
if err := validateCertificateCoverage(certificate, domains); err != nil {
|
||||
return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
}
|
||||
supportFiles = append(supportFiles,
|
||||
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
|
||||
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)},
|
||||
)
|
||||
if route.RedirectHTTP {
|
||||
builder.WriteString(renderHTTPRedirectServer(route.Domain))
|
||||
builder.WriteString(renderHTTPRedirectServer(serverNames))
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(route.Domain, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, upstreamConfig, cfg))
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, upstreamConfig, cfg))
|
||||
}
|
||||
builder.WriteString(renderHTTPSServer(route.Domain, route.OriginURL, route.OriginHost, certificate.ID, customHeaders, cacheConfig, upstreamConfig, cfg))
|
||||
builder.WriteString(renderHTTPSServer(serverNames, route.OriginURL, route.OriginHost, certificate.ID, customHeaders, cacheConfig, upstreamConfig, cfg))
|
||||
}
|
||||
return builder.String(), dedupeSupportFiles(supportFiles), nil
|
||||
}
|
||||
@@ -836,18 +883,38 @@ func nextVersionNumber(now time.Time) (string, error) {
|
||||
return fmt.Sprintf("%s-%03d", prefix, count+1), nil
|
||||
}
|
||||
|
||||
func renderHTTPProxyServer(domain string, originURL string, originHost string, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, upstreamConfig routeUpstreamConfig, cfg openRestyConfigSnapshot) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n location / {\n%s%s%s }\n}\n\n", domain, renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig))
|
||||
func renderHTTPProxyServer(serverNames string, originURL string, originHost string, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, upstreamConfig routeUpstreamConfig, cfg openRestyConfigSnapshot) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n location / {\n%s%s%s }\n}\n\n", serverNames, renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig))
|
||||
}
|
||||
|
||||
func renderHTTPRedirectServer(domain string) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", domain)
|
||||
func renderHTTPRedirectServer(serverNames string) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", serverNames)
|
||||
}
|
||||
|
||||
func renderHTTPSServer(domain string, originURL string, originHost string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, upstreamConfig routeUpstreamConfig, cfg openRestyConfigSnapshot) string {
|
||||
func renderHTTPSServer(serverNames string, originURL string, originHost string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, upstreamConfig routeUpstreamConfig, cfg openRestyConfigSnapshot) string {
|
||||
certPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateCertFileName(certificateID))
|
||||
keyPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateKeyFileName(certificateID))
|
||||
return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n\n location / {\n%s%s%s }\n}\n\n", domain, certPath, keyPath, renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig))
|
||||
return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n\n location / {\n%s%s%s }\n}\n\n", serverNames, certPath, keyPath, renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig))
|
||||
}
|
||||
|
||||
func renderServerNames(domains []string) string {
|
||||
return strings.Join(domains, " ")
|
||||
}
|
||||
|
||||
func validateCertificateCoverage(certificate *model.TLSCertificate, domains []string) error {
|
||||
if certificate == nil {
|
||||
return errors.New("certificate is nil")
|
||||
}
|
||||
leaf, err := parseLeafCertificate(certificate.CertPEM)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, domain := range domains {
|
||||
if err := leaf.VerifyHostname(domain); err != nil {
|
||||
return fmt.Errorf("certificate does not cover domain %s", domain)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func renderConnectionUpgradeMap() string {
|
||||
@@ -1022,6 +1089,10 @@ func buildUpstreamProxyPassURI(parsed *url.URL) string {
|
||||
}
|
||||
|
||||
func buildRouteUpstreamName(route *model.ProxyRoute) string {
|
||||
identity := strings.TrimSpace(route.SiteName)
|
||||
if identity == "" {
|
||||
identity = route.Domain
|
||||
}
|
||||
sanitized := strings.Map(func(r rune) rune {
|
||||
switch {
|
||||
case r >= 'a' && r <= 'z':
|
||||
@@ -1033,7 +1104,7 @@ func buildRouteUpstreamName(route *model.ProxyRoute) string {
|
||||
default:
|
||||
return '_'
|
||||
}
|
||||
}, route.Domain)
|
||||
}, identity)
|
||||
sanitized = strings.Trim(sanitized, "_")
|
||||
if sanitized == "" {
|
||||
sanitized = "backend"
|
||||
|
||||
@@ -115,6 +115,29 @@ func TestCreateProxyRouteRejectsHTTPSWithoutCertificate(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProxyRouteSupportsWebsiteDomains(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "main-site",
|
||||
Domains: []string{"app.example.com", "www.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if route.SiteName != "main-site" {
|
||||
t.Fatalf("unexpected site name: %s", route.SiteName)
|
||||
}
|
||||
if route.Domain != "app.example.com" {
|
||||
t.Fatalf("expected primary domain mirror, got %s", route.Domain)
|
||||
}
|
||||
if !strings.Contains(route.Domains, "www.example.com") {
|
||||
t.Fatalf("expected domains payload to contain alias, got %s", route.Domains)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionRendersCustomHeaders(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
if err := model.UpdateOption("OpenRestyWebsocketEnabled", "true"); err != nil {
|
||||
@@ -300,6 +323,94 @@ func TestPublishConfigVersionRendersMultipleUpstreams(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionRendersMultiDomainWebsite(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
certPEM, keyPEM := generateCertificatePair(t, []string{"app.example.com", "www.example.com"})
|
||||
certificate, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "multi-domain",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
_, err = CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "marketing-site",
|
||||
Domains: []string{"app.example.com", "www.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
EnableHTTPS: true,
|
||||
CertID: &certificate.ID,
|
||||
RedirectHTTP: true,
|
||||
CacheEnabled: true,
|
||||
CachePolicy: proxyRouteCachePolicyPathPrefix,
|
||||
CacheRules: []string{"/assets"},
|
||||
CustomHeaders: []ProxyRouteCustomHeaderInput{{Key: "X-Site", Value: "marketing"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root")
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "server_name app.example.com www.example.com;") {
|
||||
t.Fatal("expected rendered config to include all domains in one server_name")
|
||||
}
|
||||
if strings.Contains(result.Version.RenderedConfig, "server_name app.example.com;") {
|
||||
t.Fatal("expected rendered config to avoid standalone primary-domain server block")
|
||||
}
|
||||
if strings.Contains(result.Version.RenderedConfig, "server_name www.example.com;") {
|
||||
t.Fatal("expected rendered config to avoid standalone alias server block")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"site_name":"marketing-site"`) {
|
||||
t.Fatal("expected snapshot to include site_name")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"domains":["app.example.com","www.example.com"]`) {
|
||||
t.Fatal("expected snapshot to include domain list")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffConfigVersionTracksAddedDomainWithinWebsite(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "main-site",
|
||||
Domains: []string{"app.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if _, err := PublishConfigVersion("root"); err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
|
||||
if _, err := UpdateProxyRoute(route.ID, ProxyRouteInput{
|
||||
SiteName: "main-site",
|
||||
Domains: []string{"app.example.com", "www.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
diff, err := DiffConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("DiffConfigVersion failed: %v", err)
|
||||
}
|
||||
if len(diff.AddedDomains) != 1 || diff.AddedDomains[0] != "www.example.com" {
|
||||
t.Fatalf("unexpected added domains: %#v", diff.AddedDomains)
|
||||
}
|
||||
if len(diff.ModifiedDomains) != 1 || diff.ModifiedDomains[0] != "app.example.com" {
|
||||
t.Fatalf("unexpected modified domains: %#v", diff.ModifiedDomains)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionRendersHostnameLoadBalancingUpstream(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"openflare/model"
|
||||
"regexp"
|
||||
@@ -26,7 +27,9 @@ type ProxyRouteCustomHeaderInput struct {
|
||||
}
|
||||
|
||||
type ProxyRouteInput struct {
|
||||
SiteName string `json:"site_name"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains"`
|
||||
OriginID *uint `json:"origin_id"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginScheme string `json:"origin_scheme"`
|
||||
@@ -91,7 +94,13 @@ func DeleteProxyRoute(id uint) error {
|
||||
}
|
||||
|
||||
func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.ProxyRoute, error) {
|
||||
domain := strings.ToLower(strings.TrimSpace(input.Domain))
|
||||
domains, err := normalizeProxyRouteDomainsInput(route, input.Domain, input.Domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domain := domains[0]
|
||||
siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain)
|
||||
|
||||
originURL, originID, err := resolveProxyRoutePrimaryOrigin(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -123,11 +132,16 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if domain == "" {
|
||||
return nil, errors.New("域名不能为空")
|
||||
domainsJSON, err := json.Marshal(domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, errors.New("域名格式不合法")
|
||||
|
||||
if err := validateProxyRouteSiteName(siteName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteIdentityUniqueness(route, siteName, domains); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateOriginHost(originHost); err != nil {
|
||||
return nil, err
|
||||
@@ -147,10 +161,13 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
if input.RedirectHTTP && !input.EnableHTTPS {
|
||||
return nil, errors.New("仅启用 HTTPS 后才能开启 HTTP 重定向")
|
||||
}
|
||||
|
||||
if route == nil {
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = string(domainsJSON)
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
@@ -167,6 +184,115 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteSiteNameInput(route *model.ProxyRoute, raw string, primaryDomain string) string {
|
||||
siteName := strings.TrimSpace(raw)
|
||||
if siteName != "" {
|
||||
return siteName
|
||||
}
|
||||
if route != nil && strings.TrimSpace(route.SiteName) != "" {
|
||||
return strings.TrimSpace(route.SiteName)
|
||||
}
|
||||
return primaryDomain
|
||||
}
|
||||
|
||||
func normalizeProxyRouteDomainValue(raw string) string {
|
||||
return strings.ToLower(strings.TrimSpace(raw))
|
||||
}
|
||||
|
||||
func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string, rawDomains []string) ([]string, error) {
|
||||
if len(rawDomains) > 0 {
|
||||
domains, err := normalizeProxyRouteDomains(rawDomains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain != "" && domain != domains[0] {
|
||||
return nil, errors.New("domain must match domains[0]")
|
||||
}
|
||||
return domains, nil
|
||||
}
|
||||
|
||||
if route != nil {
|
||||
existingDomains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if err == nil && len(existingDomains) > 0 {
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain == "" || domain == existingDomains[0] {
|
||||
return existingDomains, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return normalizeProxyRouteDomains([]string{rawDomain})
|
||||
}
|
||||
|
||||
func normalizeProxyRouteDomains(rawDomains []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(rawDomains))
|
||||
seen := make(map[string]struct{}, len(rawDomains))
|
||||
for _, rawDomain := range rawDomains {
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, errors.New("域名格式不合法")
|
||||
}
|
||||
if _, ok := seen[domain]; ok {
|
||||
continue
|
||||
}
|
||||
seen[domain] = struct{}{}
|
||||
normalized = append(normalized, domain)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, errors.New("至少填写一个域名")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func validateProxyRouteSiteName(siteName string) error {
|
||||
if strings.TrimSpace(siteName) == "" {
|
||||
return errors.New("站点标识不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateProxyRouteIdentityUniqueness(route *model.ProxyRoute, siteName string, domains []string) error {
|
||||
routes, err := model.ListProxyRoutes()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
currentID := uint(0)
|
||||
if route != nil {
|
||||
currentID = route.ID
|
||||
}
|
||||
|
||||
for _, item := range routes {
|
||||
if item == nil || item.ID == currentID {
|
||||
continue
|
||||
}
|
||||
existingSiteName := normalizeProxyRouteSiteNameInput(item, item.SiteName, item.Domain)
|
||||
if existingSiteName == siteName {
|
||||
return errors.New("站点标识已存在")
|
||||
}
|
||||
|
||||
existingDomains, err := decodeStoredDomains(item.Domains, item.Domain)
|
||||
if err != nil {
|
||||
return fmt.Errorf("existing route %d domains are invalid: %w", item.ID, err)
|
||||
}
|
||||
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("域名 %s 已存在", domain)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func resolveProxyRoutePrimaryOrigin(input ProxyRouteInput) (string, *uint, error) {
|
||||
if hasStructuredOriginInput(input) {
|
||||
scheme, err := normalizeOriginScheme(input.OriginScheme)
|
||||
@@ -446,6 +572,18 @@ func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, erro
|
||||
return normalizeUpstreams(fallbackOriginURL, upstreams)
|
||||
}
|
||||
|
||||
func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return normalizeProxyRouteDomains([]string{fallbackDomain})
|
||||
}
|
||||
var domains []string
|
||||
if err := json.Unmarshal([]byte(text), &domains); err != nil {
|
||||
return nil, errors.New("域名配置格式不合法")
|
||||
}
|
||||
return normalizeProxyRouteDomains(domains)
|
||||
}
|
||||
|
||||
func validateOriginURL(raw string) error {
|
||||
if raw == "" {
|
||||
return errors.New("源站地址不能为空")
|
||||
|
||||
Reference in New Issue
Block a user