[功能] 支持为 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
+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")