diff --git a/docs/design.md b/docs/design.md index 56d7bccc..531cfcd7 100644 --- a/docs/design.md +++ b/docs/design.md @@ -109,7 +109,7 @@ Origin * `proxy_routes` 至少包含一个上游地址;为兼容历史数据保留 `origin_url` 主上游字段,也允许在同一规则内补充多个上游做负载均衡 * `proxy_routes` 上游统一渲染为带 keepalive 的 named `upstream`;单上游可附带 base path 或 query 并在 `proxy_pass` 中追加,多上游仍限定为纯 `scheme://host[:port]` * `proxy_routes.origin_host` 为可选字段,用于回源时覆盖 `Host` 请求头;未设置时默认透传访问域名 -* 网站级流量限制、反向代理、HTTPS 与缓存配置当前按站点共享,不在同一网站内做域名级差异化配置 +* 网站级流量限制、反向代理、HTTPS 与缓存配置当前按站点共享,不在同一网站内做域名级差异化配置;但 HTTPS 允许为同一站点绑定多张证书,由服务端在同一 `server` 块内联合渲染并按域名完成覆盖校验 * 发布渲染时必须将 `proxy_routes.domains` 中的全部域名一并纳入同一站点配置,避免同站点在版本快照中被拆散 * 所有上游地址都必须为合法 `http://` 或 `https://` * `config_versions` 必须保存完整快照、渲染结果与 `checksum` diff --git a/docs/development-guidelines.md b/docs/development-guidelines.md index 0fafe879..2436a7d1 100644 --- a/docs/development-guidelines.md +++ b/docs/development-guidelines.md @@ -124,7 +124,7 @@ * `proxy_routes` 如关联 `origins`,必须同时保存可直接渲染的 `origin_url`;源站地址变更时,由 service 负责同步更新引用该源站的规则快照 * `proxy_routes` 的上游统一使用 named `upstream` + keepalive;单上游如带 base path 或 query,应在 `proxy_pass` 上补回 URI,多上游仅允许纯 `scheme://host[:port]` * `proxy_routes.origin_host` 为可选字段,仅用于覆盖回源 `Host` 请求头,不引入新的平台化对象 -* 流量限制、反向代理、HTTPS 与缓存配置当前都归属站点级 `proxy_routes`,同一网站内不拆分域名级差异配置 +* 流量限制、反向代理、HTTPS 与缓存配置当前都归属站点级 `proxy_routes`,同一网站内不拆分域名级差异配置;其中 HTTPS 可绑定一张或多张证书,但证书选择仍属于站点级配置而非域名级配置 * `config_versions` 必须保存完整快照与渲染结果 * 全局同时只能有一个激活版本 * 回滚通过重新激活旧版本实现 diff --git a/docs/website-configuration-redesign.md b/docs/website-configuration-redesign.md index 1a16283e..cb34ae85 100644 --- a/docs/website-configuration-redesign.md +++ b/docs/website-configuration-redesign.md @@ -146,10 +146,12 @@ HTTPS 分区负责维护站点级 TLS 行为,要求如下: * 支持开启或关闭 HTTPS -* 支持选择证书 +* 支持选择一张或多张证书 * 支持保留现有 `HTTP -> HTTPS` 跳转能力 * 当 HTTPS 开启时必须明确证书来源 -* 应校验证书是否覆盖当前网站的全部域名;若无法覆盖,应阻止保存或给出不可忽略的错误提示 +* 若只选择一张证书,则该证书必须覆盖当前网站的全部域名 +* 若选择多张证书,则所选证书集合必须联合覆盖当前网站的全部域名;任一域名至少要被其中一张证书覆盖 +* 发布渲染时应在同一 `server` 块内输出多组 `ssl_certificate` / `ssl_certificate_key`,并保证证书顺序稳定、文件输出可复用 ### 5.7 缓存 diff --git a/openflare_server/model/database_schema_version.go b/openflare_server/model/database_schema_version.go index c2627bb7..7e257595 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 = 6 + currentDatabaseSchemaVersion = 7 databaseSchemaVersionRowID = 1 ) diff --git a/openflare_server/model/main_test.go b/openflare_server/model/main_test.go index 9c14a349..37beb432 100644 --- a/openflare_server/model/main_test.go +++ b/openflare_server/model/main_test.go @@ -60,6 +60,35 @@ func (legacyProxyRouteV5) TableName() string { return "proxy_routes" } +type legacyProxyRouteV6 struct { + ID uint `gorm:"primaryKey"` + SiteName string `gorm:"size:255;not null;default:''"` + Domain string `gorm:"uniqueIndex;size:255;not null"` + Domains string `gorm:"type:text;not null;default:'[]'"` + 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"` + LimitConnPerServer int `gorm:"not null;default:0"` + LimitConnPerIP int `gorm:"not null;default:0"` + LimitRate string `gorm:"size:32;not null;default:''"` + 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 (legacyProxyRouteV6) TableName() string { + return "proxy_routes" +} + func openBareTestSQLiteDB(t *testing.T, name string) *gorm.DB { t.Helper() @@ -649,6 +678,82 @@ func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteRateLimitFields(t *testing.T) } } +func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteCertificateListFields(t *testing.T) { + db := openBareTestSQLiteDB(t, "legacy-proxy-route-cert-ids.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(&legacyProxyRouteV6{}); err != nil { + t.Fatalf("auto migrate legacy proxy_routes v6: %v", err) + } + + now := time.Now().UTC() + certID := uint(9) + if err := db.Create(&legacyProxyRouteV6{ + SiteName: "secure-site", + Domain: "secure.example.com", + Domains: `["secure.example.com","www.secure.example.com"]`, + OriginURL: "https://origin-secure.internal:8443", + Upstreams: `["https://origin-secure.internal:8443"]`, + Enabled: true, + EnableHTTPS: true, + CertID: &certID, + RedirectHTTP: true, + LimitConnPerServer: 120, + LimitConnPerIP: 12, + LimitRate: "512k", + CacheEnabled: false, + CachePolicy: "", + CacheRules: `[]`, + CustomHeaders: `[]`, + CreatedAt: now, + UpdatedAt: now, + }).Error; err != nil { + t.Fatalf("seed legacy proxy route v6: %v", err) + } + if err := saveDatabaseSchemaVersion(db, 6); 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.CertID == nil || *route.CertID != certID { + t.Fatalf("expected cert_id mirror to be preserved, got %+v", route.CertID) + } + + var certIDs []uint + if err := json.Unmarshal([]byte(route.CertIDs), &certIDs); err != nil { + t.Fatalf("decode migrated cert_ids: %v", err) + } + if len(certIDs) != 1 || certIDs[0] != certID { + t.Fatalf("unexpected migrated cert_ids: %#v", certIDs) + } +} + 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 bb494af1..b32c14b8 100644 --- a/openflare_server/model/migrations.go +++ b/openflare_server/model/migrations.go @@ -330,6 +330,85 @@ func ensureProxyRouteSiteNameUniqueIndex(db *gorm.DB) error { return db.Exec(`CREATE UNIQUE INDEX IF NOT EXISTS idx_proxy_routes_site_name ON proxy_routes(site_name)`).Error } +func decodeProxyRouteCertIDsForMigration(raw string, fallbackCertID *uint) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + if fallbackCertID == nil || *fallbackCertID == 0 { + return []uint{}, nil + } + return []uint{*fallbackCertID}, nil + } + + var certIDs []uint + if err := json.Unmarshal([]byte(text), &certIDs); err != nil { + return nil, fmt.Errorf("decode proxy route cert_ids failed: %w", err) + } + + normalized := make([]uint, 0, len(certIDs)) + seen := make(map[uint]struct{}, len(certIDs)) + for _, certID := range certIDs { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + normalized = append(normalized, certID) + } + if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 { + return []uint{*fallbackCertID}, nil + } + return normalized, nil +} + +func backfillProxyRouteCertificateFields(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{}, "cert_ids") { + return nil + } + + var routes []ProxyRoute + if err := db.Order("id asc").Find(&routes).Error; err != nil { + return fmt.Errorf("list proxy routes for certificate field backfill failed: %w", err) + } + for _, route := range routes { + certIDs, err := decodeProxyRouteCertIDsForMigration(route.CertIDs, route.CertID) + if err != nil { + return fmt.Errorf("normalize proxy route %d cert_ids failed: %w", route.ID, err) + } + certIDsJSON, err := json.Marshal(certIDs) + if err != nil { + return fmt.Errorf("encode proxy route %d cert_ids failed: %w", route.ID, err) + } + + var primaryCertID *uint + if len(certIDs) > 0 { + primaryCertID = &certIDs[0] + } + + updates := make(map[string]any, 2) + if strings.TrimSpace(route.CertIDs) != string(certIDsJSON) { + updates["cert_ids"] = string(certIDsJSON) + } + if (route.CertID == nil) != (primaryCertID == nil) || (route.CertID != nil && primaryCertID != nil && *route.CertID != *primaryCertID) { + updates["cert_id"] = primaryCertID + } + 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 certificate fields failed: %w", route.ID, err) + } + } + return nil +} + func validateDatabaseSchemaV5(db *gorm.DB, backend string) error { if err := validateDatabaseSchemaV4(db, backend); err != nil { return err @@ -400,6 +479,42 @@ func validateDatabaseSchemaV6(db *gorm.DB, backend string) error { return nil } +func validateDatabaseSchemaV7(db *gorm.DB, backend string) error { + if err := validateDatabaseSchemaV6(db, backend); err != nil { + return err + } + if !db.Migrator().HasColumn(&ProxyRoute{}, "cert_ids") { + return fmt.Errorf("column proxy_routes.cert_ids is missing") + } + + var routes []ProxyRoute + if err := db.Order("id asc").Find(&routes).Error; err != nil { + return fmt.Errorf("list proxy routes for certificate validation failed: %w", err) + } + for _, route := range routes { + certIDs, err := decodeProxyRouteCertIDsForMigration(route.CertIDs, route.CertID) + if err != nil { + return fmt.Errorf("proxy route %d cert_ids are invalid: %w", route.ID, err) + } + if route.EnableHTTPS && len(certIDs) == 0 { + return fmt.Errorf("proxy route %d has https enabled without cert_ids", route.ID) + } + if !route.EnableHTTPS && route.RedirectHTTP { + return fmt.Errorf("proxy route %d enables redirect_http without https", route.ID) + } + if len(certIDs) == 0 { + if route.CertID != nil { + return fmt.Errorf("proxy route %d primary cert_id mirror is invalid", route.ID) + } + continue + } + if route.CertID == nil || *route.CertID != certIDs[0] { + return fmt.Errorf("proxy route %d primary cert_id mirror is invalid", route.ID) + } + } + return nil +} + func renameLegacyObservabilityShardTables(db *gorm.DB) error { for _, baseTable := range shardedObservabilityBaseTables() { for _, table := range observabilityShardTables(baseTable) { @@ -768,6 +883,24 @@ func migrateV6(db *gorm.DB, backend string) error { return ensureProxyRouteSiteNameUniqueIndex(db) } +// migrateV7 adds structured website-level certificate lists to proxy_routes +// while keeping cert_id as the primary certificate compatibility mirror. +func migrateV7(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 + } + if err := ensureProxyRouteSiteNameUniqueIndex(db); err != nil { + return err + } + return backfillProxyRouteCertificateFields(db) +} + func databaseSchemaMigrations() []databaseSchemaMigration { return []databaseSchemaMigration{ {fromVersion: 1, toVersion: 2, migrate: migrateV2, validate: validateDatabaseSchemaV2}, @@ -775,6 +908,7 @@ func databaseSchemaMigrations() []databaseSchemaMigration { {fromVersion: 3, toVersion: 4, migrate: migrateV4, validate: validateDatabaseSchemaV4}, {fromVersion: 4, toVersion: 5, migrate: migrateV5, validate: validateDatabaseSchemaV5}, {fromVersion: 5, toVersion: 6, migrate: migrateV6, validate: validateDatabaseSchemaV6}, + {fromVersion: 6, toVersion: 7, migrate: migrateV7, validate: validateDatabaseSchemaV7}, } } @@ -851,7 +985,10 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error { if err := ensureProxyRouteSiteNameUniqueIndex(db); err != nil { return err } - if err := validateDatabaseSchemaV6(db, backend); err != nil { + if err := backfillProxyRouteCertificateFields(db); err != nil { + return err + } + if err := validateDatabaseSchemaV7(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 4d844d85..fefe50cb 100644 --- a/openflare_server/model/proxy_route.go +++ b/openflare_server/model/proxy_route.go @@ -14,6 +14,7 @@ type ProxyRoute struct { Enabled bool `json:"enabled" gorm:"not null;default:true"` EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"` CertID *uint `json:"cert_id"` + CertIDs string `json:"cert_ids" gorm:"type:text;not null;default:'[]'"` RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"` LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"` LimitConnPerIP int `json:"limit_conn_per_ip" gorm:"not null;default:0"` @@ -64,6 +65,7 @@ func (route *ProxyRoute) Update() error { "enabled": route.Enabled, "enable_https": route.EnableHTTPS, "cert_id": route.CertID, + "cert_ids": route.CertIDs, "redirect_http": route.RedirectHTTP, "limit_conn_per_server": route.LimitConnPerServer, "limit_conn_per_ip": route.LimitConnPerIP, diff --git a/openflare_server/service/config_version.go b/openflare_server/service/config_version.go index ac255636..b67bfa8e 100644 --- a/openflare_server/service/config_version.go +++ b/openflare_server/service/config_version.go @@ -73,6 +73,7 @@ type snapshotRoute struct { Enabled bool `json:"enabled"` EnableHTTPS bool `json:"enable_https"` CertID *uint `json:"cert_id,omitempty"` + CertIDs []uint `json:"cert_ids,omitempty"` RedirectHTTP bool `json:"redirect_http"` LimitConnPerServer int `json:"limit_conn_per_server,omitempty"` LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"` @@ -470,6 +471,7 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) { Enabled: route.Enabled, EnableHTTPS: route.EnableHTTPS, CertID: route.CertID, + CertIDs: mustDecodeSnapshotCertIDs(route), RedirectHTTP: route.RedirectHTTP, LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, @@ -484,6 +486,17 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) { return items, nil } +func mustDecodeSnapshotCertIDs(route *model.ProxyRoute) []uint { + if route == nil { + return []uint{} + } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return []uint{} + } + return certIDs +} + func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) { text := strings.TrimSpace(snapshotJSON) if text == "" { @@ -526,6 +539,11 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute { if err == nil { routes[index].CustomHeaders = normalizedHeaders } + normalizedCertIDs, primaryCertID, err := normalizeSnapshotCertificateIDs(routes[index].CertID, routes[index].CertIDs) + if err == nil { + routes[index].CertID = primaryCertID + routes[index].CertIDs = normalizedCertIDs + } normalizedUpstreams, err := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams) if err == nil { routes[index].OriginURL = normalizedUpstreams[0] @@ -565,7 +583,7 @@ func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRo } func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool { - 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.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || left.LimitRate != right.LimitRate || 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.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || !uintSliceEqual(left.CertIDs, right.CertIDs) { return false } if len(left.Domains) != len(right.Domains) { @@ -792,6 +810,32 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot) builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) continue } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return "", nil, fmt.Errorf("route %s cert_ids are invalid: %w", route.Domain, err) + } + if len(certIDs) > 0 { + certificates, err := loadTLSCertificates(certIDs) + if err != nil { + return "", nil, fmt.Errorf("route %s certificate lookup failed: %w", route.Domain, err) + } + if err := validateCertificateCoverageSet(certificates, domains); err != nil { + return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err) + } + for _, certificate := range certificates { + 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(serverNames)) + } else { + builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) + } + builder.WriteString(renderHTTPSServerWithCertificates(serverNames, route.OriginURL, route.OriginHost, certIDs, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) + continue + } if route.CertID == nil || *route.CertID == 0 { return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain) } @@ -917,6 +961,33 @@ func onOff(value bool) string { return "off" } +func normalizeSnapshotCertificateIDs(primaryCertID *uint, certIDs []uint) ([]uint, *uint, error) { + candidates := make([]uint, 0, len(certIDs)+1) + if primaryCertID != nil && *primaryCertID != 0 { + candidates = append(candidates, *primaryCertID) + } + candidates = append(candidates, certIDs...) + + normalized := make([]uint, 0, len(candidates)) + seen := make(map[uint]struct{}, len(candidates)) + for _, certID := range candidates { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + normalized = append(normalized, certID) + } + + var normalizedPrimary *uint + if len(normalized) > 0 { + normalizedPrimary = &normalized[0] + } + return normalized, normalizedPrimary, nil +} + func uintPointerEqual(left *uint, right *uint) bool { if left == nil || right == nil { return left == nil && right == nil @@ -924,6 +995,18 @@ func uintPointerEqual(left *uint, right *uint) bool { return *left == *right } +func uintSliceEqual(left []uint, right []uint) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] != right[index] { + return false + } + } + return true +} + func checksum(content string) string { sum := sha256.Sum256([]byte(content)) return hex.EncodeToString(sum[:]) @@ -971,6 +1054,17 @@ func renderHTTPSServer(serverNames string, originURL string, originHost string, 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%s }\n}\n\n", serverNames, certPath, keyPath, renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig)) } +func renderHTTPSServerWithCertificates(serverNames string, originURL string, originHost string, certificateIDs []uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, cfg openRestyConfigSnapshot) string { + var certificateBlock strings.Builder + for _, certificateID := range certificateIDs { + certPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateCertFileName(certificateID)) + keyPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateKeyFileName(certificateID)) + certificateBlock.WriteString(fmt.Sprintf(" ssl_certificate %s;\n", certPath)) + certificateBlock.WriteString(fmt.Sprintf(" ssl_certificate_key %s;\n", keyPath)) + } + return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n%s\n location / {\n%s%s%s%s }\n}\n\n", serverNames, certificateBlock.String(), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig)) +} + func renderServerNames(domains []string) string { return strings.Join(domains, " ") } @@ -991,6 +1085,48 @@ func validateCertificateCoverage(certificate *model.TLSCertificate, domains []st return nil } +func validateCertificateCoverageSet(certificates []*model.TLSCertificate, domains []string) error { + if len(certificates) == 0 { + return errors.New("certificate set is empty") + } + leaves := make([]interface{ VerifyHostname(string) error }, 0, len(certificates)) + for _, certificate := range certificates { + if certificate == nil { + return errors.New("certificate is nil") + } + leaf, err := parseLeafCertificate(certificate.CertPEM) + if err != nil { + return err + } + leaves = append(leaves, leaf) + } + for _, domain := range domains { + covered := false + for _, leaf := range leaves { + if leaf.VerifyHostname(domain) == nil { + covered = true + break + } + } + if !covered { + return fmt.Errorf("certificate does not cover domain %s", domain) + } + } + return nil +} + +func loadTLSCertificates(certIDs []uint) ([]*model.TLSCertificate, error) { + certificates := make([]*model.TLSCertificate, 0, len(certIDs)) + for _, certID := range certIDs { + certificate, err := model.GetTLSCertificateByID(certID) + if err != nil { + return nil, err + } + certificates = append(certificates, certificate) + } + return certificates, nil +} + func renderConnectionUpgradeMap() string { return " map $http_upgrade $connection_upgrade {\n default upgrade;\n '' \"\";\n }\n\n" } diff --git a/openflare_server/service/https_phase1_test.go b/openflare_server/service/https_phase1_test.go index 9b2c794f..9bdcea6b 100644 --- a/openflare_server/service/https_phase1_test.go +++ b/openflare_server/service/https_phase1_test.go @@ -374,6 +374,73 @@ func TestPublishConfigVersionRendersMultiDomainWebsite(t *testing.T) { } } +func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t *testing.T) { + setupServiceTestDB(t) + + appCertPEM, appKeyPEM := generateCertificatePair(t, []string{"app.example.com"}) + appCertificate, err := CreateTLSCertificate(TLSCertificateInput{ + Name: "app-only", + CertPEM: appCertPEM, + KeyPEM: appKeyPEM, + }) + if err != nil { + t.Fatalf("CreateTLSCertificate app-only failed: %v", err) + } + + wwwCertPEM, wwwKeyPEM := generateCertificatePair(t, []string{"www.example.com"}) + wwwCertificate, err := CreateTLSCertificate(TLSCertificateInput{ + Name: "www-only", + CertPEM: wwwCertPEM, + KeyPEM: wwwKeyPEM, + }) + if err != nil { + t.Fatalf("CreateTLSCertificate www-only failed: %v", err) + } + + route, err := CreateProxyRoute(ProxyRouteInput{ + SiteName: "marketing-site", + Domains: []string{"app.example.com", "www.example.com"}, + OriginURL: "https://origin.internal", + Enabled: true, + EnableHTTPS: true, + CertIDs: []uint{appCertificate.ID, wwwCertificate.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) + } + if route.CertID == nil || *route.CertID != appCertificate.ID { + t.Fatalf("expected primary cert mirror to point at first certificate, got %#v", route.CertID) + } + if len(route.CertIDs) != 2 || route.CertIDs[0] != appCertificate.ID || route.CertIDs[1] != wwwCertificate.ID { + t.Fatalf("expected cert_ids to persist in order, got %#v", route.CertIDs) + } + + result, err := PublishConfigVersion("root") + if err != nil { + t.Fatalf("PublishConfigVersion failed: %v", err) + } + if strings.Count(result.Version.RenderedConfig, "ssl_certificate __OPENFLARE_CERT_DIR__/") != 2 { + t.Fatalf("expected rendered config to include two ssl_certificate directives, got %s", result.Version.RenderedConfig) + } + if strings.Count(result.Version.RenderedConfig, "ssl_certificate_key __OPENFLARE_CERT_DIR__/") != 2 { + t.Fatalf("expected rendered config to include two ssl_certificate_key directives, got %s", result.Version.RenderedConfig) + } + if !strings.Contains(result.Version.SupportFilesJSON, certificateCertFileName(appCertificate.ID)) { + t.Fatal("expected support files to include first certificate") + } + if !strings.Contains(result.Version.SupportFilesJSON, certificateCertFileName(wwwCertificate.ID)) { + t.Fatal("expected support files to include second certificate") + } + if !strings.Contains(result.Version.SnapshotJSON, `"cert_ids":[`) { + t.Fatal("expected snapshot to include cert_ids") + } +} + func TestDiffConfigVersionTracksAddedDomainWithinWebsite(t *testing.T) { setupServiceTestDB(t) diff --git a/openflare_server/service/proxy_route.go b/openflare_server/service/proxy_route.go index 403536fe..5ac99557 100644 --- a/openflare_server/service/proxy_route.go +++ b/openflare_server/service/proxy_route.go @@ -43,6 +43,7 @@ type ProxyRouteInput struct { Enabled bool `json:"enabled"` EnableHTTPS bool `json:"enable_https"` CertID *uint `json:"cert_id"` + CertIDs []uint `json:"cert_ids"` RedirectHTTP bool `json:"redirect_http"` LimitConnPerServer int `json:"limit_conn_per_server"` LimitConnPerIP int `json:"limit_conn_per_ip"` @@ -69,6 +70,7 @@ type ProxyRouteView struct { Enabled bool `json:"enabled"` EnableHTTPS bool `json:"enable_https"` CertID *uint `json:"cert_id"` + CertIDs []uint `json:"cert_ids"` RedirectHTTP bool `json:"redirect_http"` LimitConnPerServer int `json:"limit_conn_per_server"` LimitConnPerIP int `json:"limit_conn_per_ip"` @@ -192,6 +194,14 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro if err != nil { return nil, err } + certIDs, err := normalizeProxyRouteCertificateIDs(input.EnableHTTPS, input.CertID, input.CertIDs) + if err != nil { + return nil, err + } + certIDsJSON, err := json.Marshal(certIDs) + if err != nil { + return nil, err + } domainsJSON, err := json.Marshal(domains) if err != nil { return nil, err @@ -209,14 +219,11 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro if !input.EnableHTTPS { input.RedirectHTTP = false input.CertID = nil + input.CertIDs = nil } - if input.EnableHTTPS { - if input.CertID == nil || *input.CertID == 0 { - return nil, errors.New("must select a certificate when HTTPS is enabled") - } - if _, err := model.GetTLSCertificateByID(*input.CertID); err != nil { - return nil, errors.New("selected certificate does not exist") - } + input.CertIDs = certIDs + if len(certIDs) > 0 { + input.CertID = &certIDs[0] } if input.RedirectHTTP && !input.EnableHTTPS { return nil, errors.New("redirect_http requires enable_https") @@ -235,6 +242,7 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro route.Enabled = input.Enabled route.EnableHTTPS = input.EnableHTTPS route.CertID = input.CertID + route.CertIDs = string(certIDsJSON) route.RedirectHTTP = input.RedirectHTTP route.LimitConnPerServer = limitConnPerServer route.LimitConnPerIP = limitConnPerIP @@ -279,6 +287,14 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) { if err != nil { return nil, err } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return nil, err + } + var certID *uint + if len(certIDs) > 0 { + certID = &certIDs[0] + } primaryDomain := domains[0] return &ProxyRouteView{ ID: route.ID, @@ -294,7 +310,8 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) { UpstreamList: upstreams, Enabled: route.Enabled, EnableHTTPS: route.EnableHTTPS, - CertID: route.CertID, + CertID: certID, + CertIDs: certIDs, RedirectHTTP: route.RedirectHTTP, LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, @@ -427,6 +444,38 @@ func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) { return value, nil } +func normalizeProxyRouteCertificateIDs(enableHTTPS bool, certID *uint, certIDs []uint) ([]uint, error) { + if !enableHTTPS { + return []uint{}, nil + } + + candidates := make([]uint, 0, len(certIDs)+1) + if certID != nil && *certID != 0 { + candidates = append(candidates, *certID) + } + candidates = append(candidates, certIDs...) + + normalized := make([]uint, 0, len(candidates)) + seen := make(map[uint]struct{}, len(candidates)) + for _, item := range candidates { + if item == 0 { + continue + } + if _, ok := seen[item]; ok { + continue + } + if _, err := model.GetTLSCertificateByID(item); err != nil { + return nil, errors.New("selected certificate does not exist") + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + if len(normalized) == 0 { + return nil, errors.New("must select a certificate when HTTPS is enabled") + } + return normalized, nil +} + func normalizeProxyRouteLimitRate(raw string) (string, error) { normalized := strings.ToLower(strings.TrimSpace(raw)) if normalized == "" || normalized == "0" { @@ -732,6 +781,36 @@ func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) { return normalizeProxyRouteDomains(domains) } +func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + if fallbackCertID == nil || *fallbackCertID == 0 { + return []uint{}, nil + } + return []uint{*fallbackCertID}, nil + } + var certIDs []uint + if err := json.Unmarshal([]byte(text), &certIDs); err != nil { + return nil, errors.New("cert_ids payload is invalid") + } + normalized := make([]uint, 0, len(certIDs)) + seen := make(map[uint]struct{}, len(certIDs)) + for _, certID := range certIDs { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + normalized = append(normalized, certID) + } + if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 { + return []uint{*fallbackCertID}, nil + } + return normalized, nil +} + func validateOriginURL(raw string) error { if raw == "" { return errors.New("origin URL cannot be empty") diff --git a/openflare_server/service/tls_certificate.go b/openflare_server/service/tls_certificate.go index a1e40178..0e061853 100644 --- a/openflare_server/service/tls_certificate.go +++ b/openflare_server/service/tls_certificate.go @@ -2,6 +2,7 @@ package service import ( "crypto/tls" + "encoding/json" "errors" "fmt" "mime/multipart" @@ -54,7 +55,7 @@ func CreateTLSCertificate(input TLSCertificateInput) (*model.TLSCertificate, err } if err = certificate.Insert(); err != nil { if isUniqueConstraintError(err) { - return nil, errors.New("证书名称已存在") + return nil, errors.New("certificate name already exists") } return nil, err } @@ -63,7 +64,7 @@ func CreateTLSCertificate(input TLSCertificateInput) (*model.TLSCertificate, err func CreateTLSCertificateFromFiles(name string, certFile *multipart.FileHeader, keyFile *multipart.FileHeader, remark string) (*model.TLSCertificate, error) { if certFile == nil || keyFile == nil { - return nil, errors.New("证书文件和私钥文件不能为空") + return nil, errors.New("certificate file and key file cannot be empty") } certContent, err := readMultipartFile(certFile) if err != nil { @@ -101,13 +102,31 @@ func UpdateTLSCertificate(id uint, input TLSCertificateInput) (*model.TLSCertifi } func DeleteTLSCertificate(id uint) error { - var routeCount int64 - if err := model.DB.Model(&model.ProxyRoute{}).Where("cert_id = ?", id).Count(&routeCount).Error; err != nil { + routes, err := model.ListProxyRoutes() + if err != nil { return err } - if routeCount > 0 { - return errors.New("证书仍被反代规则引用,无法删除") + for _, route := range routes { + if route == nil { + continue + } + if route.CertID != nil && *route.CertID == id { + return errors.New("certificate is still referenced by proxy routes") + } + if strings.TrimSpace(route.CertIDs) == "" { + continue + } + var certIDs []uint + if err := json.Unmarshal([]byte(route.CertIDs), &certIDs); err != nil { + return fmt.Errorf("proxy route %d cert_ids payload is invalid: %w", route.ID, err) + } + for _, certID := range certIDs { + if certID == id { + return errors.New("certificate is still referenced by proxy routes") + } + } } + certificate, err := model.GetTLSCertificateByID(id) if err != nil { return err @@ -121,17 +140,17 @@ func buildTLSCertificate(existing *model.TLSCertificate, input TLSCertificateInp keyPEM := strings.TrimSpace(input.KeyPEM) remark := strings.TrimSpace(input.Remark) if name == "" { - return nil, errors.New("证书名称不能为空") + return nil, errors.New("certificate name cannot be empty") } if certPEM == "" || keyPEM == "" { - return nil, errors.New("证书内容和私钥内容不能为空") + return nil, errors.New("certificate content and key content cannot be empty") } parsed, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) if err != nil { - return nil, fmt.Errorf("证书或私钥格式不合法: %w", err) + return nil, fmt.Errorf("certificate or key format is invalid: %w", err) } if len(parsed.Certificate) == 0 { - return nil, errors.New("证书内容不合法") + return nil, errors.New("certificate content is invalid") } leaf, err := parseLeafCertificate(certPEM) if err != nil { diff --git a/openflare_server/web/features/config-versions/components/config-versions-page.tsx b/openflare_server/web/features/config-versions/components/config-versions-page.tsx index b11f845c..ca733d5a 100644 --- a/openflare_server/web/features/config-versions/components/config-versions-page.tsx +++ b/openflare_server/web/features/config-versions/components/config-versions-page.tsx @@ -300,9 +300,6 @@ function PublishPreviewCard({

