From 649287a775562278be199c9209d68fa12fe1424a Mon Sep 17 00:00:00 2001 From: ryan Date: Sun, 31 May 2026 20:47:03 +0800 Subject: [PATCH] =?UTF-8?q?[=E4=BC=98=E5=8C=96]=20=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- openflare_server/model/node_access_log.go | 78 +------------------ .../model/node_metric_snapshot.go | 29 +++---- openflare_server/model/node_request_report.go | 29 +++---- openflare_server/model/sharding.go | 29 ------- openflare_server/service/agent.go | 4 - openflare_server/service/agent_ws.go | 6 -- openflare_server/service/auth_source.go | 9 +-- openflare_server/service/config_version.go | 15 ---- openflare_server/service/proxy_route.go | 2 +- openflare_server/service/tls_certificate.go | 69 ++++++---------- openflare_server/utils/slice.go | 32 +++++++- 11 files changed, 79 insertions(+), 223 deletions(-) diff --git a/openflare_server/model/node_access_log.go b/openflare_server/model/node_access_log.go index c1ece99e..0973cb6b 100644 --- a/openflare_server/model/node_access_log.go +++ b/openflare_server/model/node_access_log.go @@ -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" diff --git a/openflare_server/model/node_metric_snapshot.go b/openflare_server/model/node_metric_snapshot.go index 70d9ace7..e71d4bcc 100644 --- a/openflare_server/model/node_metric_snapshot.go +++ b/openflare_server/model/node_metric_snapshot.go @@ -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) { diff --git a/openflare_server/model/node_request_report.go b/openflare_server/model/node_request_report.go index 262c0241..e343a9ac 100644 --- a/openflare_server/model/node_request_report.go +++ b/openflare_server/model/node_request_report.go @@ -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) { diff --git a/openflare_server/model/sharding.go b/openflare_server/model/sharding.go index 0f232d5a..d5d64791 100644 --- a/openflare_server/model/sharding.go +++ b/openflare_server/model/sharding.go @@ -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]) - }) -} diff --git a/openflare_server/service/agent.go b/openflare_server/service/agent.go index 5cbcbb1e..3bcd54e6 100644 --- a/openflare_server/service/agent.go +++ b/openflare_server/service/agent.go @@ -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 diff --git a/openflare_server/service/agent_ws.go b/openflare_server/service/agent_ws.go index ea81644b..f39c459c 100644 --- a/openflare_server/service/agent_ws.go +++ b/openflare_server/service/agent_ws.go @@ -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 diff --git a/openflare_server/service/auth_source.go b/openflare_server/service/auth_source.go index e59e3ebd..3fefcf54 100644 --- a/openflare_server/service/auth_source.go +++ b/openflare_server/service/auth_source.go @@ -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 == "" { diff --git a/openflare_server/service/config_version.go b/openflare_server/service/config_version.go index c4b0bca2..b863ef46 100644 --- a/openflare_server/service/config_version.go +++ b/openflare_server/service/config_version.go @@ -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 { diff --git a/openflare_server/service/proxy_route.go b/openflare_server/service/proxy_route.go index 9fab4543..b7dffa01 100644 --- a/openflare_server/service/proxy_route.go +++ b/openflare_server/service/proxy_route.go @@ -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" diff --git a/openflare_server/service/tls_certificate.go b/openflare_server/service/tls_certificate.go index b3a89ef6..9f0a563b 100644 --- a/openflare_server/service/tls_certificate.go +++ b/openflare_server/service/tls_certificate.go @@ -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 { diff --git a/openflare_server/utils/slice.go b/openflare_server/utils/slice.go index 00128f75..ba836f59 100644 --- a/openflare_server/utils/slice.go +++ b/openflare_server/utils/slice.go @@ -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 +}