mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 07:36:37 +08:00
[功能] 添加域名证书绑定支持,允许为每个域名单独选择证书并优化相关逻辑
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user