[优化] 代码优化

This commit is contained in:
ryan
2026-05-31 20:47:03 +08:00
parent 2514e7edc4
commit 649287a775
11 changed files with 79 additions and 223 deletions
+1 -77
View File
@@ -1,7 +1,6 @@
package model
import (
"fmt"
"sort"
"strings"
"time"
@@ -115,7 +114,7 @@ type NodeAccessLogTrendPointRow struct {
RequestCount int64 `json:"request_count"`
}
func (log *NodeAccessLog) BeforeCreate(tx *gorm.DB) error {
func (log *NodeAccessLog) BeforeCreate(*gorm.DB) error {
return assignObservabilityID(&log.ID)
}
@@ -326,16 +325,6 @@ func DeleteNodeAccessLogsByNodeBefore(db *gorm.DB, nodeID string, before time.Ti
})
}
func buildNodeAccessLogQuery(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB {
if db == nil {
db = DB.Model(&NodeAccessLog{})
}
if db.Statement == nil || db.Statement.Model == nil {
db = db.Model(&NodeAccessLog{})
}
return applyNodeAccessLogFilters(db, query)
}
func applyNodeAccessLogFilters(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB {
if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" {
db = db.Where("node_id LIKE ?", "%"+trimmed+"%")
@@ -748,71 +737,6 @@ func compareUint(left uint, right uint) int {
}
}
func buildNodeAccessLogSortClause(sortBy string, sortOrder string) string {
column := "logged_at"
switch strings.TrimSpace(sortBy) {
case "status_code":
column = "status_code"
case "remote_addr":
column = "remote_addr"
case "host":
column = "host"
case "path":
column = "path"
}
order := normalizeSortOrder(sortOrder)
if column == "logged_at" {
return fmt.Sprintf("%s %s, id %s", column, order, order)
}
return fmt.Sprintf("%s %s, logged_at desc, id desc", column, order)
}
func buildNodeAccessLogBucketSortClause(sortBy string, sortOrder string) string {
order := normalizeSortOrder(sortOrder)
switch strings.TrimSpace(sortBy) {
case "request_count":
return fmt.Sprintf("request_count %s, bucket_epoch desc", order)
default:
return fmt.Sprintf("bucket_epoch %s", order)
}
}
func buildNodeAccessLogIPSummarySortClause(sortBy string, sortOrder string) string {
order := normalizeSortOrder(sortOrder)
switch strings.TrimSpace(sortBy) {
case "recent_requests":
return fmt.Sprintf("recent_requests %s, last_seen_epoch desc, remote_addr asc", order)
case "last_seen_at":
return fmt.Sprintf("last_seen_epoch %s, total_requests desc, remote_addr asc", order)
case "remote_addr":
return fmt.Sprintf("remote_addr %s", order)
default:
return fmt.Sprintf("total_requests %s, last_seen_epoch desc, remote_addr asc", order)
}
}
func accessLogBucketEpochExpr(bucketMinutes int) string {
bucketSeconds := bucketMinutes * 60
if bucketSeconds <= 0 {
bucketSeconds = 180
}
switch DB.Dialector.Name() {
case "postgres":
return fmt.Sprintf("CAST(floor(extract(epoch from logged_at) / %d) * %d AS BIGINT)", bucketSeconds, bucketSeconds)
default:
return fmt.Sprintf("CAST((strftime('%%s', logged_at) / %d) * %d AS INTEGER)", bucketSeconds, bucketSeconds)
}
}
func accessLogEpochExpr(expression string) string {
switch DB.Dialector.Name() {
case "postgres":
return fmt.Sprintf("CAST(extract(epoch from %s) AS BIGINT)", expression)
default:
return fmt.Sprintf("CAST(strftime('%%s', %s) AS INTEGER)", expression)
}
}
func normalizeSortOrder(sortOrder string) string {
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
return "asc"
+11 -18
View File
@@ -1,7 +1,7 @@
package model
import (
"sort"
"openflare/utils"
"time"
"gorm.io/gorm"
@@ -26,6 +26,14 @@ type NodeMetricSnapshot struct {
CreatedAt time.Time `json:"created_at"`
}
func (snapshot *NodeMetricSnapshot) GetID() uint {
return snapshot.ID
}
func (snapshot *NodeMetricSnapshot) GetTime() time.Time {
return snapshot.CapturedAt
}
func (snapshot *NodeMetricSnapshot) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&snapshot.ID)
}
@@ -52,16 +60,7 @@ func ListNodeMetricSnapshots(nodeID string, since time.Time, limit int) (snapsho
if err != nil {
return nil, err
}
sort.Slice(rows, func(i int, j int) bool {
if rows[i].CapturedAt.Equal(rows[j].CapturedAt) {
return rows[i].ID > rows[j].ID
}
return rows[i].CapturedAt.After(rows[j].CapturedAt)
})
if limit > 0 && len(rows) > limit {
rows = rows[:limit]
}
return rows, nil
return utils.SortAndLimitRecords(rows, limit), nil
}
func ListMetricSnapshotsSince(since time.Time) (snapshots []*NodeMetricSnapshot, err error) {
@@ -79,13 +78,7 @@ func ListMetricSnapshotsSince(since time.Time) (snapshots []*NodeMetricSnapshot,
if err != nil {
return nil, err
}
sort.Slice(rows, func(i int, j int) bool {
if rows[i].CapturedAt.Equal(rows[j].CapturedAt) {
return rows[i].ID > rows[j].ID
}
return rows[i].CapturedAt.After(rows[j].CapturedAt)
})
return rows, nil
return utils.SortAndLimitRecords(rows, 0), nil
}
func NodeMetricSnapshotExists(db *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
+11 -18
View File
@@ -1,7 +1,7 @@
package model
import (
"sort"
"openflare/utils"
"time"
"gorm.io/gorm"
@@ -21,6 +21,14 @@ type NodeRequestReport struct {
CreatedAt time.Time `json:"created_at"`
}
func (report *NodeRequestReport) GetID() uint {
return report.ID
}
func (report *NodeRequestReport) GetTime() time.Time {
return report.WindowEndedAt
}
func (report *NodeRequestReport) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&report.ID)
}
@@ -47,16 +55,7 @@ func ListNodeRequestReports(nodeID string, since time.Time, limit int) (reports
if err != nil {
return nil, err
}
sort.Slice(rows, func(i int, j int) bool {
if rows[i].WindowEndedAt.Equal(rows[j].WindowEndedAt) {
return rows[i].ID > rows[j].ID
}
return rows[i].WindowEndedAt.After(rows[j].WindowEndedAt)
})
if limit > 0 && len(rows) > limit {
rows = rows[:limit]
}
return rows, nil
return utils.SortAndLimitRecords(rows, limit), nil
}
func ListRequestReportsSince(since time.Time) (reports []*NodeRequestReport, err error) {
@@ -74,13 +73,7 @@ func ListRequestReportsSince(since time.Time) (reports []*NodeRequestReport, err
if err != nil {
return nil, err
}
sort.Slice(rows, func(i int, j int) bool {
if rows[i].WindowEndedAt.Equal(rows[j].WindowEndedAt) {
return rows[i].ID > rows[j].ID
}
return rows[i].WindowEndedAt.After(rows[j].WindowEndedAt)
})
return rows, nil
return utils.SortAndLimitRecords(rows, 0), nil
}
func NodeRequestReportExists(db *gorm.DB, nodeID string, windowStartedAt time.Time, windowEndedAt time.Time) (bool, error) {
-29
View File
@@ -2,7 +2,6 @@ package model
import (
"fmt"
"sort"
"strconv"
"strings"
"sync"
@@ -129,10 +128,6 @@ func observabilityShardSuffixForValue(value any) (string, error) {
}
}
func observabilityShardTableForID(baseTable string, id uint) string {
return baseTable + observabilityShardSuffixForID(id)
}
func legacyObservabilityShardTableName(tableName string) string {
return tableName + "_legacy_v2_to_v3"
}
@@ -144,24 +139,6 @@ func normalizeShardedDB(db *gorm.DB) *gorm.DB {
return DB
}
func sessionIgnoringSharding(db *gorm.DB) *gorm.DB {
db = normalizeShardedDB(db)
if db == nil {
return nil
}
return db.Session(&gorm.Session{}).Set(sharding.ShardingIgnoreStoreKey, true)
}
func baseDialector(db *gorm.DB) gorm.Dialector {
if db == nil {
return nil
}
if dialector, ok := db.Dialector.(sharding.ShardingDialector); ok {
return dialector.Dialector
}
return db.Dialector
}
func nextObservabilityID() (uint, error) {
observabilityIDNodeOnce.Do(func() {
observabilityIDNode, observabilityIDNodeErr = snowflake.NewNode(0)
@@ -223,9 +200,3 @@ func deleteAcrossShards(db *gorm.DB, baseTable string, model any, apply func(tx
}
return deleted, nil
}
func sortShardRows[T any](items []T, less func(left T, right T) bool) {
sort.Slice(items, func(i int, j int) bool {
return less(items[i], items[j])
})
}
-4
View File
@@ -144,10 +144,6 @@ type NodeView struct {
UpdatedAt time.Time `json:"updated_at"`
}
func RegisterNode(node *model.Node, payload AgentNodePayload) (*AgentRegistrationResponse, error) {
return RegisterNodeWithAgentToken(node, payload)
}
func HeartbeatNode(node *model.Node, payload AgentNodePayload) (*HeartbeatResponse, error) {
slog.Debug("agent heartbeat received", "node_id", node.NodeID, "current_version", strings.TrimSpace(payload.CurrentVersion))
payload.NodeID = node.NodeID
-6
View File
@@ -157,12 +157,6 @@ func IsAgentWSConnected(nodeID string) bool {
}
}
func AgentWSClientCount() int {
defaultAgentWSHub.mu.RLock()
defer defaultAgentWSHub.mu.RUnlock()
return len(defaultAgentWSHub.clients)
}
func SendAgentWSSettings(nodeID string, settings *AgentSettings) bool {
if settings == nil {
return false
+1 -8
View File
@@ -190,14 +190,6 @@ func CompleteOAuthLogin(source *model.AuthSource, profile *OAuthProfile, current
return &OAuthCallbackResult{Status: "linked", User: user}, nil, nil
}
if common.RegisterEnabled {
user, err := createUserFromOAuthProfile(source, profile)
if err != nil {
return nil, nil, err
}
return &OAuthCallbackResult{Status: "registered", User: user}, nil, nil
}
pending := &PendingExternalAccount{
AuthSourceID: source.ID,
ExternalID: profile.ExternalID,
@@ -244,6 +236,7 @@ func LinkPendingExternalAccount(pending *PendingExternalAccount, input LinkExist
return &user, nil
}
// CreateUserFromOAuthProfile 根据 OAuth 资料创建新用户
func createUserFromOAuthProfile(source *model.AuthSource, profile *OAuthProfile) (*model.User, error) {
displayName := strings.TrimSpace(profile.DisplayName)
if displayName == "" {
@@ -292,21 +292,6 @@ func DiffConfigVersion() (*ConfigDiffResult, error) {
return result, nil
}
func HasConfigChanges() (bool, error) {
bundle, err := buildCurrentConfigBundle(false)
if err != nil {
return false, err
}
activeVersion, err := model.GetActiveConfigVersion()
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return len(bundle.Routes) > 0, nil
}
return false, err
}
return activeVersion.Checksum != bundle.Checksum, nil
}
func PublishConfigVersion(createdBy string, force bool) (*ReleaseResult, error) {
bundle, err := buildCurrentConfigBundle(true)
if err != nil {
+1 -1
View File
@@ -16,7 +16,7 @@ import (
)
var proxyHeaderKeyPattern = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
var proxyRouteLimitRatePattern = regexp.MustCompile(`^\d+(?:[kKmM])?$`)
var proxyRouteLimitRatePattern = regexp.MustCompile(`^\d+[kKmM]?$`)
const (
proxyRouteCachePolicyURL = "url"
+23 -46
View File
@@ -194,25 +194,29 @@ func DeleteTLSCertificate(id uint) error {
return certificate.Delete()
}
func fillAcmeCertificateFields(cert *model.TLSCertificate, input TLSApplyInput) {
cert.Name = strings.TrimSpace(input.Name)
cert.Remark = strings.TrimSpace(input.Remark)
cert.AcmeAccountID = input.AcmeAccountID
cert.DnsAccountID = input.DnsAccountID
cert.KeyAlgorithm = input.KeyAlgorithm
cert.AutoRenew = input.AutoRenew
cert.PrimaryDomain = strings.TrimSpace(input.PrimaryDomain)
cert.OtherDomains = strings.TrimSpace(input.OtherDomains)
cert.DisableCNAME = input.DisableCNAME
cert.SkipDNS = input.SkipDNS
cert.DNS1 = strings.TrimSpace(input.DNS1)
cert.DNS2 = strings.TrimSpace(input.DNS2)
cert.ApplyStatus = "applying"
}
func ApplyTLSCertificate(input TLSApplyInput) (*model.TLSCertificate, error) {
cert := &model.TLSCertificate{
Name: strings.TrimSpace(input.Name),
Remark: strings.TrimSpace(input.Remark),
Provider: "acme",
AcmeAccountID: input.AcmeAccountID,
DnsAccountID: input.DnsAccountID,
KeyAlgorithm: input.KeyAlgorithm,
AutoRenew: input.AutoRenew,
PrimaryDomain: strings.TrimSpace(input.PrimaryDomain),
OtherDomains: strings.TrimSpace(input.OtherDomains),
DisableCNAME: input.DisableCNAME,
SkipDNS: input.SkipDNS,
DNS1: strings.TrimSpace(input.DNS1),
DNS2: strings.TrimSpace(input.DNS2),
ApplyStatus: "applying",
CertPEM: " ", // Temporary empty value, since gorm may prevent empty insert
KeyPEM: " ", // Temporary empty value
Provider: "acme",
CertPEM: " ", // Temporary empty value, since gorm may prevent empty insert
KeyPEM: " ", // Temporary empty value
}
fillAcmeCertificateFields(cert, input)
if cert.Name == "" {
return nil, errors.New("certificate name cannot be empty")
@@ -242,24 +246,11 @@ func UpdateAcmeCertificate(id uint, input TLSApplyInput) (*model.TLSCertificate,
return nil, errors.New("only acme certificates can be updated via this endpoint")
}
cert.Name = strings.TrimSpace(input.Name)
fillAcmeCertificateFields(cert, input)
if cert.Name == "" {
return nil, errors.New("certificate name cannot be empty")
}
cert.Remark = strings.TrimSpace(input.Remark)
cert.AcmeAccountID = input.AcmeAccountID
cert.DnsAccountID = input.DnsAccountID
cert.KeyAlgorithm = input.KeyAlgorithm
cert.AutoRenew = input.AutoRenew
cert.PrimaryDomain = strings.TrimSpace(input.PrimaryDomain)
cert.OtherDomains = strings.TrimSpace(input.OtherDomains)
cert.DisableCNAME = input.DisableCNAME
cert.SkipDNS = input.SkipDNS
cert.DNS1 = strings.TrimSpace(input.DNS1)
cert.DNS2 = strings.TrimSpace(input.DNS2)
cert.ApplyStatus = "applying"
if err := cert.Update(); err != nil {
if model.IsUniqueConstraintError(err) {
return nil, errors.New("certificate name already exists")
@@ -287,24 +278,10 @@ func ConvertTLSCertificateToAcme(id uint, input TLSApplyInput) (*model.TLSCertif
return nil, errors.New("certificate is already applying")
}
name := strings.TrimSpace(input.Name)
if name == "" {
fillAcmeCertificateFields(cert, input)
if cert.Name == "" {
return nil, errors.New("certificate name cannot be empty")
}
cert.Name = name
cert.Remark = strings.TrimSpace(input.Remark)
cert.AcmeAccountID = input.AcmeAccountID
cert.DnsAccountID = input.DnsAccountID
cert.KeyAlgorithm = input.KeyAlgorithm
cert.AutoRenew = input.AutoRenew
cert.PrimaryDomain = strings.TrimSpace(input.PrimaryDomain)
cert.OtherDomains = strings.TrimSpace(input.OtherDomains)
cert.DisableCNAME = input.DisableCNAME
cert.SkipDNS = input.SkipDNS
cert.DNS1 = strings.TrimSpace(input.DNS1)
cert.DNS2 = strings.TrimSpace(input.DNS2)
cert.ApplyStatus = "applying"
cert.ApplyMessage = ""
if err := cert.Update(); err != nil {
+31 -1
View File
@@ -1,6 +1,10 @@
package utils
import "strings"
import (
"sort"
"strings"
"time"
)
// Unique returns a new slice containing only the unique elements of the input slice,
// preserving their original order.
@@ -44,3 +48,29 @@ func UniqueAndCleanStringSlice(slice []string) []string {
}
return result
}
// IdentifiableTimeRecord represents a database record that has a unique ID and a primary timestamp field.
type IdentifiableTimeRecord interface {
GetID() uint
GetTime() time.Time
}
// SortAndLimitRecords sorts a slice of IdentifiableTimeRecord descendingly by their timestamp (and ID as a tie-breaker),
// and limits the slice to the specified size if limit > 0.
func SortAndLimitRecords[T IdentifiableTimeRecord](rows []T, limit int) []T {
if len(rows) == 0 {
return rows
}
sort.Slice(rows, func(i, j int) bool {
ti := rows[i].GetTime()
tj := rows[j].GetTime()
if ti.Equal(tj) {
return rows[i].GetID() > rows[j].GetID()
}
return ti.After(tj)
})
if limit > 0 && len(rows) > limit {
rows = rows[:limit]
}
return rows
}