mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
fix lint
This commit is contained in:
@@ -0,0 +1,178 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package proxy_route provides helpers for building proxy route configurations.
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
type proxyRouteJSONFields struct {
|
||||
cacheRulesJSON string
|
||||
upstreamsJSON string
|
||||
customHeadersJSON string
|
||||
certIDsJSON string
|
||||
domainCertIDsJSON string
|
||||
domainsJSON string
|
||||
}
|
||||
|
||||
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) {
|
||||
switch upstreamType {
|
||||
case proxyRouteUpstreamTypeTunnel, proxyRouteUpstreamTypePages:
|
||||
if upstreamType == proxyRouteUpstreamTypePages {
|
||||
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
|
||||
return "", nil, nil, err
|
||||
}
|
||||
}
|
||||
originURL := "http://127.0.0.1"
|
||||
return originURL, nil, []string{originURL}, nil
|
||||
default:
|
||||
originURL, originID, err := resolveProxyRoutePrimaryOrigin(ctx, input)
|
||||
if err != nil {
|
||||
return "", nil, nil, err
|
||||
}
|
||||
upstreams, err := normalizeUpstreams(originURL, input.Upstreams)
|
||||
if err != nil {
|
||||
return "", nil, nil, err
|
||||
}
|
||||
return originURL, originID, upstreams, nil
|
||||
}
|
||||
}
|
||||
|
||||
func marshalProxyRouteJSONFields(
|
||||
domains []string,
|
||||
upstreams []string,
|
||||
cacheRules []string,
|
||||
customHeaders []CustomHeaderInput,
|
||||
certIDs []uint,
|
||||
domainCertIDs []uint,
|
||||
) (*proxyRouteJSONFields, error) {
|
||||
cacheRulesJSON, err := json.Marshal(cacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeadersJSON, err := json.Marshal(customHeaders)
|
||||
if 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
|
||||
}
|
||||
return &proxyRouteJSONFields{
|
||||
cacheRulesJSON: string(cacheRulesJSON),
|
||||
upstreamsJSON: string(upstreamsJSON),
|
||||
customHeadersJSON: string(customHeadersJSON),
|
||||
certIDsJSON: string(certIDsJSON),
|
||||
domainCertIDsJSON: string(domainCertIDsJSON),
|
||||
domainsJSON: string(domainsJSON),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteHTTPSInput(input *Input) {
|
||||
if input.EnableHTTPS {
|
||||
return
|
||||
}
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteBasicAuth(input *Input) error {
|
||||
if !input.BasicAuthEnabled {
|
||||
input.BasicAuthUsername = ""
|
||||
input.BasicAuthPassword = ""
|
||||
return nil
|
||||
}
|
||||
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
|
||||
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
|
||||
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
|
||||
return errors.New(errProxyRouteBasicAuth)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func populateProxyRouteFields(
|
||||
route *model.ProxyRoute,
|
||||
input Input,
|
||||
siteName, domain string,
|
||||
jsonFields *proxyRouteJSONFields,
|
||||
originID *uint,
|
||||
upstreams []string,
|
||||
originHost, remark, cachePolicy string,
|
||||
limitConnPerServer, limitConnPerIP int,
|
||||
limitRate, upstreamType string,
|
||||
) {
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = jsonFields.domainsJSON
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = jsonFields.upstreamsJSON
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.CertIDs = jsonFields.certIDsJSON
|
||||
route.DomainCertIDs = jsonFields.domainCertIDsJSON
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.LimitConnPerServer = limitConnPerServer
|
||||
route.LimitConnPerIP = limitConnPerIP
|
||||
route.LimitRate = limitRate
|
||||
route.CacheEnabled = input.CacheEnabled
|
||||
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
|
||||
route.CacheRules = jsonFields.cacheRulesJSON
|
||||
route.CustomHeaders = jsonFields.customHeadersJSON
|
||||
route.BasicAuthEnabled = input.BasicAuthEnabled
|
||||
route.BasicAuthUsername = input.BasicAuthUsername
|
||||
route.BasicAuthPassword = input.BasicAuthPassword
|
||||
route.Remark = remark
|
||||
route.UpstreamType = upstreamType
|
||||
}
|
||||
|
||||
func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, upstreamType string, input Input) error {
|
||||
switch upstreamType {
|
||||
case proxyRouteUpstreamTypeTunnel:
|
||||
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
|
||||
return err
|
||||
}
|
||||
route.TunnelNodeID = tunnelNodeID
|
||||
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
|
||||
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
|
||||
route.PagesProjectID = nil
|
||||
case proxyRouteUpstreamTypePages:
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = input.PagesProjectID
|
||||
default:
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
func normalizeExplicitDomainCertIDs(ctx context.Context, domains []string, rawDomainCertIDs []uint) ([]uint, []uint, *uint, error) {
|
||||
if len(rawDomainCertIDs) != len(domains) {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
|
||||
}
|
||||
|
||||
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 := lookupTLSCertificateByID(ctx, item); err != nil {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
|
||||
}
|
||||
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(errProxyRouteCertRequired)
|
||||
}
|
||||
|
||||
primaryCertID := &uniqueCertIDs[0]
|
||||
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
|
||||
}
|
||||
|
||||
func normalizeDerivedDomainCertIDs(
|
||||
ctx context.Context,
|
||||
domains []string,
|
||||
normalizedCertIDs []uint,
|
||||
) ([]uint, []uint, *uint, error) {
|
||||
switch {
|
||||
case len(normalizedCertIDs) == 0:
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
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(ctx, domains, normalizedCertIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
}
|
||||
}
|
||||
@@ -42,9 +42,9 @@ const (
|
||||
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
|
||||
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
|
||||
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
|
||||
errProxyRoutePagesNotFound = "Pages 项目不存在"
|
||||
errProxyRoutePagesDisabled = "Pages 项目未启用"
|
||||
errProxyRoutePagesNoDeploy = "Pages 项目没有激活部署"
|
||||
errProxyRoutePagesNotFound = "pages 项目不存在"
|
||||
errProxyRoutePagesDisabled = "pages 项目未启用"
|
||||
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
|
||||
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
|
||||
errProxyRouteOriginPort = "端口格式不合法"
|
||||
errProxyRouteOriginPortEmpty = "端口不能为空"
|
||||
|
||||
@@ -30,6 +30,13 @@ const (
|
||||
proxyRouteCachePolicySuffix = "suffix"
|
||||
proxyRouteCachePolicyPathPrefix = "path_prefix"
|
||||
proxyRouteCachePolicyPathExact = "path_exact"
|
||||
proxyRouteSchemeHTTP = "http"
|
||||
proxyRouteSchemeHTTPS = "https"
|
||||
proxyRouteUpstreamTypeTunnel = "tunnel"
|
||||
proxyRouteUpstreamTypePages = "pages"
|
||||
|
||||
maxOriginHostnameLength = 253
|
||||
originURIPathQueryParts = 2
|
||||
)
|
||||
|
||||
type tlsCertificateRow struct {
|
||||
@@ -100,7 +107,7 @@ func validateOriginAddress(address string) error {
|
||||
if ip := net.ParseIP(address); ip != nil {
|
||||
return nil
|
||||
}
|
||||
if len(address) > 253 {
|
||||
if len(address) > maxOriginHostnameLength {
|
||||
return errors.New(errProxyRouteOriginInvalid)
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
@@ -136,7 +143,7 @@ func normalizeOriginPort(raw string) (string, error) {
|
||||
func normalizeOriginScheme(raw string) (string, error) {
|
||||
scheme := strings.ToLower(strings.TrimSpace(raw))
|
||||
switch scheme {
|
||||
case "http", "https":
|
||||
case proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
|
||||
return scheme, nil
|
||||
default:
|
||||
return "", errors.New(errProxyRouteOriginSchemeOnly)
|
||||
@@ -187,7 +194,7 @@ func buildOriginURLFromParts(scheme, address, port, uri string) (string, error)
|
||||
if strings.HasPrefix(normalizedURI, "?") {
|
||||
parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?")
|
||||
} else {
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", 2)
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", originURIPathQueryParts)
|
||||
parsed.Path = pathQuery[0]
|
||||
if len(pathQuery) > 1 {
|
||||
parsed.RawQuery = pathQuery[1]
|
||||
@@ -465,65 +472,14 @@ func normalizeProxyRouteDomainCertificateIDs(
|
||||
}
|
||||
|
||||
if len(rawDomainCertIDs) > 0 {
|
||||
if len(rawDomainCertIDs) != len(domains) {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
|
||||
}
|
||||
|
||||
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 := lookupTLSCertificateByID(ctx, item); err != nil {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
|
||||
}
|
||||
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(errProxyRouteCertRequired)
|
||||
}
|
||||
|
||||
primaryCertID := &uniqueCertIDs[0]
|
||||
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
|
||||
return normalizeExplicitDomainCertIDs(ctx, domains, rawDomainCertIDs)
|
||||
}
|
||||
|
||||
normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(ctx, enableHTTPS, certID, certIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
switch {
|
||||
case len(normalizedCertIDs) == 0:
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
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(ctx, domains, normalizedCertIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
}
|
||||
return normalizeDerivedDomainCertIDs(ctx, domains, normalizedCertIDs)
|
||||
}
|
||||
|
||||
func validateProxyRouteDomainCertificateCoverage(ctx context.Context, domains []string, domainCertIDs []uint) error {
|
||||
@@ -637,65 +593,6 @@ func hasStructuredOriginInput(input Input) bool {
|
||||
strings.TrimSpace(input.OriginURI) != ""
|
||||
}
|
||||
|
||||
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
|
||||
if hasStructuredOriginInput(input) {
|
||||
scheme, err := normalizeOriginScheme(input.OriginScheme)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
port, err := normalizeOriginPort(input.OriginPort)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
uri, err := normalizeOriginURI(input.OriginURI)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
if input.OriginID != nil && *input.OriginID != 0 {
|
||||
origin, err := model.GetOriginByID(ctx, *input.OriginID)
|
||||
if err != nil {
|
||||
return "", nil, errors.New(errProxyRouteOriginNotFound)
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
address := normalizeOriginAddress(input.OriginAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, err := getOrCreateOriginByAddress(ctx, address)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
originURL := strings.TrimSpace(input.OriginURL)
|
||||
if originURL == "" {
|
||||
return "", nil, errors.New(errProxyRouteOriginEmpty)
|
||||
}
|
||||
address, err := extractOriginAddress(originURL)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, findErr := model.GetOriginByAddress(ctx, address)
|
||||
if findErr == nil {
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
|
||||
return "", nil, findErr
|
||||
}
|
||||
return originURL, nil, nil
|
||||
}
|
||||
|
||||
func normalizeCustomHeaders(headers []CustomHeaderInput) ([]CustomHeaderInput, error) {
|
||||
if len(headers) == 0 {
|
||||
return []CustomHeaderInput{}, nil
|
||||
@@ -942,7 +839,7 @@ func validateOriginURL(raw string) error {
|
||||
if err != nil {
|
||||
return errors.New(errProxyRouteOriginInvalid)
|
||||
}
|
||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
||||
if parsed.Scheme != proxyRouteSchemeHTTP && parsed.Scheme != proxyRouteSchemeHTTPS {
|
||||
return errors.New(errProxyRouteOriginScheme)
|
||||
}
|
||||
if parsed.Host == "" {
|
||||
@@ -996,7 +893,7 @@ func validateTunnelRouteInput(ctx context.Context, tunnelNodeID *uint, targetAdd
|
||||
return errors.New(errProxyRouteTunnelAddrReq)
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(targetProtocol)) {
|
||||
case "", "http", "https":
|
||||
case "", proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
|
||||
return nil
|
||||
default:
|
||||
return errors.New(errProxyRouteTunnelProtocol)
|
||||
@@ -1025,10 +922,10 @@ func validatePagesRouteInput(ctx context.Context, projectID *uint) error {
|
||||
|
||||
func normalizeUpstreamType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "tunnel":
|
||||
return "tunnel"
|
||||
case "pages":
|
||||
return "pages"
|
||||
case proxyRouteUpstreamTypeTunnel:
|
||||
return proxyRouteUpstreamTypeTunnel
|
||||
case proxyRouteUpstreamTypePages:
|
||||
return proxyRouteUpstreamTypePages
|
||||
default:
|
||||
return "direct"
|
||||
}
|
||||
@@ -1036,9 +933,9 @@ func normalizeUpstreamType(raw string) string {
|
||||
|
||||
func normalizeTunnelTargetProtocol(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "https":
|
||||
return "https"
|
||||
case proxyRouteSchemeHTTPS:
|
||||
return proxyRouteSchemeHTTPS
|
||||
default:
|
||||
return "http"
|
||||
return proxyRouteSchemeHTTP
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -168,28 +167,9 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain)
|
||||
|
||||
upstreamType := normalizeUpstreamType(input.UpstreamType)
|
||||
var originURL string
|
||||
var originID *uint
|
||||
var upstreams []string
|
||||
|
||||
if upstreamType == "tunnel" {
|
||||
originURL = "http://127.0.0.1"
|
||||
upstreams = []string{originURL}
|
||||
} else if upstreamType == "pages" {
|
||||
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
originURL = "http://127.0.0.1"
|
||||
upstreams = []string{originURL}
|
||||
} else {
|
||||
originURL, originID, err = resolveProxyRoutePrimaryOrigin(ctx, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreams, err = normalizeUpstreams(originURL, input.Upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
originHost := strings.TrimSpace(input.OriginHost)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
@@ -215,25 +195,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cacheRulesJSON, err := json.Marshal(cacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeadersJSON, err := json.Marshal(customHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !input.EnableHTTPS {
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
normalizeProxyRouteHTTPSInput(&input)
|
||||
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
|
||||
ctx,
|
||||
domains,
|
||||
@@ -248,15 +210,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
if err := validateProxyRouteDomainCertificateCoverage(ctx, 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)
|
||||
jsonFields, err := marshalProxyRouteJSONFields(domains, upstreams, cacheRules, customHeaders, certIDs, domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -277,67 +231,31 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
return nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
}
|
||||
|
||||
if input.BasicAuthEnabled {
|
||||
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
|
||||
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
|
||||
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
|
||||
return nil, errors.New(errProxyRouteBasicAuth)
|
||||
}
|
||||
} else {
|
||||
input.BasicAuthUsername = ""
|
||||
input.BasicAuthPassword = ""
|
||||
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if route == nil {
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = string(domainsJSON)
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = string(upstreamsJSON)
|
||||
route.Enabled = input.Enabled
|
||||
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
|
||||
route.LimitRate = limitRate
|
||||
route.CacheEnabled = input.CacheEnabled
|
||||
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
|
||||
route.CacheRules = string(cacheRulesJSON)
|
||||
route.CustomHeaders = string(customHeadersJSON)
|
||||
route.BasicAuthEnabled = input.BasicAuthEnabled
|
||||
route.BasicAuthUsername = input.BasicAuthUsername
|
||||
route.BasicAuthPassword = input.BasicAuthPassword
|
||||
route.Remark = remark
|
||||
route.UpstreamType = upstreamType
|
||||
if upstreamType == "tunnel" {
|
||||
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
route.TunnelNodeID = tunnelNodeID
|
||||
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
|
||||
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
|
||||
route.PagesProjectID = nil
|
||||
} else if upstreamType == "pages" {
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = input.PagesProjectID
|
||||
} else {
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = nil
|
||||
populateProxyRouteFields(
|
||||
route,
|
||||
input,
|
||||
siteName,
|
||||
domain,
|
||||
jsonFields,
|
||||
originID,
|
||||
upstreams,
|
||||
originHost,
|
||||
remark,
|
||||
cachePolicy,
|
||||
limitConnPerServer,
|
||||
limitConnPerIP,
|
||||
limitRate,
|
||||
upstreamType,
|
||||
)
|
||||
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return route, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func resolveStructuredOriginInput(ctx context.Context, input Input) (string, *uint, error) {
|
||||
scheme, err := normalizeOriginScheme(input.OriginScheme)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
port, err := normalizeOriginPort(input.OriginPort)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
uri, err := normalizeOriginURI(input.OriginURI)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
if input.OriginID != nil && *input.OriginID != 0 {
|
||||
return resolveOriginByID(ctx, scheme, port, uri, *input.OriginID)
|
||||
}
|
||||
return resolveOriginByAddress(ctx, scheme, port, uri, input.OriginAddress)
|
||||
}
|
||||
|
||||
func resolveOriginByID(ctx context.Context, scheme, port, uri string, originID uint) (string, *uint, error) {
|
||||
origin, err := model.GetOriginByID(ctx, originID)
|
||||
if err != nil {
|
||||
return "", nil, errors.New(errProxyRouteOriginNotFound)
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
func resolveOriginByAddress(ctx context.Context, scheme, port, uri, rawAddress string) (string, *uint, error) {
|
||||
address := normalizeOriginAddress(rawAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, err := getOrCreateOriginByAddress(ctx, address)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
func resolveLegacyOriginInput(ctx context.Context, originURL string) (string, *uint, error) {
|
||||
if originURL == "" {
|
||||
return "", nil, errors.New(errProxyRouteOriginEmpty)
|
||||
}
|
||||
address, err := extractOriginAddress(originURL)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, findErr := model.GetOriginByAddress(ctx, address)
|
||||
if findErr == nil {
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
|
||||
return "", nil, findErr
|
||||
}
|
||||
return originURL, nil, nil
|
||||
}
|
||||
|
||||
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
|
||||
if hasStructuredOriginInput(input) {
|
||||
return resolveStructuredOriginInput(ctx, input)
|
||||
}
|
||||
return resolveLegacyOriginInput(ctx, strings.TrimSpace(input.OriginURL))
|
||||
}
|
||||
Reference in New Issue
Block a user