From 49472b54bf83e1d873db9e2da8e2ec8b86d2798c Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 1 Apr 2026 09:57:40 +0800 Subject: [PATCH] =?UTF-8?q?[=E5=8A=9F=E8=83=BD]=20=E6=B7=BB=E5=8A=A0?= =?UTF-8?q?=E5=9F=9F=E5=90=8D=E8=AF=81=E4=B9=A6=E7=BB=91=E5=AE=9A=E6=94=AF?= =?UTF-8?q?=E6=8C=81=EF=BC=8C=E5=85=81=E8=AE=B8=E4=B8=BA=E6=AF=8F=E4=B8=AA?= =?UTF-8?q?=E5=9F=9F=E5=90=8D=E5=8D=95=E7=8B=AC=E9=80=89=E6=8B=A9=E8=AF=81?= =?UTF-8?q?=E4=B9=A6=E5=B9=B6=E4=BC=98=E5=8C=96=E7=9B=B8=E5=85=B3=E9=80=BB?= =?UTF-8?q?=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/design.md | 4 +- docs/development-guidelines.md | 4 +- docs/website-configuration-redesign.md | 33 +- .../model/database_schema_version.go | 2 +- openflare_server/model/main_test.go | 104 ++++++ openflare_server/model/migrations.go | 301 +++++++++++++++++- openflare_server/model/proxy_route.go | 2 + openflare_server/service/config_version.go | 154 +++++++-- openflare_server/service/https_phase1_test.go | 58 +++- openflare_server/service/proxy_route.go | 218 ++++++++++++- openflare_server/service/tls_certificate.go | 9 + .../components/domain-list-input.tsx | 10 + .../components/proxy-route-config-page.tsx | 20 +- .../components/proxy-route-create-drawer.tsx | 14 + .../web/features/proxy-routes/helpers.ts | 1 + .../web/features/proxy-routes/types.ts | 2 + .../web/tests/unit/proxy-routes-page.test.tsx | 6 +- 17 files changed, 873 insertions(+), 69 deletions(-) diff --git a/docs/design.md b/docs/design.md index 531cfcd7..173b9b38 100644 --- a/docs/design.md +++ b/docs/design.md @@ -109,7 +109,9 @@ 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 允许为同一站点绑定多张证书,由服务端在同一 `server` 块内联合渲染并按域名完成覆盖校验 +* 网站级流量限制、反向代理与缓存配置当前按站点共享,不在同一网站内做域名级差异化配置;但 HTTPS 允许在同一站点内按域名绑定证书 +* `proxy_routes.domain_cert_ids` 用于记录与 `domains` 平行的域名证书绑定;值为 `0` 表示该域名不启用 HTTPS,仅保留 HTTP +* 发布渲染时,带证书的域名按证书分组输出独立 `443 ssl` `server` 块;未绑定证书的域名不得被自动带入 HTTPS * 发布渲染时必须将 `proxy_routes.domains` 中的全部域名一并纳入同一站点配置,避免同站点在版本快照中被拆散 * 所有上游地址都必须为合法 `http://` 或 `https://` * `config_versions` 必须保存完整快照、渲染结果与 `checksum` diff --git a/docs/development-guidelines.md b/docs/development-guidelines.md index 2436a7d1..14f5444d 100644 --- a/docs/development-guidelines.md +++ b/docs/development-guidelines.md @@ -124,7 +124,9 @@ * `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 的启停仍由站点级 `proxy_routes` 控制,但证书绑定必须通过与 `domains` 平行的 `domain_cert_ids` 记录逐域名保存;未绑定证书的域名不得参与 HTTPS 渲染 +* `proxy_routes.cert_ids` 仅作为站点级证书集合与兼容镜像,必须由 `domain_cert_ids` 推导生成;`cert_id` 继续作为首个已使用证书的兼容镜像 * `config_versions` 必须保存完整快照与渲染结果 * 全局同时只能有一个激活版本 * 回滚通过重新激活旧版本实现 diff --git a/docs/website-configuration-redesign.md b/docs/website-configuration-redesign.md index cb34ae85..b0891041 100644 --- a/docs/website-configuration-redesign.md +++ b/docs/website-configuration-redesign.md @@ -6,7 +6,7 @@ 现阶段已经出现以下真实需求: -* 多个域名指向同一站点,并共享反向代理、HTTPS、缓存等设置 +* 多个域名指向同一站点,并共享反向代理、缓存等设置,同时允许按域名分别绑定 HTTPS 证书 * 后续希望围绕“网站”继续叠加更多功能,而不是持续在规则列表中堆积字段 * 现有抽屉式编辑界面已经不适合承载更复杂的配置结构 @@ -29,7 +29,6 @@ * 域名设置 * 流量限制 * 反向代理 -* HTTPS * 缓存 ## 4. 核心模型要求 @@ -89,8 +88,7 @@ 1. 域名设置 2. 流量限制 3. 反向代理 -4. HTTPS -5. 缓存 +4. 缓存 为降低跨分区校验干扰,每个分区应支持独立保存与反馈;若采用统一保存,也必须提供未保存修改提示。 @@ -102,12 +100,16 @@ * 可维护 `domains` 列表 * 可新增、删除、排序域名 * 明确提示第一项为主域名 +* 每个域名可单独选择一张证书,形成与 `domains` 平行的 `domain_cert_ids` +* 若某个域名未选择证书,则该域名不启用 HTTPS +* `HTTP -> HTTPS` 跳转逻辑与域名证书绑定放在同一分区维护 * 保存前校验: `site_name` 非空且唯一 `domains` 非空 每个域名格式合法 域名在当前站点内不重复 域名在全局不与其他网站冲突 + 已选择证书的域名必须被对应证书覆盖 ### 5.4 流量限制 @@ -141,19 +143,7 @@ * 多上游模式下,上游项保持 `scheme://host[:port]` 形式 * 同一网站的多个上游在多上游模式下维持统一协议,降低渲染复杂度 -### 5.6 HTTPS - -HTTPS 分区负责维护站点级 TLS 行为,要求如下: - -* 支持开启或关闭 HTTPS -* 支持选择一张或多张证书 -* 支持保留现有 `HTTP -> HTTPS` 跳转能力 -* 当 HTTPS 开启时必须明确证书来源 -* 若只选择一张证书,则该证书必须覆盖当前网站的全部域名 -* 若选择多张证书,则所选证书集合必须联合覆盖当前网站的全部域名;任一域名至少要被其中一张证书覆盖 -* 发布渲染时应在同一 `server` 块内输出多组 `ssl_certificate` / `ssl_certificate_key`,并保证证书顺序稳定、文件输出可复用 - -### 5.7 缓存 +### 5.6 缓存 缓存分区负责维护站点级缓存策略,要求如下: @@ -177,6 +167,7 @@ HTTPS 分区负责维护站点级 TLS 行为,要求如下: 域名列表变更 站点级配置变更 * 发布渲染时,同一网站的全部域名必须落入同一份站点配置上下文中 +* 同一网站内,带证书的域名需按证书分组生成 HTTPS `server`;未配置证书的域名只保留 HTTP ## 7. 前端实现要求 @@ -191,8 +182,8 @@ HTTPS 分区负责维护站点级 TLS 行为,要求如下: 实施前必须准备显式数据库迁移与校验逻辑,至少包含: -1. 新增 `site_name` 与 `domains` 存储结构 -2. 将旧数据从单域名回填到站点结构 +1. 新增 `site_name`、`domains` 与 `domain_cert_ids` 存储结构 +2. 将旧数据从单域名回填到站点结构,并补齐逐域名证书映射 3. 为 `site_name` 建立唯一约束 4. 为域名唯一性建立可校验约束 5. 对迁移结果做一致性校验 @@ -243,7 +234,7 @@ HTTPS 分区负责维护站点级 TLS 行为,要求如下: * 列表页 UI 改造 * 子页面路由与布局 -* 域名设置、流量限制、反向代理、HTTPS、缓存五个分区 +* 域名设置、流量限制、反向代理、缓存四个分区 * 前端交互与表单测试 ### 阶段四:联调、发布验证与文档收口 @@ -270,6 +261,6 @@ HTTPS 分区负责维护站点级 TLS 行为,要求如下: * 原列表页已用“配置”按钮替代“编辑”按钮 * 网站配置子页面已经采用左侧菜单、右侧设置的布局 * 五个分区均可独立完成基本配置与保存 -* 发布后的渲染结果可正确覆盖同一网站的全部域名 +* 发布后的渲染结果可正确覆盖同一网站的全部域名,并只为已绑定证书的域名生成 HTTPS 配置 * Agent 同步、应用、回滚链路不被破坏 * 迁移、接口、渲染与前端关键路径均有对应测试或等效回归验证 diff --git a/openflare_server/model/database_schema_version.go b/openflare_server/model/database_schema_version.go index 7e257595..969f9922 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 = 7 + currentDatabaseSchemaVersion = 8 databaseSchemaVersionRowID = 1 ) diff --git a/openflare_server/model/main_test.go b/openflare_server/model/main_test.go index 37beb432..da389d0b 100644 --- a/openflare_server/model/main_test.go +++ b/openflare_server/model/main_test.go @@ -89,6 +89,36 @@ func (legacyProxyRouteV6) TableName() string { return "proxy_routes" } +type legacyProxyRouteV7 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 + CertIDs string `gorm:"type:text;not null;default:'[]'"` + 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 (legacyProxyRouteV7) TableName() string { + return "proxy_routes" +} + func openBareTestSQLiteDB(t *testing.T, name string) *gorm.DB { t.Helper() @@ -754,6 +784,80 @@ func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteCertificateListFields(t *test } } +func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteDomainCertificateFields(t *testing.T) { + db := openBareTestSQLiteDB(t, "legacy-proxy-route-domain-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(&legacyProxyRouteV7{}); err != nil { + t.Fatalf("auto migrate legacy proxy_routes v7: %v", err) + } + + now := time.Now().UTC() + certID := uint(9) + if err := db.Create(&legacyProxyRouteV7{ + 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, + CertIDs: `[9]`, + 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 v7: %v", err) + } + if err := saveDatabaseSchemaVersion(db, 7); 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) + } + + var domainCertIDs []uint + if err := json.Unmarshal([]byte(route.DomainCertIDs), &domainCertIDs); err != nil { + t.Fatalf("decode migrated domain_cert_ids: %v", err) + } + if len(domainCertIDs) != 2 || domainCertIDs[0] != certID || domainCertIDs[1] != certID { + t.Fatalf("unexpected migrated domain_cert_ids: %#v", domainCertIDs) + } +} + 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 b32c14b8..f5f62508 100644 --- a/openflare_server/model/migrations.go +++ b/openflare_server/model/migrations.go @@ -1,7 +1,9 @@ package model import ( + "crypto/x509" "encoding/json" + "encoding/pem" "errors" "fmt" "net" @@ -409,6 +411,218 @@ func backfillProxyRouteCertificateFields(db *gorm.DB) error { return nil } +func decodeProxyRouteDomainCertIDsForMigration( + raw string, + domainCount int, +) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []uint{}, nil + } + + var domainCertIDs []uint + if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil { + return nil, fmt.Errorf("decode proxy route domain_cert_ids failed: %w", err) + } + if len(domainCertIDs) == 0 { + return []uint{}, nil + } + if domainCount > 0 && len(domainCertIDs) != domainCount { + return nil, fmt.Errorf("proxy route domain_cert_ids length does not match domains") + } + + normalized := make([]uint, len(domainCertIDs)) + copy(normalized, domainCertIDs) + return normalized, nil +} + +func parseLeafCertificateForMigration(certPEM string) (*x509.Certificate, error) { + var firstErr error + rest := []byte(certPEM) + for len(rest) > 0 { + block, remaining := pem.Decode(rest) + if block == nil { + break + } + rest = remaining + if block.Type != "CERTIFICATE" { + continue + } + certificate, err := x509.ParseCertificate(block.Bytes) + if err == nil { + return certificate, nil + } + if firstErr == nil { + firstErr = err + } + } + if firstErr != nil { + return nil, firstErr + } + return nil, fmt.Errorf("parse certificate pem failed") +} + +func deriveProxyRouteDomainCertIDsForMigration( + db *gorm.DB, + domains []string, + certIDs []uint, +) ([]uint, error) { + if len(certIDs) == 0 { + return []uint{}, nil + } + if len(certIDs) == 1 { + result := make([]uint, len(domains)) + for index := range result { + result[index] = certIDs[0] + } + return result, nil + } + if len(certIDs) == len(domains) { + result := make([]uint, len(certIDs)) + copy(result, certIDs) + return result, nil + } + + var certificates []TLSCertificate + if err := db.Where("id IN ?", certIDs).Find(&certificates).Error; err != nil { + return nil, fmt.Errorf("load certificates for proxy route migration failed: %w", err) + } + certificateByID := make(map[uint]*x509.Certificate, len(certificates)) + for index := range certificates { + leaf, err := parseLeafCertificateForMigration(certificates[index].CertPEM) + if err != nil { + return nil, fmt.Errorf("parse certificate %d for proxy route migration failed: %w", certificates[index].ID, err) + } + certificateByID[certificates[index].ID] = leaf + } + + result := make([]uint, len(domains)) + for domainIndex, domain := range domains { + if domainIndex < len(certIDs) { + certificate := certificateByID[certIDs[domainIndex]] + if certificate != nil && certificate.VerifyHostname(domain) == nil { + result[domainIndex] = certIDs[domainIndex] + continue + } + } + + assigned := uint(0) + for _, certID := range certIDs { + certificate := certificateByID[certID] + if certificate != nil && certificate.VerifyHostname(domain) == nil { + assigned = certID + break + } + } + if assigned == 0 { + return nil, fmt.Errorf("no certificate covers domain %s", domain) + } + result[domainIndex] = assigned + } + return result, nil +} + +func uniqueProxyRouteCertIDsFromDomainAssignments(domainCertIDs []uint) []uint { + unique := make([]uint, 0, len(domainCertIDs)) + seen := make(map[uint]struct{}, len(domainCertIDs)) + for _, certID := range domainCertIDs { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + unique = append(unique, certID) + } + return unique +} + +func backfillProxyRouteDomainCertificateFields(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{}, "domain_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 domain certificate 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) + } + certIDs, err := decodeProxyRouteCertIDsForMigration(route.CertIDs, route.CertID) + if err != nil { + return fmt.Errorf("normalize proxy route %d cert_ids failed: %w", route.ID, err) + } + + domainCertIDs, err := decodeProxyRouteDomainCertIDsForMigration( + route.DomainCertIDs, + len(domains), + ) + if err != nil { + return fmt.Errorf("normalize proxy route %d domain_cert_ids failed: %w", route.ID, err) + } + if len(domainCertIDs) == 0 && len(certIDs) > 0 { + domainCertIDs, err = deriveProxyRouteDomainCertIDsForMigration( + db, + domains, + certIDs, + ) + if err != nil { + return fmt.Errorf("derive proxy route %d domain_cert_ids failed: %w", route.ID, err) + } + } + if !route.EnableHTTPS { + domainCertIDs = []uint{} + certIDs = []uint{} + } + + domainCertIDsJSON, err := json.Marshal(domainCertIDs) + if err != nil { + return fmt.Errorf("encode proxy route %d domain_cert_ids failed: %w", route.ID, err) + } + normalizedCertIDs := uniqueProxyRouteCertIDsFromDomainAssignments(domainCertIDs) + if len(domainCertIDs) == 0 { + normalizedCertIDs = []uint{} + } + certIDsJSON, err := json.Marshal(normalizedCertIDs) + if err != nil { + return fmt.Errorf("encode proxy route %d cert_ids failed: %w", route.ID, err) + } + + var primaryCertID *uint + if len(normalizedCertIDs) > 0 { + primaryCertID = &normalizedCertIDs[0] + } + + updates := make(map[string]any, 3) + if strings.TrimSpace(route.DomainCertIDs) != string(domainCertIDsJSON) { + updates["domain_cert_ids"] = string(domainCertIDsJSON) + } + 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 domain 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 @@ -515,6 +729,66 @@ func validateDatabaseSchemaV7(db *gorm.DB, backend string) error { return nil } +func validateDatabaseSchemaV8(db *gorm.DB, backend string) error { + if err := validateDatabaseSchemaV7(db, backend); err != nil { + return err + } + if !db.Migrator().HasColumn(&ProxyRoute{}, "domain_cert_ids") { + return fmt.Errorf("column proxy_routes.domain_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 domain certificate validation failed: %w", err) + } + 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) + } + domainCertIDs, err := decodeProxyRouteDomainCertIDsForMigration(route.DomainCertIDs, len(domains)) + if err != nil { + return fmt.Errorf("proxy route %d domain_cert_ids are invalid: %w", route.ID, err) + } + 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 { + if len(domainCertIDs) != 0 { + return fmt.Errorf("proxy route %d has domain_cert_ids while https is disabled", route.ID) + } + continue + } + if len(domainCertIDs) != len(domains) { + return fmt.Errorf("proxy route %d domain_cert_ids length is invalid", route.ID) + } + normalizedCertIDs := uniqueProxyRouteCertIDsFromDomainAssignments(domainCertIDs) + if len(normalizedCertIDs) == 0 { + return fmt.Errorf("proxy route %d has https enabled without domain certificate assignments", route.ID) + } + if !uintSlicesEqualForMigration(certIDs, normalizedCertIDs) { + return fmt.Errorf("proxy route %d cert_ids mirror is invalid", route.ID) + } + if route.CertID == nil || *route.CertID != normalizedCertIDs[0] { + return fmt.Errorf("proxy route %d primary cert_id mirror is invalid", route.ID) + } + } + return nil +} + +func uintSlicesEqualForMigration(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 renameLegacyObservabilityShardTables(db *gorm.DB) error { for _, baseTable := range shardedObservabilityBaseTables() { for _, table := range observabilityShardTables(baseTable) { @@ -901,6 +1175,27 @@ func migrateV7(db *gorm.DB, backend string) error { return backfillProxyRouteCertificateFields(db) } +// migrateV8 adds per-domain certificate assignments to proxy_routes while +// keeping cert_ids as the website-level compatibility mirror. +func migrateV8(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 + } + if err := backfillProxyRouteCertificateFields(db); err != nil { + return err + } + return backfillProxyRouteDomainCertificateFields(db) +} + func databaseSchemaMigrations() []databaseSchemaMigration { return []databaseSchemaMigration{ {fromVersion: 1, toVersion: 2, migrate: migrateV2, validate: validateDatabaseSchemaV2}, @@ -909,6 +1204,7 @@ func databaseSchemaMigrations() []databaseSchemaMigration { {fromVersion: 4, toVersion: 5, migrate: migrateV5, validate: validateDatabaseSchemaV5}, {fromVersion: 5, toVersion: 6, migrate: migrateV6, validate: validateDatabaseSchemaV6}, {fromVersion: 6, toVersion: 7, migrate: migrateV7, validate: validateDatabaseSchemaV7}, + {fromVersion: 7, toVersion: 8, migrate: migrateV8, validate: validateDatabaseSchemaV8}, } } @@ -988,7 +1284,10 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error { if err := backfillProxyRouteCertificateFields(db); err != nil { return err } - if err := validateDatabaseSchemaV7(db, backend); err != nil { + if err := backfillProxyRouteDomainCertificateFields(db); err != nil { + return err + } + if err := validateDatabaseSchemaV8(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 fefe50cb..f34208bc 100644 --- a/openflare_server/model/proxy_route.go +++ b/openflare_server/model/proxy_route.go @@ -15,6 +15,7 @@ type ProxyRoute struct { 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:'[]'"` + DomainCertIDs string `json:"domain_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"` @@ -66,6 +67,7 @@ func (route *ProxyRoute) Update() error { "enable_https": route.EnableHTTPS, "cert_id": route.CertID, "cert_ids": route.CertIDs, + "domain_cert_ids": route.DomainCertIDs, "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 b67bfa8e..356e9cd0 100644 --- a/openflare_server/service/config_version.go +++ b/openflare_server/service/config_version.go @@ -74,6 +74,7 @@ type snapshotRoute struct { EnableHTTPS bool `json:"enable_https"` CertID *uint `json:"cert_id,omitempty"` CertIDs []uint `json:"cert_ids,omitempty"` + DomainCertIDs []uint `json:"domain_cert_ids,omitempty"` RedirectHTTP bool `json:"redirect_http"` LimitConnPerServer int `json:"limit_conn_per_server,omitempty"` LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"` @@ -472,6 +473,7 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) { EnableHTTPS: route.EnableHTTPS, CertID: route.CertID, CertIDs: mustDecodeSnapshotCertIDs(route), + DomainCertIDs: mustDecodeSnapshotDomainCertIDs(route, domains), RedirectHTTP: route.RedirectHTTP, LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, @@ -497,6 +499,24 @@ func mustDecodeSnapshotCertIDs(route *model.ProxyRoute) []uint { return certIDs } +func mustDecodeSnapshotDomainCertIDs( + route *model.ProxyRoute, + domains []string, +) []uint { + if route == nil { + return []uint{} + } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return []uint{} + } + domainCertIDs, err := resolveProxyRouteDomainCertIDs(route, domains, certIDs) + if err != nil { + return []uint{} + } + return domainCertIDs +} + func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) { text := strings.TrimSpace(snapshotJSON) if text == "" { @@ -544,6 +564,14 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute { routes[index].CertID = primaryCertID routes[index].CertIDs = normalizedCertIDs } + normalizedDomainCertIDs, err := normalizeSnapshotDomainCertificateIDs( + routes[index].Domains, + routes[index].CertIDs, + routes[index].DomainCertIDs, + ) + if err == nil { + routes[index].DomainCertIDs = normalizedDomainCertIDs + } normalizedUpstreams, err := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams) if err == nil { routes[index].OriginURL = normalizedUpstreams[0] @@ -583,7 +611,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 || !uintSliceEqual(left.CertIDs, right.CertIDs) { + 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) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) { return false } if len(left.Domains) != len(right.Domains) { @@ -814,48 +842,79 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot) 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 + domainCertIDs, err := resolveProxyRouteDomainCertIDs(route, domains, certIDs) + if err != nil { + return "", nil, fmt.Errorf("route %s domain_cert_ids are invalid: %w", route.Domain, err) } if route.CertID == nil || *route.CertID == 0 { return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain) } - certificate, err := model.GetTLSCertificateByID(*route.CertID) + if len(certIDs) == 0 { + return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain) + } + certificates, err := loadTLSCertificates(certIDs) if err != nil { - return "", nil, fmt.Errorf("路由 %s 关联证书不存在", route.Domain) + return "", nil, fmt.Errorf("route %s certificate lookup failed: %w", route.Domain, err) } - if err := validateCertificateCoverage(certificate, domains); err != nil { - return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err) + certificateByID := make(map[uint]*model.TLSCertificate, len(certificates)) + for _, certificate := range certificates { + if certificate == nil { + continue + } + certificateByID[certificate.ID] = certificate + supportFiles = append(supportFiles, + SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)}, + SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)}, + ) } - supportFiles = append(supportFiles, - SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)}, - SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)}, - ) + + httpOnlyDomains := make([]string, 0, len(domains)) + domainsByCertID := make(map[uint][]string, len(certIDs)) + for index, domain := range domains { + if index >= len(domainCertIDs) || domainCertIDs[index] == 0 { + httpOnlyDomains = append(httpOnlyDomains, domain) + continue + } + domainsByCertID[domainCertIDs[index]] = append( + domainsByCertID[domainCertIDs[index]], + domain, + ) + } + for _, certID := range certIDs { + assignedDomains := domainsByCertID[certID] + if len(assignedDomains) == 0 { + continue + } + certificate := certificateByID[certID] + if certificate == nil { + return "", nil, fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID) + } + if err := validateCertificateCoverage(certificate, assignedDomains); err != nil { + return "", nil, fmt.Errorf("site %s certificate validation failed: %w", displayName, err) + } + } + if route.RedirectHTTP { - builder.WriteString(renderHTTPRedirectServer(serverNames)) + if len(httpOnlyDomains) > 0 { + builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) + } + for _, certID := range certIDs { + assignedDomains := domainsByCertID[certID] + if len(assignedDomains) == 0 { + continue + } + builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains))) + } } else { builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) } - builder.WriteString(renderHTTPSServer(serverNames, route.OriginURL, route.OriginHost, certificate.ID, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) + for _, certID := range certIDs { + assignedDomains := domainsByCertID[certID] + if len(assignedDomains) == 0 { + continue + } + builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), route.OriginURL, route.OriginHost, certID, customHeaders, cacheConfig, limitConfig, upstreamConfig, cfg)) + } } return builder.String(), dedupeSupportFiles(supportFiles), nil } @@ -988,6 +1047,37 @@ func normalizeSnapshotCertificateIDs(primaryCertID *uint, certIDs []uint) ([]uin return normalized, normalizedPrimary, nil } +func normalizeSnapshotDomainCertificateIDs( + domains []string, + certIDs []uint, + domainCertIDs []uint, +) ([]uint, error) { + if len(domainCertIDs) > 0 { + if len(domains) > 0 && len(domainCertIDs) != len(domains) { + return nil, errors.New("snapshot domain_cert_ids length is invalid") + } + normalized := make([]uint, len(domainCertIDs)) + copy(normalized, domainCertIDs) + return normalized, nil + } + if len(certIDs) == 0 { + return []uint{}, nil + } + if len(certIDs) == 1 { + normalized := make([]uint, len(domains)) + for index := range normalized { + normalized[index] = certIDs[0] + } + return normalized, nil + } + if len(certIDs) == len(domains) { + normalized := make([]uint, len(certIDs)) + copy(normalized, certIDs) + return normalized, nil + } + return []uint{}, nil +} + func uintPointerEqual(left *uint, right *uint) bool { if left == nil || right == nil { return left == nil && right == nil diff --git a/openflare_server/service/https_phase1_test.go b/openflare_server/service/https_phase1_test.go index 9bdcea6b..d63809fc 100644 --- a/openflare_server/service/https_phase1_test.go +++ b/openflare_server/service/https_phase1_test.go @@ -403,7 +403,7 @@ func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t OriginURL: "https://origin.internal", Enabled: true, EnableHTTPS: true, - CertIDs: []uint{appCertificate.ID, wwwCertificate.ID}, + DomainCertIDs: []uint{appCertificate.ID, wwwCertificate.ID}, RedirectHTTP: true, CacheEnabled: true, CachePolicy: proxyRouteCachePolicyPathPrefix, @@ -419,6 +419,9 @@ func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t 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) } + if len(route.DomainCertIDs) != 2 || route.DomainCertIDs[0] != appCertificate.ID || route.DomainCertIDs[1] != wwwCertificate.ID { + t.Fatalf("expected domain_cert_ids to persist per domain, got %#v", route.DomainCertIDs) + } result, err := PublishConfigVersion("root") if err != nil { @@ -439,6 +442,59 @@ func TestPublishConfigVersionRendersMultipleCertificatesForMultiDomainWebsite(t if !strings.Contains(result.Version.SnapshotJSON, `"cert_ids":[`) { t.Fatal("expected snapshot to include cert_ids") } + if !strings.Contains(result.Version.SnapshotJSON, `"domain_cert_ids":[`) { + t.Fatal("expected snapshot to include domain_cert_ids") + } +} + +func TestPublishConfigVersionSkipsHTTPSForDomainsWithoutCertificate(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) + } + + route, err := CreateProxyRoute(ProxyRouteInput{ + SiteName: "partial-https-site", + Domains: []string{"app.example.com", "www.example.com"}, + OriginURL: "https://origin.internal", + Enabled: true, + EnableHTTPS: true, + DomainCertIDs: []uint{appCertificate.ID, 0}, + RedirectHTTP: true, + }) + if err != nil { + t.Fatalf("CreateProxyRoute failed: %v", err) + } + if len(route.CertIDs) != 1 || route.CertIDs[0] != appCertificate.ID { + t.Fatalf("expected website cert_ids to keep used certificates only, got %#v", route.CertIDs) + } + if len(route.DomainCertIDs) != 2 || route.DomainCertIDs[0] != appCertificate.ID || route.DomainCertIDs[1] != 0 { + t.Fatalf("expected domain_cert_ids to preserve unassigned domains, got %#v", route.DomainCertIDs) + } + + result, err := PublishConfigVersion("root") + if err != nil { + t.Fatalf("PublishConfigVersion failed: %v", err) + } + if strings.Contains(result.Version.RenderedConfig, "listen 443 ssl;\n http2 on;\n server_name app.example.com www.example.com;") { + t.Fatal("expected https server block to exclude domains without certificate") + } + if !strings.Contains(result.Version.RenderedConfig, "listen 443 ssl;\n http2 on;\n server_name app.example.com;") { + t.Fatal("expected https server block to contain only the certified domain") + } + if !strings.Contains(result.Version.RenderedConfig, "listen 80;\n server_name app.example.com;\n\n return 301 https://$host$request_uri;") { + t.Fatal("expected certified domain to keep http redirect") + } + if !strings.Contains(result.Version.RenderedConfig, "listen 80;\n server_name www.example.com;") { + t.Fatal("expected non-certified domain to stay on plain http") + } } func TestDiffConfigVersionTracksAddedDomainWithinWebsite(t *testing.T) { diff --git a/openflare_server/service/proxy_route.go b/openflare_server/service/proxy_route.go index 5ac99557..405cde5d 100644 --- a/openflare_server/service/proxy_route.go +++ b/openflare_server/service/proxy_route.go @@ -44,6 +44,7 @@ type ProxyRouteInput struct { EnableHTTPS bool `json:"enable_https"` CertID *uint `json:"cert_id"` CertIDs []uint `json:"cert_ids"` + DomainCertIDs []uint `json:"domain_cert_ids"` RedirectHTTP bool `json:"redirect_http"` LimitConnPerServer int `json:"limit_conn_per_server"` LimitConnPerIP int `json:"limit_conn_per_ip"` @@ -71,6 +72,7 @@ type ProxyRouteView struct { EnableHTTPS bool `json:"enable_https"` CertID *uint `json:"cert_id"` CertIDs []uint `json:"cert_ids"` + DomainCertIDs []uint `json:"domain_cert_ids"` RedirectHTTP bool `json:"redirect_http"` LimitConnPerServer int `json:"limit_conn_per_server"` LimitConnPerIP int `json:"limit_conn_per_ip"` @@ -194,14 +196,33 @@ 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 !input.EnableHTTPS { + input.RedirectHTTP = false + input.CertID = nil + input.CertIDs = nil + input.DomainCertIDs = nil + } + domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs( + domains, + input.EnableHTTPS, + input.DomainCertIDs, + input.CertID, + input.CertIDs, + ) if err != nil { return nil, err } + if err := validateProxyRouteDomainCertificateCoverage(domains, domainCertIDs); err != nil { + return nil, err + } certIDsJSON, err := json.Marshal(certIDs) if err != nil { return nil, err } + domainCertIDsJSON, err := json.Marshal(domainCertIDs) + if err != nil { + return nil, err + } domainsJSON, err := json.Marshal(domains) if err != nil { return nil, err @@ -216,15 +237,9 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro if err := validateOriginHost(originHost); err != nil { return nil, err } - if !input.EnableHTTPS { - input.RedirectHTTP = false - input.CertID = nil - input.CertIDs = nil - } + input.DomainCertIDs = domainCertIDs input.CertIDs = certIDs - if len(certIDs) > 0 { - input.CertID = &certIDs[0] - } + input.CertID = primaryCertID if input.RedirectHTTP && !input.EnableHTTPS { return nil, errors.New("redirect_http requires enable_https") } @@ -243,6 +258,7 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro route.EnableHTTPS = input.EnableHTTPS route.CertID = input.CertID route.CertIDs = string(certIDsJSON) + route.DomainCertIDs = string(domainCertIDsJSON) route.RedirectHTTP = input.RedirectHTTP route.LimitConnPerServer = limitConnPerServer route.LimitConnPerIP = limitConnPerIP @@ -291,6 +307,10 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) { if err != nil { return nil, err } + domainCertIDs, err := resolveProxyRouteDomainCertIDs(route, domains, certIDs) + if err != nil { + return nil, err + } var certID *uint if len(certIDs) > 0 { certID = &certIDs[0] @@ -312,6 +332,7 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) { EnableHTTPS: route.EnableHTTPS, CertID: certID, CertIDs: certIDs, + DomainCertIDs: domainCertIDs, RedirectHTTP: route.RedirectHTTP, LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, @@ -476,6 +497,185 @@ func normalizeProxyRouteCertificateIDs(enableHTTPS bool, certID *uint, certIDs [ return normalized, nil } +func normalizeProxyRouteDomainCertificateIDs( + domains []string, + enableHTTPS bool, + rawDomainCertIDs []uint, + certID *uint, + certIDs []uint, +) ([]uint, []uint, *uint, error) { + if !enableHTTPS { + return []uint{}, []uint{}, nil, nil + } + + if len(rawDomainCertIDs) > 0 { + if len(rawDomainCertIDs) != len(domains) { + return nil, nil, nil, errors.New("domain_cert_ids must match domains length") + } + + normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs)) + uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs)) + seen := make(map[uint]struct{}, len(rawDomainCertIDs)) + hasAssignedCertificate := false + for index, item := range rawDomainCertIDs { + if item == 0 { + continue + } + if _, err := model.GetTLSCertificateByID(item); err != nil { + return nil, nil, nil, errors.New("selected certificate does not exist") + } + normalizedDomainCertIDs[index] = item + hasAssignedCertificate = true + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + uniqueCertIDs = append(uniqueCertIDs, item) + } + if !hasAssignedCertificate { + return nil, nil, nil, errors.New("must select a certificate when HTTPS is enabled") + } + + primaryCertID := &uniqueCertIDs[0] + return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil + } + + normalizedCertIDs, err := normalizeProxyRouteCertificateIDs( + enableHTTPS, + certID, + certIDs, + ) + if err != nil { + return nil, nil, nil, err + } + + switch { + case len(normalizedCertIDs) == 0: + return nil, nil, nil, errors.New("must select a certificate when HTTPS is enabled") + case len(normalizedCertIDs) == 1: + domainCertIDs := make([]uint, len(domains)) + for index := range domainCertIDs { + domainCertIDs[index] = normalizedCertIDs[0] + } + primaryCertID := &normalizedCertIDs[0] + return domainCertIDs, normalizedCertIDs, primaryCertID, nil + case len(normalizedCertIDs) == len(domains): + domainCertIDs := make([]uint, len(normalizedCertIDs)) + copy(domainCertIDs, normalizedCertIDs) + primaryCertID := &normalizedCertIDs[0] + return domainCertIDs, normalizedCertIDs, primaryCertID, nil + default: + domainCertIDs, err := deriveDomainCertIDsFromCertificateSet( + domains, + normalizedCertIDs, + ) + if err != nil { + return nil, nil, nil, err + } + primaryCertID := &normalizedCertIDs[0] + return domainCertIDs, normalizedCertIDs, primaryCertID, nil + } +} + +func validateProxyRouteDomainCertificateCoverage( + domains []string, + domainCertIDs []uint, +) error { + if len(domainCertIDs) == 0 { + return nil + } + + domainsByCertID := make(map[uint][]string) + for index, certID := range domainCertIDs { + if certID == 0 { + continue + } + domainsByCertID[certID] = append(domainsByCertID[certID], domains[index]) + } + + for certID, assignedDomains := range domainsByCertID { + certificate, err := model.GetTLSCertificateByID(certID) + if err != nil { + return errors.New("selected certificate does not exist") + } + if err := validateCertificateCoverage(certificate, assignedDomains); err != nil { + return err + } + } + return nil +} + +func deriveDomainCertIDsFromCertificateSet( + domains []string, + certIDs []uint, +) ([]uint, error) { + certificates, err := loadTLSCertificates(certIDs) + if err != nil { + return nil, err + } + + result := make([]uint, len(domains)) + for domainIndex, domain := range domains { + if domainIndex < len(certificates) && + certificates[domainIndex] != nil && + validateCertificateCoverage(certificates[domainIndex], []string{domain}) == nil { + result[domainIndex] = certificates[domainIndex].ID + continue + } + + assigned := uint(0) + for _, certificate := range certificates { + if certificate != nil && + validateCertificateCoverage(certificate, []string{domain}) == nil { + assigned = certificate.ID + break + } + } + if assigned == 0 { + return nil, fmt.Errorf("certificate does not cover domain %s", domain) + } + result[domainIndex] = assigned + } + return result, nil +} + +func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []uint{}, nil + } + + var domainCertIDs []uint + if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil { + return nil, errors.New("domain_cert_ids payload is invalid") + } + if len(domainCertIDs) == 0 { + return []uint{}, nil + } + if domainCount > 0 && len(domainCertIDs) != domainCount { + return nil, errors.New("domain_cert_ids length does not match domains") + } + + normalized := make([]uint, len(domainCertIDs)) + copy(normalized, domainCertIDs) + return normalized, nil +} + +func resolveProxyRouteDomainCertIDs( + route *model.ProxyRoute, + domains []string, + certIDs []uint, +) ([]uint, error) { + domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, len(domains)) + if err != nil { + return nil, err + } + if len(domainCertIDs) > 0 || len(certIDs) == 0 { + return domainCertIDs, nil + } + return deriveDomainCertIDsFromCertificateSet(domains, certIDs) +} + func normalizeProxyRouteLimitRate(raw string) (string, error) { normalized := strings.ToLower(strings.TrimSpace(raw)) if normalized == "" || normalized == "0" { diff --git a/openflare_server/service/tls_certificate.go b/openflare_server/service/tls_certificate.go index 0e061853..6bcbfa05 100644 --- a/openflare_server/service/tls_certificate.go +++ b/openflare_server/service/tls_certificate.go @@ -125,6 +125,15 @@ func DeleteTLSCertificate(id uint) error { return errors.New("certificate is still referenced by proxy routes") } } + domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, 0) + if err != nil { + return fmt.Errorf("proxy route %d domain_cert_ids payload is invalid: %w", route.ID, err) + } + for _, certID := range domainCertIDs { + if certID == id { + return errors.New("certificate is still referenced by proxy routes") + } + } } certificate, err := model.GetTLSCertificateByID(id) diff --git a/openflare_server/web/features/proxy-routes/components/domain-list-input.tsx b/openflare_server/web/features/proxy-routes/components/domain-list-input.tsx index 9f7905a6..47a6df1d 100644 --- a/openflare_server/web/features/proxy-routes/components/domain-list-input.tsx +++ b/openflare_server/web/features/proxy-routes/components/domain-list-input.tsx @@ -90,12 +90,22 @@ function buildDomainSuggestions( export function buildDomainRowsFromRoute( domains: string[], + domainCertIDs: number[], certIDs: number[], ): DomainListRow[] { if (domains.length === 0) { return ensureRows([]); } + if (domainCertIDs.length === domains.length) { + return domains.map((domain, index) => ({ + domain, + certificateId: domainCertIDs[index] + ? String(domainCertIDs[index]) + : '', + })); + } + if (certIDs.length === 0) { return domains.map((domain) => ({ domain, certificateId: '' })); } 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 1bb13f8e..6e8a2bb5 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 @@ -216,12 +216,24 @@ function normalizeSelectedCertificateIDs(rows: DomainListRow[]) { return Array.from( new Set( rows + .filter((item) => item.domain.trim() !== '') .map((item) => Number(item.certificateId)) .filter((item) => Number.isFinite(item) && item > 0), ), ); } +function buildDomainCertificateIDs(rows: DomainListRow[]) { + return rows + .filter((item) => item.domain.trim() !== '') + .map((item) => { + const certificateID = Number(item.certificateId); + return Number.isFinite(certificateID) && certificateID > 0 + ? certificateID + : 0; + }); +} + function buildDomainRows(route: ProxyRouteItem) { const selectedCertIDs = route.cert_ids.length > 0 @@ -230,7 +242,11 @@ function buildDomainRows(route: ProxyRouteItem) { ? [route.cert_id] : []; - return buildDomainRowsFromRoute(route.domains, selectedCertIDs); + return buildDomainRowsFromRoute( + route.domains, + route.domain_cert_ids, + selectedCertIDs, + ); } function ConfigSectionShell({ @@ -311,6 +327,7 @@ function DomainSettingsSection({ const domains = values.domain_rows .map((item) => item.domain.trim().toLowerCase()) .filter(Boolean); + const domainCertIDs = buildDomainCertificateIDs(values.domain_rows); const certIDs = normalizeSelectedCertificateIDs(values.domain_rows); onSave( @@ -322,6 +339,7 @@ function DomainSettingsSection({ enable_https: certIDs.length > 0, cert_id: certIDs[0] ?? null, cert_ids: certIDs, + domain_cert_ids: domainCertIDs, redirect_http: certIDs.length > 0 ? values.redirect_http : false, }), { message: '域名设置已保存。' }, diff --git a/openflare_server/web/features/proxy-routes/components/proxy-route-create-drawer.tsx b/openflare_server/web/features/proxy-routes/components/proxy-route-create-drawer.tsx index 302bf5a9..9cd51aca 100644 --- a/openflare_server/web/features/proxy-routes/components/proxy-route-create-drawer.tsx +++ b/openflare_server/web/features/proxy-routes/components/proxy-route-create-drawer.tsx @@ -95,12 +95,24 @@ function normalizeSelectedCertificateIDs(rows: DomainListRow[]) { return Array.from( new Set( rows + .filter((item) => item.domain.trim() !== '') .map((item) => Number(item.certificateId)) .filter((item) => Number.isFinite(item) && item > 0), ), ); } +function buildDomainCertificateIDs(rows: DomainListRow[]) { + return rows + .filter((item) => item.domain.trim() !== '') + .map((item) => { + const certificateID = Number(item.certificateId); + return Number.isFinite(certificateID) && certificateID > 0 + ? certificateID + : 0; + }); +} + export function ProxyRouteCreateDrawer({ open, onOpenChange, @@ -143,6 +155,7 @@ export function ProxyRouteCreateDrawer({ const domains = values.domain_rows .map((item) => item.domain.trim().toLowerCase()) .filter(Boolean); + const domainCertIDs = buildDomainCertificateIDs(values.domain_rows); const selectedCertIDs = normalizeSelectedCertificateIDs(values.domain_rows); const { urls } = parseOriginUrls(values.origin_urls_text); const primaryOrigin = parseOriginUrl(urls[0]); @@ -168,6 +181,7 @@ export function ProxyRouteCreateDrawer({ enable_https: selectedCertIDs.length > 0, cert_id: selectedCertIDs[0] ?? null, cert_ids: selectedCertIDs, + domain_cert_ids: domainCertIDs, redirect_http: selectedCertIDs.length > 0 ? values.redirect_http : false, limit_conn_per_server: 0, limit_conn_per_ip: 0, diff --git a/openflare_server/web/features/proxy-routes/helpers.ts b/openflare_server/web/features/proxy-routes/helpers.ts index e9d53a11..dfbbf96d 100644 --- a/openflare_server/web/features/proxy-routes/helpers.ts +++ b/openflare_server/web/features/proxy-routes/helpers.ts @@ -275,6 +275,7 @@ export function buildPayloadFromRoute( enable_https: route.enable_https, cert_id: route.cert_id, cert_ids: route.cert_ids, + domain_cert_ids: route.domain_cert_ids, redirect_http: route.redirect_http, limit_conn_per_server: route.limit_conn_per_server, limit_conn_per_ip: route.limit_conn_per_ip, diff --git a/openflare_server/web/features/proxy-routes/types.ts b/openflare_server/web/features/proxy-routes/types.ts index 371ea6bb..bb0b59c7 100644 --- a/openflare_server/web/features/proxy-routes/types.ts +++ b/openflare_server/web/features/proxy-routes/types.ts @@ -19,6 +19,7 @@ export interface ProxyRouteItem { enable_https: boolean; cert_id: number | null; cert_ids: number[]; + domain_cert_ids: number[]; redirect_http: boolean; limit_conn_per_server: number; limit_conn_per_ip: number; @@ -50,6 +51,7 @@ export interface ProxyRouteMutationPayload { enable_https: boolean; cert_id: number | null; cert_ids?: number[]; + domain_cert_ids?: number[]; redirect_http: boolean; limit_conn_per_server?: number; limit_conn_per_ip?: number; diff --git a/openflare_server/web/tests/unit/proxy-routes-page.test.tsx b/openflare_server/web/tests/unit/proxy-routes-page.test.tsx index a1b08b56..ebd72419 100644 --- a/openflare_server/web/tests/unit/proxy-routes-page.test.tsx +++ b/openflare_server/web/tests/unit/proxy-routes-page.test.tsx @@ -50,6 +50,7 @@ function buildRoute(overrides: Record = {}) { enable_https: true, cert_id: 1, cert_ids: [1], + domain_cert_ids: [1, 0], redirect_http: true, limit_conn_per_server: 120, limit_conn_per_ip: 12, @@ -185,6 +186,7 @@ describe('Proxy route website pages', () => { enable_https: payload.enable_https, cert_id: payload.cert_id, cert_ids: payload.cert_ids ?? [], + domain_cert_ids: payload.domain_cert_ids ?? [], redirect_http: payload.redirect_http, limit_conn_per_server: 0, limit_conn_per_ip: 0, @@ -337,6 +339,7 @@ describe('Proxy route website pages', () => { enable_https: payload.enable_https, cert_id: payload.cert_id, cert_ids: payload.cert_ids, + domain_cert_ids: payload.domain_cert_ids, redirect_http: payload.redirect_http, }), }), @@ -404,7 +407,7 @@ describe('Proxy route website pages', () => { await user.type(secondaryDomainInput, 'www.brand.example.com'); await user.selectOptions(screen.getByLabelText('证书 1'), '1'); - await user.selectOptions(screen.getByLabelText('证书 2'), '1'); + await user.selectOptions(screen.getByLabelText('证书 2'), ''); const saveButton = document.querySelector( 'button[form="proxy-route-domains-form"]', @@ -427,6 +430,7 @@ describe('Proxy route website pages', () => { enable_https: true, cert_id: 1, cert_ids: [1], + domain_cert_ids: [1, 0], redirect_http: true, }); });