mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 14:46:36 +08:00
[功能] 添加站点名称和多域名支持到代理路由,更新相关逻辑和测试
This commit is contained in:
@@ -4,7 +4,7 @@ import "time"
|
||||
|
||||
const (
|
||||
legacyDatabaseSchemaVersion = 1
|
||||
currentDatabaseSchemaVersion = 4
|
||||
currentDatabaseSchemaVersion = 5
|
||||
databaseSchemaVersionRowID = 1
|
||||
)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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