Pending Main Config

-

- {`Checksum: ${preview.checksum}`} -

{preview.main_config} diff --git a/openflare_server/web/features/proxy-routes/components/proxy-route-config-page.tsx b/openflare_server/web/features/proxy-routes/components/proxy-route-config-page.tsx index cb95e832..23795538 100644 --- a/openflare_server/web/features/proxy-routes/components/proxy-route-config-page.tsx +++ b/openflare_server/web/features/proxy-routes/components/proxy-route-config-page.tsx @@ -159,14 +159,14 @@ const reverseProxySchema = z const httpsSchema = z .object({ enable_https: z.boolean(), - cert_id: z.string(), + cert_ids: z.array(z.string()), redirect_http: z.boolean(), }) .superRefine((value, context) => { - if (value.enable_https && !value.cert_id.trim()) { + if (value.enable_https && value.cert_ids.length === 0) { context.addIssue({ code: z.ZodIssueCode.custom, - path: ['cert_id'], + path: ['cert_ids'], message: '启用 HTTPS 时必须选择证书', }); } @@ -542,7 +542,12 @@ function HTTPSSection({ resolver: zodResolver(httpsSchema), defaultValues: { enable_https: route.enable_https, - cert_id: route.cert_id ? String(route.cert_id) : '', + cert_ids: + route.cert_ids.length > 0 + ? route.cert_ids.map((certID) => String(certID)) + : route.cert_id + ? [String(route.cert_id)] + : [], redirect_http: route.redirect_http, }, }); @@ -550,12 +555,18 @@ function HTTPSSection({ useEffect(() => { form.reset({ enable_https: route.enable_https, - cert_id: route.cert_id ? String(route.cert_id) : '', + cert_ids: + route.cert_ids.length > 0 + ? route.cert_ids.map((certID) => String(certID)) + : route.cert_id + ? [String(route.cert_id)] + : [], redirect_http: route.redirect_http, }); }, [form, route]); const watchedEnableHTTPS = form.watch('enable_https'); + const watchedCertIDs = form.watch('cert_ids'); return ( Number(value) > 0) + ? Number( + values.cert_ids.find((value) => Number(value) > 0) ?? 0, + ) + : null, + cert_ids: values.enable_https + ? values.cert_ids + .map((value) => Number(value)) + .filter((value) => Number.isFinite(value) && value > 0) + : [], redirect_http: values.enable_https ? values.redirect_http : false, }), { message: 'HTTPS 设置已保存。' }, @@ -585,7 +607,7 @@ function HTTPSSection({ onChange={(checked) => { form.setValue('enable_https', checked, { shouldDirty: true }); if (!checked) { - form.setValue('cert_id', '', { shouldDirty: true }); + form.setValue('cert_ids', [], { shouldDirty: true }); form.setValue('redirect_http', false, { shouldDirty: true }); } }} @@ -593,20 +615,30 @@ function HTTPSSection({ {certificates.map((certificate) => ( ))} + {watchedEnableHTTPS && watchedCertIDs.length > 0 ? ( +

+ 已选择 {watchedCertIDs.length} 张证书,发布时会校验证书集合是否覆盖全部域名。 +

+ ) : null}