[功能] 添加源站管理功能,包括源站的创建、更新、删除及列表展示

This commit is contained in:
ryan
2026-03-20 20:01:42 +08:00
parent edd31da527
commit afd891f0f6
26 changed files with 2046 additions and 105 deletions
+241
View File
@@ -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
}
+192
View File
@@ -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
}
+105
View File
@@ -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")
}
}
+84 -1
View File
@@ -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