[优化] 代码优化

This commit is contained in:
ryan
2026-05-31 13:37:40 +08:00
parent c2fcd2eddf
commit 8894620b92
12 changed files with 125 additions and 57 deletions
+8
View File
@@ -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")
}
+10 -1
View File
@@ -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
}
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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
+3 -3
View File
@@ -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
+5 -20
View File
@@ -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 {
+5 -5
View File
@@ -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
+5 -20
View File
@@ -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
}
+19
View File
@@ -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
}