mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 16:16:37 +08:00
[功能] 添加源站管理功能,包括源站的创建、更新、删除及列表展示
This commit is contained in:
@@ -0,0 +1,241 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type OriginInput struct {
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
type OriginRouteSummary struct {
|
||||
ID uint `json:"id"`
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
type OriginView struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
Remark string `json:"remark"`
|
||||
RouteCount int64 `json:"route_count"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type OriginDetailView struct {
|
||||
OriginView
|
||||
Routes []OriginRouteSummary `json:"routes"`
|
||||
}
|
||||
|
||||
func ListOrigins() ([]OriginView, error) {
|
||||
origins, err := model.ListOrigins()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildOriginViews(origins)
|
||||
}
|
||||
|
||||
func GetOriginDetail(id uint) (*OriginDetailView, error) {
|
||||
origin, err := model.GetOriginByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views, err := buildOriginViews([]*model.Origin{origin})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
routes, err := model.ListProxyRoutesByOriginID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]OriginRouteSummary, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
items = append(items, OriginRouteSummary{
|
||||
ID: route.ID,
|
||||
Domain: route.Domain,
|
||||
OriginURL: route.OriginURL,
|
||||
Enabled: route.Enabled,
|
||||
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
return items[i].Domain < items[j].Domain
|
||||
})
|
||||
detail := &OriginDetailView{
|
||||
OriginView: views[0],
|
||||
Routes: items,
|
||||
}
|
||||
return detail, nil
|
||||
}
|
||||
|
||||
func CreateOrigin(input OriginInput) (*model.Origin, error) {
|
||||
origin, err := buildOrigin(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = origin.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("源站地址已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return origin, nil
|
||||
}
|
||||
|
||||
func UpdateOrigin(id uint, input OriginInput) (*model.Origin, error) {
|
||||
origin, err := model.GetOriginByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
previousAddress := origin.Address
|
||||
nextOrigin, err := buildOrigin(origin, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Save(nextOrigin).Error; err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return errors.New("源站地址已存在")
|
||||
}
|
||||
return err
|
||||
}
|
||||
if previousAddress == nextOrigin.Address {
|
||||
return nil
|
||||
}
|
||||
return updateRoutesForOriginAddress(tx, nextOrigin.ID, nextOrigin.Address)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nextOrigin, nil
|
||||
}
|
||||
|
||||
func DeleteOrigin(id uint) error {
|
||||
routes, err := model.ListProxyRoutesByOriginID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(routes) > 0 {
|
||||
return errors.New("该源站仍被规则引用,无法删除")
|
||||
}
|
||||
origin, err := model.GetOriginByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return origin.Delete()
|
||||
}
|
||||
|
||||
func buildOrigin(existing *model.Origin, input OriginInput) (*model.Origin, error) {
|
||||
address := normalizeOriginAddress(input.Address)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.Origin{}
|
||||
}
|
||||
existing.Address = address
|
||||
existing.Name = normalizeOriginName(input.Name, address)
|
||||
existing.Remark = strings.TrimSpace(input.Remark)
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
func getOrCreateOriginByAddress(address string) (*model.Origin, error) {
|
||||
normalizedAddress := normalizeOriginAddress(address)
|
||||
if err := validateOriginAddress(normalizedAddress); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
existing, err := model.GetOriginByAddress(normalizedAddress)
|
||||
if err == nil {
|
||||
return existing, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
origin := &model.Origin{
|
||||
Name: normalizedAddress,
|
||||
Address: normalizedAddress,
|
||||
Remark: "",
|
||||
}
|
||||
if err := origin.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return model.GetOriginByAddress(normalizedAddress)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return origin, nil
|
||||
}
|
||||
|
||||
func updateRoutesForOriginAddress(tx *gorm.DB, originID uint, address string) error {
|
||||
var routes []*model.ProxyRoute
|
||||
if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil {
|
||||
return fmt.Errorf("query routes for origin update failed: %w", err)
|
||||
}
|
||||
for _, route := range routes {
|
||||
rewrittenOriginURL, err := rewriteOriginURLAddress(route.OriginURL, address)
|
||||
if err != nil {
|
||||
return fmt.Errorf("rewrite route %d origin failed: %w", route.ID, err)
|
||||
}
|
||||
upstreams := make([]string, 0)
|
||||
if strings.TrimSpace(route.Upstreams) != "" {
|
||||
if err := json.Unmarshal([]byte(route.Upstreams), &upstreams); err != nil {
|
||||
return fmt.Errorf("decode route %d upstreams failed: %w", route.ID, err)
|
||||
}
|
||||
}
|
||||
if len(upstreams) == 0 {
|
||||
upstreams = append(upstreams, rewrittenOriginURL)
|
||||
} else {
|
||||
upstreams[0] = rewrittenOriginURL
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err)
|
||||
}
|
||||
if err := tx.Model(&model.ProxyRoute{}).
|
||||
Where("id = ?", route.ID).
|
||||
Updates(map[string]any{
|
||||
"origin_url": rewrittenOriginURL,
|
||||
"upstreams": string(upstreamsJSON),
|
||||
}).Error; err != nil {
|
||||
return fmt.Errorf("update route %d origin address failed: %w", route.ID, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildOriginViews(origins []*model.Origin) ([]OriginView, error) {
|
||||
countRows, err := model.ListOriginRouteCounts()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
countMap := make(map[uint]int64, len(countRows))
|
||||
for _, row := range countRows {
|
||||
countMap[row.OriginID] = row.RouteCount
|
||||
}
|
||||
views := make([]OriginView, 0, len(origins))
|
||||
for _, origin := range origins {
|
||||
views = append(views, OriginView{
|
||||
ID: origin.ID,
|
||||
Name: origin.Name,
|
||||
Address: origin.Address,
|
||||
Remark: origin.Remark,
|
||||
RouteCount: countMap[origin.ID],
|
||||
CreatedAt: origin.CreatedAt,
|
||||
UpdatedAt: origin.UpdatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
func normalizeOriginAddress(raw string) string {
|
||||
return strings.ToLower(strings.TrimSpace(raw))
|
||||
}
|
||||
|
||||
func validateOriginAddress(address string) error {
|
||||
if address == "" {
|
||||
return errors.New("源站地址不能为空")
|
||||
}
|
||||
if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") {
|
||||
return errors.New("源站地址无需包含 IPv6 方括号")
|
||||
}
|
||||
if ip := net.ParseIP(address); ip != nil {
|
||||
return nil
|
||||
}
|
||||
if len(address) > 253 {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
for _, label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
if label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
for _, r := range label {
|
||||
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
|
||||
continue
|
||||
}
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeOriginName(name string, address string) string {
|
||||
normalized := strings.TrimSpace(name)
|
||||
if normalized != "" {
|
||||
return normalized
|
||||
}
|
||||
return address
|
||||
}
|
||||
|
||||
func normalizeOriginPort(raw string) (string, error) {
|
||||
port := strings.TrimSpace(raw)
|
||||
if port == "" {
|
||||
return "", errors.New("端口不能为空")
|
||||
}
|
||||
value, err := strconv.Atoi(port)
|
||||
if err != nil || value < 1 || value > 65535 {
|
||||
return "", errors.New("端口格式不合法")
|
||||
}
|
||||
return strconv.Itoa(value), nil
|
||||
}
|
||||
|
||||
func normalizeOriginScheme(raw string) (string, error) {
|
||||
scheme := strings.ToLower(strings.TrimSpace(raw))
|
||||
switch scheme {
|
||||
case "http", "https":
|
||||
return scheme, nil
|
||||
default:
|
||||
return "", errors.New("源站协议仅支持 http 或 https")
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeOriginURI(raw string) (string, error) {
|
||||
uri := strings.TrimSpace(raw)
|
||||
if uri == "" {
|
||||
return "", nil
|
||||
}
|
||||
if strings.Contains(uri, "://") {
|
||||
return "", errors.New("源站路径不能包含协议")
|
||||
}
|
||||
if !strings.HasPrefix(uri, "/") && !strings.HasPrefix(uri, "?") {
|
||||
return "", errors.New("源站路径需以 / 或 ? 开头")
|
||||
}
|
||||
return uri, nil
|
||||
}
|
||||
|
||||
func formatOriginHost(address string, port string) string {
|
||||
if ip := net.ParseIP(address); ip != nil && strings.Contains(address, ":") {
|
||||
return net.JoinHostPort(address, port)
|
||||
}
|
||||
return net.JoinHostPort(address, port)
|
||||
}
|
||||
|
||||
func buildOriginURLFromParts(
|
||||
scheme string,
|
||||
address string,
|
||||
port string,
|
||||
uri string,
|
||||
) (string, error) {
|
||||
normalizedScheme, err := normalizeOriginScheme(scheme)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
normalizedAddress := normalizeOriginAddress(address)
|
||||
if err := validateOriginAddress(normalizedAddress); err != nil {
|
||||
return "", err
|
||||
}
|
||||
normalizedPort, err := normalizeOriginPort(port)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
normalizedURI, err := normalizeOriginURI(uri)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
parsed := &url.URL{
|
||||
Scheme: normalizedScheme,
|
||||
Host: formatOriginHost(normalizedAddress, normalizedPort),
|
||||
}
|
||||
if normalizedURI != "" {
|
||||
if strings.HasPrefix(normalizedURI, "?") {
|
||||
parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?")
|
||||
} else {
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", 2)
|
||||
parsed.Path = pathQuery[0]
|
||||
if len(pathQuery) > 1 {
|
||||
parsed.RawQuery = pathQuery[1]
|
||||
}
|
||||
}
|
||||
}
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
func extractOriginAddress(rawURL string) (string, error) {
|
||||
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("源站地址格式不合法: %w", err)
|
||||
}
|
||||
address := normalizeOriginAddress(parsed.Hostname())
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return address, nil
|
||||
}
|
||||
|
||||
func rewriteOriginURLAddress(rawURL string, newAddress string) (string, error) {
|
||||
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("源站地址格式不合法: %w", err)
|
||||
}
|
||||
address := normalizeOriginAddress(newAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", err
|
||||
}
|
||||
port := parsed.Port()
|
||||
if port == "" {
|
||||
return "", errors.New("源站地址缺少端口")
|
||||
}
|
||||
parsed.Host = formatOriginHost(address, port)
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
func splitOriginURL(rawURL string) (scheme string, address string, port string, uri string, err error) {
|
||||
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
|
||||
if err != nil {
|
||||
return "", "", "", "", err
|
||||
}
|
||||
scheme = parsed.Scheme
|
||||
address = normalizeOriginAddress(parsed.Hostname())
|
||||
port = parsed.Port()
|
||||
uri = parsed.EscapedPath()
|
||||
if uri == "" {
|
||||
uri = parsed.Path
|
||||
}
|
||||
if parsed.RawQuery != "" {
|
||||
if uri == "" {
|
||||
uri = "?" + parsed.RawQuery
|
||||
} else {
|
||||
uri = uri + "?" + parsed.RawQuery
|
||||
}
|
||||
}
|
||||
return scheme, address, port, uri, nil
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"openflare/model"
|
||||
)
|
||||
|
||||
func TestCreateProxyRouteStructuredOriginAutoCreatesOrigin(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginScheme: "https",
|
||||
OriginAddress: "origin.internal",
|
||||
OriginPort: "8443",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if route.OriginID == nil || *route.OriginID == 0 {
|
||||
t.Fatal("expected route to be linked with an auto-created origin")
|
||||
}
|
||||
if route.OriginURL != "https://origin.internal:8443" {
|
||||
t.Fatalf("unexpected route origin url: %s", route.OriginURL)
|
||||
}
|
||||
|
||||
origin, err := model.GetOriginByID(*route.OriginID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetOriginByID failed: %v", err)
|
||||
}
|
||||
if origin.Address != "origin.internal" {
|
||||
t.Fatalf("unexpected origin address: %s", origin.Address)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateOriginRewritesLinkedRouteOriginURL(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
origin, err := CreateOrigin(OriginInput{
|
||||
Name: "primary-origin",
|
||||
Address: "origin-a.internal",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateOrigin failed: %v", err)
|
||||
}
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginID: &origin.ID,
|
||||
OriginScheme: "https",
|
||||
OriginPort: "8443",
|
||||
OriginURI: "/api",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
updatedOrigin, err := UpdateOrigin(origin.ID, OriginInput{
|
||||
Name: origin.Name,
|
||||
Address: "origin-c.internal",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateOrigin failed: %v", err)
|
||||
}
|
||||
if updatedOrigin.Address != "origin-c.internal" {
|
||||
t.Fatalf("unexpected updated origin address: %s", updatedOrigin.Address)
|
||||
}
|
||||
|
||||
reloadedRoute, err := model.GetProxyRouteByID(route.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetProxyRouteByID failed: %v", err)
|
||||
}
|
||||
if reloadedRoute.OriginURL != "https://origin-c.internal:8443/api" {
|
||||
t.Fatalf("expected route origin url to be rewritten, got %s", reloadedRoute.OriginURL)
|
||||
}
|
||||
if reloadedRoute.Upstreams == "" || reloadedRoute.Upstreams == "[]" {
|
||||
t.Fatalf("expected route upstreams to be preserved, got %s", reloadedRoute.Upstreams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteOriginRejectsReferencedOrigin(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
origin, err := CreateOrigin(OriginInput{
|
||||
Address: "origin-a.internal",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateOrigin failed: %v", err)
|
||||
}
|
||||
if _, err = CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginID: &origin.ID,
|
||||
OriginScheme: "https",
|
||||
OriginPort: "443",
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
if err = DeleteOrigin(origin.ID); err == nil {
|
||||
t.Fatal("expected referenced origin deletion to fail")
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,8 @@ import (
|
||||
"openflare/model"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var proxyHeaderKeyPattern = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
|
||||
@@ -25,7 +27,12 @@ type ProxyRouteCustomHeaderInput struct {
|
||||
|
||||
type ProxyRouteInput struct {
|
||||
Domain string `json:"domain"`
|
||||
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"`
|
||||
@@ -85,7 +92,10 @@ func DeleteProxyRoute(id uint) error {
|
||||
|
||||
func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.ProxyRoute, error) {
|
||||
domain := strings.ToLower(strings.TrimSpace(input.Domain))
|
||||
originURL := strings.TrimSpace(input.OriginURL)
|
||||
originURL, originID, err := resolveProxyRoutePrimaryOrigin(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
originHost := strings.TrimSpace(input.OriginHost)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
upstreams, err := normalizeUpstreams(originURL, input.Upstreams)
|
||||
@@ -141,6 +151,7 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.Domain = domain
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = string(upstreamsJSON)
|
||||
@@ -156,6 +167,78 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func resolveProxyRoutePrimaryOrigin(input ProxyRouteInput) (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(*input.OriginID)
|
||||
if err != nil {
|
||||
return "", nil, errors.New("所选源站不存在")
|
||||
}
|
||||
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(address)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
originURL := strings.TrimSpace(input.OriginURL)
|
||||
if originURL == "" {
|
||||
return "", nil, errors.New("源站地址不能为空")
|
||||
}
|
||||
address, err := extractOriginAddress(originURL)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, findErr := model.GetOriginByAddress(address)
|
||||
if findErr == nil {
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
|
||||
return "", nil, findErr
|
||||
}
|
||||
return originURL, nil, nil
|
||||
}
|
||||
|
||||
func hasStructuredOriginInput(input ProxyRouteInput) bool {
|
||||
return (input.OriginID != nil && *input.OriginID != 0) ||
|
||||
strings.TrimSpace(input.OriginScheme) != "" ||
|
||||
strings.TrimSpace(input.OriginAddress) != "" ||
|
||||
strings.TrimSpace(input.OriginPort) != "" ||
|
||||
strings.TrimSpace(input.OriginURI) != ""
|
||||
}
|
||||
|
||||
func normalizeCustomHeaders(headers []ProxyRouteCustomHeaderInput) ([]ProxyRouteCustomHeaderInput, error) {
|
||||
if len(headers) == 0 {
|
||||
return []ProxyRouteCustomHeaderInput{}, nil
|
||||
|
||||
Reference in New Issue
Block a user