Files
OpenFlare/internal/apps/openflare/origin/logics.go
T
ryan 943818f7d4 refactor(repository): 收敛 model/repository 分层为唯一持久化入口
将 OpenFlare 与平台业务的数据访问从 model 与 apps 直连迁入 repository,
model 仅保留实体与无 IO 规则;补充 code-check 架构守卫与开发规范。
2026-07-24 17:00:17 +08:00

233 lines
6.2 KiB
Go

// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package origin
import (
"context"
"encoding/json"
"errors"
"fmt"
"sort"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm"
)
// Input 源站创建/更新请求。
type Input struct {
Name string `json:"name"`
Address string `json:"address"`
Remark string `json:"remark"`
}
// RouteSummary 源站详情中的代理规则摘要。
type RouteSummary struct {
ID uint `json:"id"`
Domain string `json:"domain"`
OriginURL string `json:"origin_url"`
Enabled bool `json:"enabled"`
UpdatedAt string `json:"updated_at"`
}
// View 源站列表项。
type View struct {
ID uint `json:"id"`
Name string `json:"name"`
Address string `json:"address"`
Remark string `json:"remark"`
RouteCount int64 `json:"route_count"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
// DetailView 源站详情。
type DetailView struct {
View
Routes []RouteSummary `json:"routes"`
}
// ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]View, error) {
origins, err := repository.ListOrigins(ctx)
if err != nil {
return nil, err
}
return buildOriginViews(ctx, origins)
}
// GetOriginDetail 获取源站详情。
func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
origin, err := repository.GetOriginByID(ctx, id)
if err != nil {
return nil, err
}
views, err := buildOriginViews(ctx, []model.Origin{*origin})
if err != nil {
return nil, err
}
routes, err := repository.ListProxyRoutesByOriginID(ctx, id)
if err != nil {
return nil, err
}
items := make([]RouteSummary, 0, len(routes))
for _, route := range routes {
domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
domain := ""
if len(domains) > 0 {
domain = domains[0].Domain
}
items = append(items, RouteSummary{
ID: route.ID,
Domain: domain,
OriginURL: route.OriginURL,
Enabled: route.Enabled,
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
})
}
sort.Slice(items, func(i, j int) bool {
return items[i].Domain < items[j].Domain
})
return &DetailView{
View: views[0],
Routes: items,
}, nil
}
// CreateOrigin 创建源站。
func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) {
origin, err := buildOrigin(nil, input)
if err != nil {
return nil, err
}
if err = repository.CreateOriginRecord(ctx, origin); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errOriginAddressExists)
}
return nil, err
}
return origin, nil
}
// UpdateOrigin 更新源站。
func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, error) {
origin, err := repository.GetOriginByID(ctx, id)
if err != nil {
return nil, err
}
previousAddress := origin.Address
nextOrigin, err := buildOrigin(origin, input)
if err != nil {
return nil, err
}
err = repository.WithOriginTx(ctx, func(tx *gorm.DB) error {
if err := repository.SaveOriginTx(tx, nextOrigin); err != nil {
if isUniqueConstraintError(err) {
return errors.New(errOriginAddressExists)
}
return err
}
if previousAddress == nextOrigin.Address {
return nil
}
return updateRoutesForOriginAddress(ctx, tx, nextOrigin.ID, nextOrigin.Address)
})
if err != nil {
return nil, err
}
return nextOrigin, nil
}
// DeleteOrigin 删除源站。
func DeleteOrigin(ctx context.Context, id uint) error {
count, err := repository.CountProxyRoutesByOriginID(ctx, id)
if err != nil {
return err
}
if count > 0 {
return errors.New(errOriginDeleteReferenced)
}
if _, err = repository.GetOriginByID(ctx, id); err != nil {
return err
}
return repository.DeleteOriginRecord(ctx, id)
}
func buildOrigin(existing *model.Origin, input Input) (*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 buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, error) {
countRows, err := repository.ListOriginRouteCounts(ctx)
if err != nil {
return nil, err
}
countMap := make(map[uint]int64, len(countRows))
for _, row := range countRows {
countMap[row.OriginID] = row.RouteCount
}
views := make([]View, 0, len(origins))
for _, origin := range origins {
views = append(views, View{
ID: origin.ID,
Name: origin.Name,
Address: origin.Address,
Remark: origin.Remark,
RouteCount: countMap[origin.ID],
CreatedAt: origin.CreatedAt.Format("2006-01-02T15:04:05Z07:00"),
UpdatedAt: origin.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
})
}
return views, nil
}
func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uint, address string) error {
if !repository.HasProxyRoutesTable(ctx) {
return nil
}
routes, err := repository.ListProxyRoutesByOriginIDAscTx(tx, originID)
if 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 := repository.UpdateProxyRouteOriginAddressTx(tx, route.ID, rewrittenOriginURL, string(upstreamsJSON)); err != nil {
return fmt.Errorf("update route %d origin address failed: %w", route.ID, err)
}
}
return nil
}