From e39a8995f6e35591b2539f52ff847b1252ef9d84 Mon Sep 17 00:00:00 2001 From: ryan Date: Mon, 30 Mar 2026 14:11:30 +0800 Subject: [PATCH] =?UTF-8?q?[=E5=8A=9F=E8=83=BD]=20=E6=B7=BB=E5=8A=A0?= =?UTF-8?q?=E7=AB=99=E7=82=B9=E5=90=8D=E7=A7=B0=E5=92=8C=E5=A4=9A=E5=9F=9F?= =?UTF-8?q?=E5=90=8D=E6=94=AF=E6=8C=81=E5=88=B0=E4=BB=A3=E7=90=86=E8=B7=AF?= =?UTF-8?q?=E7=94=B1=EF=BC=8C=E6=9B=B4=E6=96=B0=E7=9B=B8=E5=85=B3=E9=80=BB?= =?UTF-8?q?=E8=BE=91=E5=92=8C=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../model/database_schema_version.go | 2 +- openflare_server/model/main_test.go | 96 +++++++++ openflare_server/model/migrations.go | 204 +++++++++++++++++- openflare_server/model/proxy_route.go | 6 +- openflare_server/service/config_version.go | 115 ++++++++-- openflare_server/service/https_phase1_test.go | 111 ++++++++++ openflare_server/service/proxy_route.go | 148 ++++++++++++- 7 files changed, 649 insertions(+), 33 deletions(-) diff --git a/openflare_server/model/database_schema_version.go b/openflare_server/model/database_schema_version.go index 0eaeef9e..d3d869be 100644 --- a/openflare_server/model/database_schema_version.go +++ b/openflare_server/model/database_schema_version.go @@ -4,7 +4,7 @@ import "time" const ( legacyDatabaseSchemaVersion = 1 - currentDatabaseSchemaVersion = 4 + currentDatabaseSchemaVersion = 5 databaseSchemaVersionRowID = 1 ) diff --git a/openflare_server/model/main_test.go b/openflare_server/model/main_test.go index d1a0360e..e8b9e7ad 100644 --- a/openflare_server/model/main_test.go +++ b/openflare_server/model/main_test.go @@ -10,6 +10,30 @@ import ( "gorm.io/gorm" ) +type legacyProxyRouteV4 struct { + ID uint `gorm:"primaryKey"` + Domain string `gorm:"uniqueIndex;size:255;not null"` + OriginID *uint `gorm:"index"` + OriginURL string `gorm:"size:2048;not null"` + OriginHost string `gorm:"size:255"` + Upstreams string `gorm:"type:text;not null;default:'[]'"` + Enabled bool `gorm:"not null;default:true"` + EnableHTTPS bool `gorm:"column:enable_https;not null;default:false"` + CertID *uint + RedirectHTTP bool `gorm:"not null;default:false"` + CacheEnabled bool `gorm:"not null;default:false"` + CachePolicy string `gorm:"size:32;not null;default:''"` + CacheRules string `gorm:"type:text;not null;default:'[]'"` + CustomHeaders string `gorm:"type:text;not null;default:'[]'"` + Remark string `gorm:"size:255"` + CreatedAt time.Time + UpdatedAt time.Time +} + +func (legacyProxyRouteV4) TableName() string { + return "proxy_routes" +} + func openBareTestSQLiteDB(t *testing.T, name string) *gorm.DB { t.Helper() @@ -464,6 +488,78 @@ func TestMigrateOriginsSchemaBackfillsOrigins(t *testing.T) { } } +func TestEnsureDatabaseSchemaUpToDateBackfillsProxyRouteSiteFields(t *testing.T) { + db := openBareTestSQLiteDB(t, "legacy-proxy-route-sites.db") + if err := registerSharding(db, "sqlite"); err != nil { + t.Fatalf("register sharding: %v", err) + } + if err := autoMigrateSchemaMetadata(db); err != nil { + t.Fatalf("auto migrate schema metadata: %v", err) + } + + for _, item := range registeredModels() { + if _, ok := item.(*ProxyRoute); ok { + continue + } + if err := db.AutoMigrate(item); err != nil { + t.Fatalf("auto migrate supporting table: %v", err) + } + } + if err := db.AutoMigrate(&legacyProxyRouteV4{}); err != nil { + t.Fatalf("auto migrate legacy proxy_routes: %v", err) + } + + now := time.Now().UTC() + if err := db.Create(&legacyProxyRouteV4{ + Domain: "app.example.com", + OriginURL: "https://origin-a.internal:8443", + Upstreams: `["https://origin-a.internal:8443","https://origin-b.internal:8443"]`, + Enabled: true, + EnableHTTPS: false, + RedirectHTTP: false, + CacheEnabled: false, + CachePolicy: "", + CacheRules: `[]`, + CustomHeaders: `[]`, + CreatedAt: now, + UpdatedAt: now, + }).Error; err != nil { + t.Fatalf("seed legacy proxy route: %v", err) + } + if err := saveDatabaseSchemaVersion(db, 4); err != nil { + t.Fatalf("save schema version: %v", err) + } + + previousDB := DB + DB = db + t.Cleanup(func() { + DB = previousDB + }) + + if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil { + t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err) + } + + var route ProxyRoute + if err := db.First(&route).Error; err != nil { + t.Fatalf("query migrated proxy route: %v", err) + } + if route.SiteName != "app.example.com" { + t.Fatalf("unexpected site_name after migration: %s", route.SiteName) + } + if route.Domain != "app.example.com" { + t.Fatalf("unexpected domain mirror after migration: %s", route.Domain) + } + + var domains []string + if err := json.Unmarshal([]byte(route.Domains), &domains); err != nil { + t.Fatalf("decode migrated domains: %v", err) + } + if len(domains) != 1 || domains[0] != "app.example.com" { + t.Fatalf("unexpected migrated domains: %#v", domains) + } +} + func TestRunDatabaseSchemaMigrationDoesNotAdvanceVersionWhenValidationFails(t *testing.T) { db := openBareTestSQLiteDB(t, "failed-validation.db") diff --git a/openflare_server/model/migrations.go b/openflare_server/model/migrations.go index d83c3926..32317010 100644 --- a/openflare_server/model/migrations.go +++ b/openflare_server/model/migrations.go @@ -211,6 +211,179 @@ func validateDatabaseSchemaV4(db *gorm.DB, backend string) error { return nil } +func normalizeProxyRouteDomainForMigration(raw string) string { + return strings.ToLower(strings.TrimSpace(raw)) +} + +func normalizeProxyRouteSiteNameForMigration(raw string, primaryDomain string) string { + siteName := strings.TrimSpace(raw) + if siteName != "" { + return siteName + } + return primaryDomain +} + +func decodeProxyRouteDomainsForMigration(raw string, fallbackDomain string) ([]string, error) { + primaryDomain := normalizeProxyRouteDomainForMigration(fallbackDomain) + text := strings.TrimSpace(raw) + if text == "" { + if primaryDomain == "" { + return nil, fmt.Errorf("proxy route primary domain is empty") + } + return []string{primaryDomain}, nil + } + + var domains []string + if err := json.Unmarshal([]byte(text), &domains); err != nil { + return nil, fmt.Errorf("decode proxy route domains failed: %w", err) + } + + normalized := make([]string, 0, len(domains)) + seen := make(map[string]struct{}, len(domains)) + for _, domain := range domains { + item := normalizeProxyRouteDomainForMigration(domain) + if item == "" { + continue + } + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + if len(normalized) == 0 { + if primaryDomain == "" { + return nil, fmt.Errorf("proxy route domains are empty") + } + return []string{primaryDomain}, nil + } + if primaryDomain == "" { + primaryDomain = normalized[0] + } + if normalized[0] != primaryDomain { + rest := make([]string, 0, len(normalized)) + for _, domain := range normalized { + if domain == primaryDomain { + continue + } + rest = append(rest, domain) + } + normalized = append([]string{primaryDomain}, rest...) + } + return normalized, nil +} + +func backfillProxyRouteSiteFields(db *gorm.DB) error { + if db == nil { + return fmt.Errorf("database handle is nil") + } + if !db.Migrator().HasTable(&ProxyRoute{}) { + return nil + } + if !db.Migrator().HasColumn(&ProxyRoute{}, "site_name") || !db.Migrator().HasColumn(&ProxyRoute{}, "domains") { + return nil + } + + var routes []ProxyRoute + if err := db.Order("id asc").Find(&routes).Error; err != nil { + return fmt.Errorf("list proxy routes for site field backfill failed: %w", err) + } + for _, route := range routes { + domains, err := decodeProxyRouteDomainsForMigration(route.Domains, route.Domain) + if err != nil { + return fmt.Errorf("normalize proxy route %d domains failed: %w", route.ID, err) + } + domainsJSON, err := json.Marshal(domains) + if err != nil { + return fmt.Errorf("encode proxy route %d domains failed: %w", route.ID, err) + } + + primaryDomain := domains[0] + siteName := normalizeProxyRouteSiteNameForMigration(route.SiteName, primaryDomain) + updates := make(map[string]any, 3) + if route.Domain != primaryDomain { + updates["domain"] = primaryDomain + } + if route.SiteName != siteName { + updates["site_name"] = siteName + } + if strings.TrimSpace(route.Domains) != string(domainsJSON) { + updates["domains"] = string(domainsJSON) + } + if len(updates) == 0 { + continue + } + if err := db.Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(updates).Error; err != nil { + return fmt.Errorf("update proxy route %d site fields failed: %w", route.ID, err) + } + } + return nil +} + +func ensureProxyRouteSiteNameUniqueIndex(db *gorm.DB) error { + if db == nil { + return fmt.Errorf("database handle is nil") + } + if !db.Migrator().HasTable(&ProxyRoute{}) || !db.Migrator().HasColumn(&ProxyRoute{}, "site_name") { + return nil + } + return db.Exec(`CREATE UNIQUE INDEX IF NOT EXISTS idx_proxy_routes_site_name ON proxy_routes(site_name)`).Error +} + +func validateDatabaseSchemaV5(db *gorm.DB, backend string) error { + if err := validateDatabaseSchemaV4(db, backend); err != nil { + return err + } + if !db.Migrator().HasColumn(&ProxyRoute{}, "site_name") { + return fmt.Errorf("column proxy_routes.site_name is missing") + } + if !db.Migrator().HasColumn(&ProxyRoute{}, "domains") { + return fmt.Errorf("column proxy_routes.domains is missing") + } + + var routes []ProxyRoute + if err := db.Order("id asc").Find(&routes).Error; err != nil { + return fmt.Errorf("list proxy routes for validation failed: %w", err) + } + + siteNames := make(map[string]uint, len(routes)) + domainOwners := make(map[string]uint, len(routes)) + for _, route := range routes { + domains, err := decodeProxyRouteDomainsForMigration(route.Domains, route.Domain) + if err != nil { + return fmt.Errorf("proxy route %d domains are invalid: %w", route.ID, err) + } + if len(domains) == 0 { + return fmt.Errorf("proxy route %d domains are empty", route.ID) + } + if route.Domain != domains[0] { + return fmt.Errorf("proxy route %d primary domain mirror is invalid", route.ID) + } + + siteName := normalizeProxyRouteSiteNameForMigration(route.SiteName, domains[0]) + if siteName == "" { + return fmt.Errorf("proxy route %d site_name is empty", route.ID) + } + if existingID, ok := siteNames[siteName]; ok && existingID != route.ID { + return fmt.Errorf("proxy route site_name %s is duplicated", siteName) + } + siteNames[siteName] = route.ID + + localSeen := make(map[string]struct{}, len(domains)) + for _, domain := range domains { + if _, ok := localSeen[domain]; ok { + return fmt.Errorf("proxy route %d contains duplicated domain %s", route.ID, domain) + } + localSeen[domain] = struct{}{} + if existingID, ok := domainOwners[domain]; ok && existingID != route.ID { + return fmt.Errorf("proxy route domain %s is duplicated", domain) + } + domainOwners[domain] = route.ID + } + } + return nil +} + func renameLegacyObservabilityShardTables(db *gorm.DB) error { for _, baseTable := range shardedObservabilityBaseTables() { for _, table := range observabilityShardTables(baseTable) { @@ -549,11 +722,28 @@ func migrateV4(db *gorm.DB, backend string) error { return backfillOriginsFromProxyRoutes(db) } +// migrateV5 upgrades proxy_routes to website-level identity fields by +// backfilling site_name and domains while keeping domain as the primary-domain +// compatibility mirror. +func migrateV5(db *gorm.DB, backend string) error { + if err := applyCurrentSchema(db, backend); err != nil { + return err + } + if err := backfillOriginsFromProxyRoutes(db); err != nil { + return err + } + if err := backfillProxyRouteSiteFields(db); err != nil { + return err + } + return ensureProxyRouteSiteNameUniqueIndex(db) +} + func databaseSchemaMigrations() []databaseSchemaMigration { return []databaseSchemaMigration{ {fromVersion: 1, toVersion: 2, migrate: migrateV2, validate: validateDatabaseSchemaV2}, {fromVersion: 2, toVersion: 3, migrate: migrateV3, validate: validateDatabaseSchemaV3}, {fromVersion: 3, toVersion: 4, migrate: migrateV4, validate: validateDatabaseSchemaV4}, + {fromVersion: 4, toVersion: 5, migrate: migrateV5, validate: validateDatabaseSchemaV5}, } } @@ -618,13 +808,19 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error { if err := applyCurrentSchema(db, backend); err != nil { return err } - if err := backfillOriginsFromProxyRoutes(db); err != nil { - return err - } if err := migrateSQLiteDataIfNeeded(db, backend); err != nil { return err } - if err := validateDatabaseSchemaV4(db, backend); err != nil { + if err := backfillOriginsFromProxyRoutes(db); err != nil { + return err + } + if err := backfillProxyRouteSiteFields(db); err != nil { + return err + } + if err := ensureProxyRouteSiteNameUniqueIndex(db); err != nil { + return err + } + if err := validateDatabaseSchemaV5(db, backend); err != nil { return err } return saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion) diff --git a/openflare_server/model/proxy_route.go b/openflare_server/model/proxy_route.go index 40bf8110..7d5e3a2d 100644 --- a/openflare_server/model/proxy_route.go +++ b/openflare_server/model/proxy_route.go @@ -4,7 +4,9 @@ import "time" type ProxyRoute struct { ID uint `json:"id" gorm:"primaryKey"` + 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:'[]'"` 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"` @@ -28,7 +30,7 @@ func ListProxyRoutes() (routes []*ProxyRoute, err error) { } func GetEnabledProxyRoutes() (routes []*ProxyRoute, err error) { - err = DB.Where("enabled = ?", true).Order("domain asc").Find(&routes).Error + err = DB.Where("enabled = ?", true).Order("site_name asc").Order("domain asc").Find(&routes).Error return routes, err } @@ -49,7 +51,9 @@ func (route *ProxyRoute) Insert() error { func (route *ProxyRoute) Update() error { return DB.Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{ + "site_name": route.SiteName, "domain": route.Domain, + "domains": route.Domains, "origin_id": route.OriginID, "origin_url": route.OriginURL, "origin_host": route.OriginHost, diff --git a/openflare_server/service/config_version.go b/openflare_server/service/config_version.go index 3480f8f9..ecedf426 100644 --- a/openflare_server/service/config_version.go +++ b/openflare_server/service/config_version.go @@ -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" diff --git a/openflare_server/service/https_phase1_test.go b/openflare_server/service/https_phase1_test.go index 1518aafc..b6550073 100644 --- a/openflare_server/service/https_phase1_test.go +++ b/openflare_server/service/https_phase1_test.go @@ -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) diff --git a/openflare_server/service/proxy_route.go b/openflare_server/service/proxy_route.go index b690265a..6e4d86d3 100644 --- a/openflare_server/service/proxy_route.go +++ b/openflare_server/service/proxy_route.go @@ -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("源站地址不能为空")