[功能] 添加域名证书绑定支持,允许为每个域名单独选择证书并优化相关逻辑

This commit is contained in:
ryan
2026-04-01 09:57:40 +08:00
parent a002d98f3a
commit 49472b54bf
17 changed files with 873 additions and 69 deletions
@@ -4,7 +4,7 @@ import "time"
const (
legacyDatabaseSchemaVersion = 1
currentDatabaseSchemaVersion = 7
currentDatabaseSchemaVersion = 8
databaseSchemaVersionRowID = 1
)
+104
View File
@@ -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")
+300 -1
View File
@@ -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)
+2
View File
@@ -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,