mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 06:16:37 +08:00
[优化] 代码优化
This commit is contained in:
@@ -11,6 +11,7 @@ import (
|
||||
"openflare/utils/security"
|
||||
"os"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
@@ -313,3 +314,10 @@ func CloseDB() error {
|
||||
err = sqlDB.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
func IsUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
@@ -412,7 +412,7 @@ func PublishConfigVersion(createdBy string, force bool) (*ReleaseResult, error)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("版本号生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
@@ -1095,6 +1095,9 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot,
|
||||
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
|
||||
}
|
||||
powEnabled, _ := getPoWConfigForRoute(route.ID, wafSnapshot)
|
||||
if route.PoWEnabled {
|
||||
powEnabled = true
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
continue
|
||||
@@ -1849,6 +1852,12 @@ func renderPowConfigBundle(routes []*model.ProxyRoute, wafSnapshot snapshotWAFDo
|
||||
hasPow := false
|
||||
for _, route := range routes {
|
||||
powEnabled, powConfig := getPoWConfigForRoute(route.ID, wafSnapshot)
|
||||
if route.PoWEnabled {
|
||||
powEnabled = true
|
||||
if decoded, err := decodeStoredPoWConfig(route.PoWEnabled, route.PoWConfig); err == nil {
|
||||
powConfig = decoded
|
||||
}
|
||||
}
|
||||
if !powEnabled {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -46,7 +46,7 @@ func CreateManagedDomain(input ManagedDomainInput) (*model.ManagedDomain, error)
|
||||
return nil, err
|
||||
}
|
||||
if err = domain.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
@@ -64,7 +64,7 @@ func UpdateManagedDomain(id uint, input ManagedDomainInput) (*model.ManagedDomai
|
||||
return nil, err
|
||||
}
|
||||
if err = domain.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
|
||||
@@ -83,7 +83,7 @@ func CreateNode(input NodeInput) (*NodeView, error) {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("节点标识生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
@@ -454,7 +454,7 @@ func RegisterNodeWithDiscovery(payload AgentNodePayload) (*AgentRegistrationResp
|
||||
}
|
||||
applyNodeRuntime(node, payload, false)
|
||||
if err = node.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("节点标识生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
|
||||
@@ -88,7 +88,7 @@ func CreateOrigin(input OriginInput) (*model.Origin, error) {
|
||||
return nil, err
|
||||
}
|
||||
if err = origin.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("源站地址已存在")
|
||||
}
|
||||
return nil, err
|
||||
@@ -108,7 +108,7 @@ func UpdateOrigin(id uint, input OriginInput) (*model.Origin, error) {
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Save(nextOrigin).Error; err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return errors.New("源站地址已存在")
|
||||
}
|
||||
return err
|
||||
@@ -171,7 +171,7 @@ func getOrCreateOriginByAddress(address string) (*model.Origin, error) {
|
||||
Remark: "",
|
||||
}
|
||||
if err := origin.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return model.GetOriginByAddress(normalizedAddress)
|
||||
}
|
||||
return nil, err
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net"
|
||||
"net/url"
|
||||
"openflare/model"
|
||||
"openflare/utils"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -121,7 +122,7 @@ func CreateProxyRoute(input ProxyRouteInput) (*ProxyRouteView, error) {
|
||||
return nil, err
|
||||
}
|
||||
if err = route.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("proxy route identity already exists")
|
||||
}
|
||||
return nil, err
|
||||
@@ -139,7 +140,7 @@ func UpdateProxyRoute(id uint, input ProxyRouteInput) (*ProxyRouteView, error) {
|
||||
return nil, err
|
||||
}
|
||||
if err = route.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("proxy route identity already exists")
|
||||
}
|
||||
return nil, err
|
||||
@@ -438,7 +439,6 @@ func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string,
|
||||
|
||||
func normalizeProxyRouteDomains(rawDomains []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(rawDomains))
|
||||
seen := make(map[string]struct{}, len(rawDomains))
|
||||
for _, rawDomain := range rawDomains {
|
||||
domain := normalizeProxyRouteDomainValue(rawDomain)
|
||||
if domain == "" {
|
||||
@@ -447,12 +447,9 @@ func normalizeProxyRouteDomains(rawDomains []string) ([]string, error) {
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, errors.New("domain format is invalid")
|
||||
}
|
||||
if _, ok := seen[domain]; ok {
|
||||
continue
|
||||
}
|
||||
seen[domain] = struct{}{}
|
||||
normalized = append(normalized, domain)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
if len(normalized) == 0 {
|
||||
return nil, errors.New("at least one domain is required")
|
||||
}
|
||||
@@ -850,15 +847,7 @@ func normalizeUpstreams(originURL string, upstreams []string) ([]string, error)
|
||||
}
|
||||
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)
|
||||
}
|
||||
unique := utils.Unique(trimmed)
|
||||
normalized := make([]string, 0, len(unique))
|
||||
var scheme string
|
||||
multiUpstream := len(unique) > 1
|
||||
@@ -1091,10 +1080,6 @@ func validateOriginHost(raw string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
return err != nil && strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
// PoW configuration types and validation
|
||||
|
||||
type ProxyRoutePoWListConfig struct {
|
||||
|
||||
@@ -105,7 +105,7 @@ func CreateTLSCertificate(input TLSCertificateInput) (*model.TLSCertificate, err
|
||||
return nil, err
|
||||
}
|
||||
if err = certificate.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
@@ -144,7 +144,7 @@ func UpdateTLSCertificate(id uint, input TLSCertificateInput) (*model.TLSCertifi
|
||||
return nil, err
|
||||
}
|
||||
if err = certificate.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
@@ -219,7 +219,7 @@ func ApplyTLSCertificate(input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
}
|
||||
|
||||
if err := cert.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
@@ -261,7 +261,7 @@ func UpdateAcmeCertificate(id uint, input TLSApplyInput) (*model.TLSCertificate,
|
||||
cert.ApplyStatus = "applying"
|
||||
|
||||
if err := cert.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
@@ -308,7 +308,7 @@ func ConvertTLSCertificateToAcme(id uint, input TLSApplyInput) (*model.TLSCertif
|
||||
cert.ApplyMessage = ""
|
||||
|
||||
if err := cert.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"openflare/model"
|
||||
"openflare/utils"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -413,7 +414,6 @@ func loadWAFBindings() (map[uint][]uint, error) {
|
||||
|
||||
func normalizeWAFIPList(items []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.TrimSpace(raw)
|
||||
if item == "" {
|
||||
@@ -432,19 +432,15 @@ func normalizeWAFIPList(items []string) ([]string, error) {
|
||||
}
|
||||
item = addr.String()
|
||||
}
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
sort.Strings(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeWAFCountryList(items []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.ToUpper(strings.TrimSpace(raw))
|
||||
if item == "" {
|
||||
@@ -453,30 +449,23 @@ func normalizeWAFCountryList(items []string) ([]string, error) {
|
||||
if len(item) != 2 || !unicode.IsLetter(rune(item[0])) || !unicode.IsLetter(rune(item[1])) {
|
||||
return nil, fmt.Errorf("%s 不是合法国家代码", item)
|
||||
}
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
sort.Strings(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeStringList(items []string) []string {
|
||||
normalized := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.TrimSpace(raw)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
sort.Strings(normalized)
|
||||
return normalized
|
||||
}
|
||||
@@ -518,18 +507,14 @@ func normalizeWAFRuleGroupIDs(groupIDs []uint) ([]uint, error) {
|
||||
}
|
||||
|
||||
func uniqueUintIDs(ids []uint) []uint {
|
||||
seen := make(map[uint]struct{}, len(ids))
|
||||
normalized := make([]uint, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
sort.Slice(normalized, func(i, j int) bool { return normalized[i] < normalized[j] })
|
||||
return normalized
|
||||
}
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
package utils
|
||||
|
||||
// Unique returns a new slice containing only the unique elements of the input slice,
|
||||
// preserving their original order.
|
||||
func Unique[T comparable](slice []T) []T {
|
||||
if slice == nil {
|
||||
return nil
|
||||
}
|
||||
seen := make(map[T]struct{})
|
||||
result := make([]T, 0)
|
||||
for _, item := range slice {
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result
|
||||
}
|
||||
Reference in New Issue
Block a user