Files
OpenFlare/internal/apps/openflare/proxy_route/logics.go
T
ryan d615d85a26 refactor: remove unused remarks and add quick domain create
Drop remark fields from Zone, Zone domains, proxy routes, WAF rule
groups and IP groups across models, APIs, UI and DB columns (keep
certificate/origin remarks). Add quick-create domain input for short
labels, @ apex and full FQDNs when binding domains.
2026-07-12 16:16:56 +08:00

330 lines
12 KiB
Go

// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// CustomHeaderInput 自定义响应头。
type CustomHeaderInput struct {
Key string `json:"key"`
Value string `json:"value"`
}
// Input 代理规则创建/更新请求。
type Input struct {
SiteName string `json:"site_name"`
ZoneDomainIDs []uint `json:"zone_domain_ids"`
OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"`
OriginScheme string `json:"origin_scheme"`
OriginAddress string `json:"origin_address"`
OriginPort string `json:"origin_port"`
OriginURI string `json:"origin_uri"`
OriginHost string `json:"origin_host"`
Upstreams []string `json:"upstreams"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server"`
LimitConnPerIP int `json:"limit_conn_per_ip"`
LimitRate string `json:"limit_rate"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy"`
CacheRules []string `json:"cache_rules"`
CustomHeaders []CustomHeaderInput `json:"custom_headers"`
BasicAuthEnabled bool `json:"basic_auth_enabled"`
BasicAuthUsername string `json:"basic_auth_username"`
BasicAuthPassword string `json:"basic_auth_password"`
UpstreamType string `json:"upstream_type"`
TunnelNodeID *uint `json:"tunnel_node_id"`
TunnelID *uint `json:"tunnel_id"`
TunnelTargetAddr string `json:"tunnel_target_addr"`
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
PagesProjectID *uint `json:"pages_project_id"`
}
// View 代理规则视图。
type View struct {
ID uint `json:"id"`
SiteName string `json:"site_name"`
ZoneDomainIDs []uint `json:"zone_domain_ids"`
ZoneDomains []ZoneDomainView `json:"zone_domains"`
OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host"`
Upstreams string `json:"upstreams"`
UpstreamList []string `json:"upstream_list"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server"`
LimitConnPerIP int `json:"limit_conn_per_ip"`
LimitRate string `json:"limit_rate"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy"`
CacheRules string `json:"cache_rules"`
CacheRuleList []string `json:"cache_rule_list"`
CustomHeaders string `json:"custom_headers"`
CustomHeaderList []CustomHeaderInput `json:"custom_header_list"`
BasicAuthEnabled bool `json:"basic_auth_enabled"`
BasicAuthUsername string `json:"basic_auth_username"`
BasicAuthPassword string `json:"basic_auth_password"`
UpstreamType string `json:"upstream_type"`
TunnelNodeID *uint `json:"tunnel_node_id"`
TunnelID *uint `json:"tunnel_id"`
TunnelTargetAddr string `json:"tunnel_target_addr"`
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
PagesProjectID *uint `json:"pages_project_id"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// ZoneDomainView is the route-safe representation of a bound Zone domain.
type ZoneDomainView struct {
ID uint `json:"id"`
ZoneID uint `json:"zone_id"`
Domain string `json:"domain"`
CertID *uint `json:"cert_id"`
}
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*View, error) {
routes, err := model.ListProxyRoutes(ctx)
if err != nil {
return nil, err
}
return buildProxyRouteViews(ctx, routes)
}
// GetProxyRoute 获取代理规则详情。
func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
route, err := model.GetProxyRouteByID(ctx, id)
if err != nil {
return nil, err
}
return buildProxyRouteView(ctx, route)
}
// CreateProxyRoute 创建代理规则。
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
route, _, err := buildProxyRoute(ctx, nil, input)
if err != nil {
return nil, err
}
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Create(route).Error; err != nil {
return err
}
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs)
}); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errProxyRouteIdentityExists)
}
return nil, err
}
return buildProxyRouteView(ctx, route)
}
// UpdateProxyRoute 更新代理规则。
func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) {
route, err := model.GetProxyRouteByID(ctx, id)
if err != nil {
return nil, err
}
route, _, err = buildProxyRoute(ctx, route, input)
if err != nil {
return nil, err
}
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := updateProxyRouteRecord(tx, route); err != nil {
return err
}
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs)
}); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errProxyRouteIdentityExists)
}
return nil, err
}
return buildProxyRouteView(ctx, route)
}
// DeleteProxyRoute 删除代理规则。
func DeleteProxyRoute(ctx context.Context, id uint) error {
if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
return err
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", id).Update("proxy_route_id", nil).Error; err != nil {
return err
}
return tx.Delete(&model.ProxyRoute{}, id).Error
})
}
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) {
domains, err := loadProxyRouteZoneDomains(ctx, input.ZoneDomainIDs)
if err != nil {
return nil, nil, err
}
siteName := strings.TrimSpace(input.SiteName)
upstreamType := normalizeUpstreamType(input.UpstreamType)
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
if err != nil {
return nil, nil, err
}
originHost := strings.TrimSpace(input.OriginHost)
cachePolicy := strings.TrimSpace(input.CachePolicy)
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
if err != nil {
return nil, nil, err
}
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
if err != nil {
return nil, nil, err
}
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
if err != nil {
return nil, nil, err
}
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
if err != nil {
return nil, nil, err
}
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
if err != nil {
return nil, nil, err
}
if err := validateProxyRouteZoneDomainCertificates(ctx, domains, input.EnableHTTPS); err != nil {
return nil, nil, err
}
jsonFields, err := marshalProxyRouteJSONFields(upstreams, cacheRules, customHeaders)
if err != nil {
return nil, nil, err
}
if err := validateProxyRouteSiteName(siteName); err != nil {
return nil, nil, err
}
if err := validateProxyRouteSiteNameUniqueness(ctx, route, siteName); err != nil {
return nil, nil, err
}
if err := validateOriginHost(originHost); err != nil {
return nil, nil, err
}
if input.RedirectHTTP && !input.EnableHTTPS {
return nil, nil, errors.New(errProxyRouteRedirectHTTP)
}
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
return nil, nil, err
}
if route == nil {
route = &model.ProxyRoute{}
}
populateProxyRouteFields(
route,
input,
siteName,
jsonFields,
originID,
upstreams,
originHost,
cachePolicy,
limitConnPerServer,
limitConnPerIP,
limitRate,
upstreamType,
)
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
return nil, nil, err
}
return route, domains, nil
}
func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) {
views := make([]*View, 0, len(routes))
for _, route := range routes {
view, err := buildProxyRouteView(ctx, route)
if err != nil {
return nil, err
}
views = append(views, view)
}
return views, nil
}
func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, error) {
if route == nil {
return nil, errors.New("proxy route is nil")
}
domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
if err != nil {
return nil, err
}
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
if err != nil {
return nil, err
}
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
if err != nil {
return nil, err
}
zoneDomainIDs := make([]uint, 0, len(domains))
zoneDomains := make([]ZoneDomainView, 0, len(domains))
for _, domain := range domains {
zoneDomainIDs = append(zoneDomainIDs, domain.ID)
zoneDomains = append(zoneDomains, ZoneDomainView{ID: domain.ID, ZoneID: domain.ZoneID, Domain: domain.Domain, CertID: domain.CertID})
}
return &View{
ID: route.ID,
SiteName: route.SiteName,
ZoneDomainIDs: zoneDomainIDs,
ZoneDomains: zoneDomains,
OriginID: route.OriginID,
OriginURL: route.OriginURL,
OriginHost: route.OriginHost,
Upstreams: route.Upstreams,
UpstreamList: upstreams,
Enabled: route.Enabled,
EnableHTTPS: route.EnableHTTPS,
RedirectHTTP: route.RedirectHTTP,
LimitConnPerServer: route.LimitConnPerServer,
LimitConnPerIP: route.LimitConnPerIP,
LimitRate: route.LimitRate,
CacheEnabled: route.CacheEnabled,
CachePolicy: route.CachePolicy,
CacheRules: route.CacheRules,
CacheRuleList: cacheRules,
CustomHeaders: route.CustomHeaders,
CustomHeaderList: customHeaders,
BasicAuthEnabled: route.BasicAuthEnabled,
BasicAuthUsername: route.BasicAuthUsername,
BasicAuthPassword: route.BasicAuthPassword,
UpstreamType: route.UpstreamType,
TunnelNodeID: route.TunnelNodeID,
TunnelID: route.TunnelNodeID,
TunnelTargetAddr: route.TunnelTargetAddr,
TunnelTargetProtocol: route.TunnelTargetProtocol,
PagesProjectID: route.PagesProjectID,
CreatedAt: route.CreatedAt,
UpdatedAt: route.UpdatedAt,
}, nil
}