mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 07:36:37 +08:00
fix: 收敛子代理站点标识双轨逻辑
This commit is contained in:
@@ -17,6 +17,7 @@ import (
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"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"
|
||||
@@ -315,45 +316,30 @@ func loadTLSCertificates(ctx context.Context, certIDs []uint) ([]*tlsCertificate
|
||||
return certificates, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteSiteNameInput(route *model.ProxyRoute, raw, 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 normalizeProxyRouteDomains(rawDomains []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(rawDomains))
|
||||
for _, rawDomain := range rawDomains {
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, errors.New(errProxyRouteDomainInvalid)
|
||||
}
|
||||
normalized = append(normalized, domain)
|
||||
func mapRouteIdentityDomainError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
normalized = uniqueStrings(normalized)
|
||||
if len(normalized) == 0 {
|
||||
return nil, errors.New(errProxyRouteDomainRequired)
|
||||
switch err.Error() {
|
||||
case "domain is required":
|
||||
return errors.New(errProxyRouteDomainRequired)
|
||||
default:
|
||||
if strings.Contains(err.Error(), " is invalid") {
|
||||
return errors.New(errProxyRouteDomainInvalid)
|
||||
}
|
||||
return err
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string, rawDomains []string) ([]string, error) {
|
||||
if len(rawDomains) > 0 {
|
||||
domains, err := normalizeProxyRouteDomains(rawDomains)
|
||||
domains, err := routeidentity.NormalizeDomains(rawDomains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, mapRouteIdentityDomainError(err)
|
||||
}
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain != "" && domain != domains[0] {
|
||||
@@ -363,7 +349,7 @@ func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string,
|
||||
}
|
||||
|
||||
if route != nil {
|
||||
existingDomains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
existingDomains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
|
||||
if err == nil && len(existingDomains) > 0 {
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain == "" || domain == existingDomains[0] {
|
||||
@@ -372,7 +358,11 @@ func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string,
|
||||
}
|
||||
}
|
||||
|
||||
return normalizeProxyRouteDomains([]string{rawDomain})
|
||||
domains, err := routeidentity.NormalizeDomains([]string{rawDomain})
|
||||
if err != nil {
|
||||
return nil, mapRouteIdentityDomainError(err)
|
||||
}
|
||||
return domains, nil
|
||||
}
|
||||
|
||||
func validateProxyRouteSiteName(siteName string) error {
|
||||
@@ -397,15 +387,13 @@ func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.Prox
|
||||
if item == nil || item.ID == currentID {
|
||||
continue
|
||||
}
|
||||
existingSiteName := normalizeProxyRouteSiteNameInput(item, item.SiteName, item.Domain)
|
||||
if existingSiteName == siteName {
|
||||
return errors.New(errProxyRouteSiteNameExists)
|
||||
}
|
||||
|
||||
existingDomains, err := decodeStoredDomains(item.Domains, item.Domain)
|
||||
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 {
|
||||
return errors.New(errProxyRouteSiteNameExists)
|
||||
}
|
||||
existingSet := make(map[string]struct{}, len(existingDomains))
|
||||
for _, existingDomain := range existingDomains {
|
||||
existingSet[existingDomain] = struct{}{}
|
||||
@@ -789,18 +777,6 @@ 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("domains payload is invalid")
|
||||
}
|
||||
return normalizeProxyRouteDomains(domains)
|
||||
}
|
||||
|
||||
func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
@@ -164,7 +165,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
return nil, err
|
||||
}
|
||||
domain := domains[0]
|
||||
siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain)
|
||||
siteName := routeidentity.ResolveSiteName(route, input.SiteName, domain)
|
||||
|
||||
upstreamType := normalizeUpstreamType(input.UpstreamType)
|
||||
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
|
||||
@@ -276,7 +277,7 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
|
||||
if route == nil {
|
||||
return nil, errors.New("proxy route is nil")
|
||||
}
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
domains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -307,7 +308,7 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e
|
||||
primaryDomain := domains[0]
|
||||
return &View{
|
||||
ID: route.ID,
|
||||
SiteName: normalizeProxyRouteSiteNameInput(route, route.SiteName, primaryDomain),
|
||||
SiteName: routeidentity.ResolveSiteName(route, route.SiteName, primaryDomain),
|
||||
Domain: primaryDomain,
|
||||
Domains: domains,
|
||||
PrimaryDomain: primaryDomain,
|
||||
|
||||
@@ -86,3 +86,32 @@ func TestListProxyRoutes(t *testing.T) {
|
||||
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")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user