mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
[功能] 添加支持多个上游地址,优化代理路由配置和负载均衡逻辑
This commit is contained in:
@@ -62,6 +62,7 @@ type snapshotRoute struct {
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginHost string `json:"origin_host,omitempty"`
|
||||
Upstreams []string `json:"upstreams,omitempty"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id,omitempty"`
|
||||
@@ -82,7 +83,7 @@ type routeCacheConfig struct {
|
||||
type routeUpstreamConfig struct {
|
||||
Name string
|
||||
Scheme string
|
||||
Address string
|
||||
Addresses []string
|
||||
UsesNamedUpstream bool
|
||||
}
|
||||
|
||||
@@ -408,6 +409,10 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
}
|
||||
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 上游配置无效", route.Domain)
|
||||
}
|
||||
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 缓存规则无效", route.Domain)
|
||||
@@ -416,6 +421,7 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
Domain: route.Domain,
|
||||
OriginURL: route.OriginURL,
|
||||
OriginHost: route.OriginHost,
|
||||
Upstreams: upstreams,
|
||||
Enabled: route.Enabled,
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: route.CertID,
|
||||
@@ -459,6 +465,11 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
|
||||
if err == nil {
|
||||
routes[index].CustomHeaders = normalizedHeaders
|
||||
}
|
||||
normalizedUpstreams, err := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams)
|
||||
if err == nil {
|
||||
routes[index].OriginURL = normalizedUpstreams[0]
|
||||
routes[index].Upstreams = normalizedUpstreams
|
||||
}
|
||||
normalizedCacheRules, err := normalizeCacheRules(routes[index].CacheEnabled, routes[index].CachePolicy, routes[index].CacheRules)
|
||||
if err == nil {
|
||||
routes[index].CachePolicy = normalizeCachePolicy(routes[index].CacheEnabled, routes[index].CachePolicy)
|
||||
@@ -472,6 +483,14 @@ func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
if left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || !uintPointerEqual(left.CertID, right.CertID) {
|
||||
return false
|
||||
}
|
||||
if len(left.Upstreams) != len(right.Upstreams) {
|
||||
return false
|
||||
}
|
||||
for index := range left.Upstreams {
|
||||
if left.Upstreams[index] != right.Upstreams[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.CacheRules) != len(right.CacheRules) {
|
||||
return false
|
||||
}
|
||||
@@ -648,6 +667,10 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
}
|
||||
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 上游配置无效", route.Domain)
|
||||
}
|
||||
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 缓存规则无效", route.Domain)
|
||||
@@ -657,7 +680,7 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
Policy: route.CachePolicy,
|
||||
Rules: cacheRules,
|
||||
}
|
||||
upstreamConfig := buildRouteUpstreamConfig(route, cfg)
|
||||
upstreamConfig := buildRouteUpstreamConfig(route, upstreams, cfg)
|
||||
if upstreamConfig.UsesNamedUpstream {
|
||||
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
|
||||
}
|
||||
@@ -965,24 +988,55 @@ func renderProxyPassBlock(originURL string, upstreamConfig routeUpstreamConfig,
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func buildRouteUpstreamConfig(route *model.ProxyRoute, cfg openRestyConfigSnapshot) routeUpstreamConfig {
|
||||
parsed, err := url.Parse(strings.TrimSpace(route.OriginURL))
|
||||
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
|
||||
func buildRouteUpstreamConfig(route *model.ProxyRoute, upstreams []string, cfg openRestyConfigSnapshot) routeUpstreamConfig {
|
||||
if len(upstreams) == 0 {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if shouldUseRuntimeResolver(route.OriginURL, cfg.Resolvers) {
|
||||
return routeUpstreamConfig{}
|
||||
if len(upstreams) == 1 {
|
||||
parsed, err := url.Parse(strings.TrimSpace(upstreams[0]))
|
||||
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if shouldUseRuntimeResolver(upstreams[0], cfg.Resolvers) {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if strings.TrimSpace(parsed.EscapedPath()) != "" && strings.TrimSpace(parsed.EscapedPath()) != "/" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if parsed.RawQuery != "" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
return routeUpstreamConfig{
|
||||
Name: buildRouteUpstreamName(route),
|
||||
Scheme: parsed.Scheme,
|
||||
Addresses: []string{parsed.Host},
|
||||
UsesNamedUpstream: true,
|
||||
}
|
||||
}
|
||||
if strings.TrimSpace(parsed.EscapedPath()) != "" && strings.TrimSpace(parsed.EscapedPath()) != "/" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if parsed.RawQuery != "" {
|
||||
return routeUpstreamConfig{}
|
||||
addresses := make([]string, 0, len(upstreams))
|
||||
var scheme string
|
||||
for _, upstream := range upstreams {
|
||||
parsed, err := url.Parse(strings.TrimSpace(upstream))
|
||||
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if strings.TrimSpace(parsed.EscapedPath()) != "" && strings.TrimSpace(parsed.EscapedPath()) != "/" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if parsed.RawQuery != "" {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
if scheme == "" {
|
||||
scheme = parsed.Scheme
|
||||
} else if scheme != parsed.Scheme {
|
||||
return routeUpstreamConfig{}
|
||||
}
|
||||
addresses = append(addresses, parsed.Host)
|
||||
}
|
||||
return routeUpstreamConfig{
|
||||
Name: buildRouteUpstreamName(route),
|
||||
Scheme: parsed.Scheme,
|
||||
Address: parsed.Host,
|
||||
Scheme: scheme,
|
||||
Addresses: addresses,
|
||||
UsesNamedUpstream: true,
|
||||
}
|
||||
}
|
||||
@@ -1008,7 +1062,13 @@ func buildRouteUpstreamName(route *model.ProxyRoute) string {
|
||||
}
|
||||
|
||||
func renderNamedUpstreamBlock(upstreamConfig routeUpstreamConfig) string {
|
||||
return fmt.Sprintf("upstream %s {\n server %s max_fails=3 fail_timeout=10s;\n keepalive 128;\n}\n\n", upstreamConfig.Name, upstreamConfig.Address)
|
||||
var builder strings.Builder
|
||||
builder.WriteString(fmt.Sprintf("upstream %s {\n", upstreamConfig.Name))
|
||||
for _, address := range upstreamConfig.Addresses {
|
||||
builder.WriteString(fmt.Sprintf(" server %s max_fails=3 fail_timeout=10s;\n", address))
|
||||
}
|
||||
builder.WriteString(" keepalive 128;\n}\n\n")
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func shouldUseRuntimeResolver(originURL string, resolvers string) bool {
|
||||
|
||||
@@ -72,6 +72,12 @@ func TestCreateTLSCertificateAndRenderHTTPSConfig(t *testing.T) {
|
||||
if !strings.Contains(result.Version.MainConfig, "multi_accept on;") {
|
||||
t.Fatal("expected main config to default multi_accept to on")
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, "keepalive_timeout 20;") {
|
||||
t.Fatal("expected main config to default keepalive_timeout to 20")
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, "proxy_connect_timeout 3;") {
|
||||
t.Fatal("expected main config to default proxy_connect_timeout to 3")
|
||||
}
|
||||
if strings.Contains(result.Version.MainConfig, "allow 127.0.0.1;") {
|
||||
t.Fatal("expected main config to avoid hard-coded allow rules on observability server")
|
||||
}
|
||||
@@ -212,6 +218,9 @@ func TestPublishConfigVersionRendersRouteLevelCachePolicy(t *testing.T) {
|
||||
if !strings.Contains(result.Version.MainConfig, "proxy_cache_path /var/cache/openresty/openflare") {
|
||||
t.Fatal("expected main config to include cache zone when cache infra is enabled")
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, `proxy_cache_key "$scheme$host$request_uri";`) {
|
||||
t.Fatal("expected main config to default cache key to host dimension")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "proxy_cache_methods GET;") {
|
||||
t.Fatal("expected rendered config to only cache GET requests")
|
||||
}
|
||||
@@ -244,6 +253,50 @@ func TestPublishConfigVersionRendersRouteLevelCachePolicy(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionRendersMultipleUpstreams(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "lb.example.com",
|
||||
OriginURL: "http://c1:39010",
|
||||
Upstreams: []string{"http://c2:39010", "http://c3:39010"},
|
||||
Enabled: true,
|
||||
OriginHost: "lb.example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(route.Upstreams, "c2:39010") {
|
||||
t.Fatalf("expected route upstreams to persist, got %s", route.Upstreams)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root")
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "upstream backend_lb_example_com_1 {") {
|
||||
t.Fatal("expected rendered config to define upstream block for load balancing route")
|
||||
}
|
||||
if strings.Count(result.Version.RenderedConfig, "server c") < 3 {
|
||||
t.Fatal("expected rendered config to include every upstream server")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "server c1:39010 max_fails=3 fail_timeout=10s;") {
|
||||
t.Fatal("expected rendered config to include primary upstream server")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "server c2:39010 max_fails=3 fail_timeout=10s;") {
|
||||
t.Fatal("expected rendered config to include secondary upstream server")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "server c3:39010 max_fails=3 fail_timeout=10s;") {
|
||||
t.Fatal("expected rendered config to include tertiary upstream server")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "proxy_pass http://backend_lb_example_com_1;") {
|
||||
t.Fatal("expected rendered config to proxy through load balancing upstream")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"upstreams":["http://c1:39010","http://c2:39010","http://c3:39010"]`) {
|
||||
t.Fatal("expected snapshot to include upstream list")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionOverridesOriginHostHeader(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
|
||||
@@ -27,6 +27,7 @@ type ProxyRouteInput struct {
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginHost string `json:"origin_host"`
|
||||
Upstreams []string `json:"upstreams"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
@@ -87,6 +88,10 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
originURL := strings.TrimSpace(input.OriginURL)
|
||||
originHost := strings.TrimSpace(input.OriginHost)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
upstreams, err := normalizeUpstreams(originURL, input.Upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cachePolicy := strings.TrimSpace(input.CachePolicy)
|
||||
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
|
||||
if err != nil {
|
||||
@@ -100,6 +105,10 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
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
|
||||
@@ -110,9 +119,6 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, errors.New("域名格式不合法")
|
||||
}
|
||||
if err := validateOriginURL(originURL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateOriginHost(originHost); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -135,8 +141,9 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.Domain = domain
|
||||
route.OriginURL = originURL
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = string(upstreamsJSON)
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
@@ -177,6 +184,59 @@ func normalizeCustomHeaders(headers []ProxyRouteCustomHeaderInput) ([]ProxyRoute
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeUpstreams(originURL string, upstreams []string) ([]string, error) {
|
||||
candidates := make([]string, 0, len(upstreams)+1)
|
||||
if strings.TrimSpace(originURL) != "" {
|
||||
candidates = append(candidates, originURL)
|
||||
}
|
||||
candidates = append(candidates, upstreams...)
|
||||
trimmed := make([]string, 0, len(candidates))
|
||||
for _, candidate := range candidates {
|
||||
item := strings.TrimSpace(candidate)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
trimmed = append(trimmed, item)
|
||||
}
|
||||
unique := make([]string, 0, len(trimmed))
|
||||
seen := make(map[string]struct{}, len(trimmed))
|
||||
for _, item := range trimmed {
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
unique = append(unique, item)
|
||||
}
|
||||
normalized := make([]string, 0, len(unique))
|
||||
var scheme string
|
||||
multiUpstream := len(unique) > 1
|
||||
for _, item := range unique {
|
||||
if err := validateOriginURL(item); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parsed, err := url.ParseRequestURI(item)
|
||||
if err != nil {
|
||||
return nil, errors.New("源站地址格式不合法")
|
||||
}
|
||||
if multiUpstream && parsed.Path != "" && parsed.Path != "/" {
|
||||
return nil, errors.New("多上游模式暂不支持带路径的源站地址")
|
||||
}
|
||||
if multiUpstream && parsed.RawQuery != "" {
|
||||
return nil, errors.New("多上游模式暂不支持带查询参数的源站地址")
|
||||
}
|
||||
if scheme == "" {
|
||||
scheme = parsed.Scheme
|
||||
} else if scheme != parsed.Scheme {
|
||||
return nil, errors.New("同一规则的多个上游必须使用相同协议")
|
||||
}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, errors.New("至少填写一个上游地址")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeStoredCustomHeaders(raw string) ([]ProxyRouteCustomHeaderInput, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
@@ -291,6 +351,18 @@ func decodeStoredCacheRules(raw string) ([]string, error) {
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return normalizeUpstreams(fallbackOriginURL, nil)
|
||||
}
|
||||
var upstreams []string
|
||||
if err := json.Unmarshal([]byte(text), &upstreams); err != nil {
|
||||
return nil, errors.New("上游配置格式不合法")
|
||||
}
|
||||
return normalizeUpstreams(fallbackOriginURL, upstreams)
|
||||
}
|
||||
|
||||
func validateOriginURL(raw string) error {
|
||||
if raw == "" {
|
||||
return errors.New("源站地址不能为空")
|
||||
|
||||
Reference in New Issue
Block a user