[功能] 支持为 HTTPS 启用多个证书,更新相关逻辑和测试

This commit is contained in:
ryan
2026-03-31 14:16:32 +08:00
parent 97fa56b1af
commit cff815bd47
16 changed files with 618 additions and 40 deletions
+1 -1
View File
@@ -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`
+1 -1
View File
@@ -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` 必须保存完整快照与渲染结果
* 全局同时只能有一个激活版本
* 回滚通过重新激活旧版本实现
+4 -2
View File
@@ -146,10 +146,12 @@
HTTPS 分区负责维护站点级 TLS 行为,要求如下:
* 支持开启或关闭 HTTPS
* 支持选择证书
* 支持选择一张或多张证书
* 支持保留现有 `HTTP -> HTTPS` 跳转能力
* 当 HTTPS 开启时必须明确证书来源
* 应校验证书是否覆盖当前网站的全部域名;若无法覆盖,应阻止保存或给出不可忽略的错误提示
* 若只选择一张证书,则该证书必须覆盖当前网站的全部域名
* 若选择多张证书,则所选证书集合必须联合覆盖当前网站的全部域名;任一域名至少要被其中一张证书覆盖
* 发布渲染时应在同一 `server` 块内输出多组 `ssl_certificate` / `ssl_certificate_key`,并保证证书顺序稳定、文件输出可复用
### 5.7 缓存
@@ -4,7 +4,7 @@ import "time"
const (
legacyDatabaseSchemaVersion = 1
currentDatabaseSchemaVersion = 6
currentDatabaseSchemaVersion = 7
databaseSchemaVersionRowID = 1
)
+105
View File
@@ -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")
+138 -1
View File
@@ -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)
+2
View File
@@ -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,
+137 -1
View File
@@ -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"
}
@@ -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)
+87 -8
View File
@@ -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")
+29 -10
View File
@@ -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 {
@@ -300,9 +300,6 @@ function PublishPreviewCard({
<p className="text-sm font-semibold text-[var(--foreground-primary)]">
Pending Main Config
</p>
<p className="text-xs text-[var(--foreground-secondary)]">
{`Checksum: ${preview.checksum}`}
</p>
</div>
<CodeBlock className="max-h-[32rem] whitespace-pre-wrap">
{preview.main_config}
@@ -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 (
<ConfigSectionShell
@@ -571,7 +582,18 @@ function HTTPSSection({
onSave(
buildPayloadFromRoute(route, {
enable_https: values.enable_https,
cert_id: values.enable_https && values.cert_id ? Number(values.cert_id) : null,
cert_id:
values.enable_https &&
values.cert_ids.some((value) => 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({
<ResourceField
label="证书"
error={form.formState.errors.cert_id?.message}
error={form.formState.errors.cert_ids?.message}
hint="请确保该证书能覆盖当前站点的全部域名。"
>
<ResourceSelect
multiple
size={Math.min(Math.max(certificates.length, 4), 8)}
className="min-h-44"
disabled={!watchedEnableHTTPS}
{...form.register('cert_id')}
{...form.register('cert_ids')}
>
<option value="">请选择证书</option>
{certificates.map((certificate) => (
<option key={certificate.id} value={certificate.id}>
{certificate.name}
{certificate.not_after
? `${certificate.name} · ${certificate.not_after}`
: certificate.name}
</option>
))}
</ResourceSelect>
{watchedEnableHTTPS && watchedCertIDs.length > 0 ? (
<p className="text-xs leading-5 text-[var(--foreground-secondary)]">
已选择 {watchedCertIDs.length} 张证书,发布时会校验证书集合是否覆盖全部域名。
</p>
) : null}
</ResourceField>
<ToggleField
@@ -17,7 +17,7 @@ import {
parseOriginUrl,
parseOriginUrls,
validateDomains,
} from '@/features/proxy-routes/helpers';
} from '@/features/proxy-routes/helpers';buyao
import type { ProxyRouteItem } from '@/features/proxy-routes/types';
import {
PrimaryButton,
@@ -283,6 +283,7 @@ export function buildPayloadFromRoute(
enabled: route.enabled,
enable_https: route.enable_https,
cert_id: route.cert_id,
cert_ids: route.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,
@@ -316,4 +317,3 @@ export function getWebsiteStatusBadges(route: ProxyRouteItem) {
: { label: '缓存关闭', variant: 'warning' as const },
];
}
@@ -18,6 +18,7 @@ export interface ProxyRouteItem {
enabled: boolean;
enable_https: boolean;
cert_id: number | null;
cert_ids: number[];
redirect_http: boolean;
limit_conn_per_server: number;
limit_conn_per_ip: number;
@@ -48,6 +49,7 @@ export interface ProxyRouteMutationPayload {
enabled: boolean;
enable_https: boolean;
cert_id: number | null;
cert_ids?: number[];
redirect_http: boolean;
limit_conn_per_server?: number;
limit_conn_per_ip?: number;