mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 23:06:36 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -0,0 +1,608 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultAccessLogPageSize = 20
|
||||
maxAccessLogPageSize = 200
|
||||
defaultAccessLogSortBy = "logged_at"
|
||||
defaultAccessLogSortOrder = "desc"
|
||||
defaultAccessLogFoldMinute = 3
|
||||
defaultIPTrendHours = 24
|
||||
defaultIPTrendBucketMinute = 30
|
||||
maxIPTrendHours = 168
|
||||
nodeAccessLogRetentionDays = 90
|
||||
)
|
||||
|
||||
type AccessLogQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
}
|
||||
|
||||
type AccessLogView struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
NodeName string `json:"node_name"`
|
||||
LoggedAt time.Time `json:"logged_at"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Region string `json:"region"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
}
|
||||
|
||||
type AccessLogList struct {
|
||||
Items []AccessLogView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalRecord int64 `json:"total_record"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
}
|
||||
|
||||
type FoldedAccessLogView struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueIPCount int64 `json:"unique_ip_count"`
|
||||
UniqueHostCount int64 `json:"unique_host_count"`
|
||||
SuccessCount int64 `json:"success_count"`
|
||||
ClientErrorCount int64 `json:"client_error_count"`
|
||||
ServerErrorCount int64 `json:"server_error_count"`
|
||||
}
|
||||
|
||||
type FoldedAccessLogList struct {
|
||||
Items []FoldedAccessLogView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalBucket int64 `json:"total_bucket"`
|
||||
TotalRecord int64 `json:"total_record"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
}
|
||||
|
||||
type FoldedAccessLogIPQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
BucketStartedAt string `json:"bucket_started_at"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
type FoldedAccessLogIPView struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
SuccessCount int64 `json:"success_count"`
|
||||
ClientErrorCount int64 `json:"client_error_count"`
|
||||
ServerErrorCount int64 `json:"server_error_count"`
|
||||
LastSeenAt time.Time `json:"last_seen_at"`
|
||||
}
|
||||
|
||||
type FoldedAccessLogIPList struct {
|
||||
Items []FoldedAccessLogIPView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
type AccessLogIPSummaryQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
type AccessLogIPSummaryView struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
TotalRequests int64 `json:"total_requests"`
|
||||
RecentRequests int64 `json:"recent_requests"`
|
||||
LastSeenAt time.Time `json:"last_seen_at"`
|
||||
}
|
||||
|
||||
type AccessLogIPSummaryList struct {
|
||||
Items []AccessLogIPSummaryView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
type AccessLogIPTrendQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Hours int `json:"hours"`
|
||||
BucketMinutes int `json:"bucket_minutes"`
|
||||
}
|
||||
|
||||
type AccessLogIPTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
}
|
||||
|
||||
type AccessLogIPTrendView struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Hours int `json:"hours"`
|
||||
BucketMinutes int `json:"bucket_minutes"`
|
||||
Points []AccessLogIPTrendPoint `json:"points"`
|
||||
}
|
||||
|
||||
type AccessLogCleanupInput struct {
|
||||
RetentionDays int `json:"retention_days"`
|
||||
}
|
||||
|
||||
type AccessLogCleanupResult struct {
|
||||
RetentionDays int `json:"retention_days"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
Cutoff time.Time `json:"cutoff"`
|
||||
}
|
||||
|
||||
func ListAccessLogs(input AccessLogQuery) (*AccessLogList, error) {
|
||||
normalized := normalizeAccessLogQuery(input)
|
||||
modelQuery := buildModelAccessLogQuery(normalized)
|
||||
logs, err := model.ListNodeAccessLogs(modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalRecords, totalIPs, err := model.CountNodeAccessLogs(modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeNames, err := listNodeNameMap(logs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]AccessLogView, 0, len(logs))
|
||||
for _, item := range logs {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, AccessLogView{
|
||||
ID: item.ID,
|
||||
NodeID: item.NodeID,
|
||||
NodeName: nodeNames[item.NodeID],
|
||||
LoggedAt: item.LoggedAt,
|
||||
RemoteAddr: item.RemoteAddr,
|
||||
Region: item.Region,
|
||||
Host: item.Host,
|
||||
Path: item.Path,
|
||||
StatusCode: item.StatusCode,
|
||||
})
|
||||
}
|
||||
return &AccessLogList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalRecords,
|
||||
TotalRecord: totalRecords,
|
||||
TotalIP: totalIPs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ListFoldedAccessLogs(input AccessLogQuery) (*FoldedAccessLogList, error) {
|
||||
normalized := normalizeAccessLogQuery(input)
|
||||
foldMinutes, err := normalizeFoldMinutes(normalized.FoldMinutes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
modelQuery := buildModelAccessLogQuery(normalized)
|
||||
bucketQuery := model.NodeAccessLogBucketQuery{
|
||||
NodeID: modelQuery.NodeID,
|
||||
RemoteAddr: modelQuery.RemoteAddr,
|
||||
Host: modelQuery.Host,
|
||||
Path: modelQuery.Path,
|
||||
Since: modelQuery.Since,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
SortBy: normalizeFoldSortBy(input.SortBy),
|
||||
SortOrder: normalized.SortOrder,
|
||||
FoldMinutes: foldMinutes,
|
||||
}
|
||||
items, err := model.ListNodeAccessLogBuckets(bucketQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalBuckets, err := model.CountNodeAccessLogBuckets(bucketQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalRecords, totalIPs, err := model.CountNodeAccessLogs(modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]FoldedAccessLogView, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, FoldedAccessLogView{
|
||||
BucketStartedAt: time.Unix(item.BucketEpoch, 0).UTC(),
|
||||
RequestCount: item.RequestCount,
|
||||
UniqueIPCount: item.UniqueIPCount,
|
||||
UniqueHostCount: item.UniqueHostCount,
|
||||
SuccessCount: item.SuccessCount,
|
||||
ClientErrorCount: item.ClientErrorCount,
|
||||
ServerErrorCount: item.ServerErrorCount,
|
||||
})
|
||||
}
|
||||
return &FoldedAccessLogList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalBuckets,
|
||||
TotalBucket: totalBuckets,
|
||||
TotalRecord: totalRecords,
|
||||
TotalIP: totalIPs,
|
||||
FoldMinutes: foldMinutes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ListFoldedAccessLogIPs(input FoldedAccessLogIPQuery) (*FoldedAccessLogIPList, error) {
|
||||
normalized, bucketStartedAt, err := normalizeFoldedAccessLogIPQuery(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
modelQuery := model.NodeAccessLogBucketIPQuery{
|
||||
NodeID: normalized.NodeID,
|
||||
RemoteAddr: normalized.RemoteAddr,
|
||||
Host: normalized.Host,
|
||||
Path: normalized.Path,
|
||||
BucketStartedAt: bucketStartedAt,
|
||||
FoldMinutes: normalized.FoldMinutes,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}
|
||||
items, err := model.ListNodeAccessLogBucketIPs(modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalIP, err := model.CountNodeAccessLogBucketIPs(modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]FoldedAccessLogIPView, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, FoldedAccessLogIPView{
|
||||
RemoteAddr: item.RemoteAddr,
|
||||
RequestCount: item.RequestCount,
|
||||
SuccessCount: item.SuccessCount,
|
||||
ClientErrorCount: item.ClientErrorCount,
|
||||
ServerErrorCount: item.ServerErrorCount,
|
||||
LastSeenAt: time.Unix(item.LastSeenEpoch, 0).UTC(),
|
||||
})
|
||||
}
|
||||
return &FoldedAccessLogIPList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalIP,
|
||||
TotalIP: totalIP,
|
||||
BucketStartedAt: bucketStartedAt,
|
||||
FoldMinutes: normalized.FoldMinutes,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ListAccessLogIPSummaries(input AccessLogIPSummaryQuery) (*AccessLogIPSummaryList, error) {
|
||||
normalized := normalizeAccessLogIPSummaryQuery(input)
|
||||
since := time.Now().UTC().Add(-nodeAccessLogRetentionWindow)
|
||||
recentSince := time.Now().UTC().Add(-3 * time.Hour)
|
||||
query := model.NodeAccessLogIPSummaryQuery{
|
||||
NodeID: strings.TrimSpace(normalized.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(normalized.RemoteAddr),
|
||||
Host: strings.TrimSpace(normalized.Host),
|
||||
Since: since,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}
|
||||
items, err := model.ListNodeAccessLogIPSummaries(query, recentSince)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalIP, err := model.CountNodeAccessLogIPSummaries(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]AccessLogIPSummaryView, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, AccessLogIPSummaryView{
|
||||
RemoteAddr: item.RemoteAddr,
|
||||
TotalRequests: item.TotalRequests,
|
||||
RecentRequests: item.RecentRequests,
|
||||
LastSeenAt: time.Unix(item.LastSeenEpoch, 0).UTC(),
|
||||
})
|
||||
}
|
||||
return &AccessLogIPSummaryList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalIP,
|
||||
TotalIP: totalIP,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func GetAccessLogIPTrend(input AccessLogIPTrendQuery) (*AccessLogIPTrendView, error) {
|
||||
normalized, err := normalizeAccessLogIPTrendQuery(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
points, err := model.ListNodeAccessLogIPTrend(model.NodeAccessLogIPTrendQuery{
|
||||
NodeID: strings.TrimSpace(normalized.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(normalized.RemoteAddr),
|
||||
Host: strings.TrimSpace(normalized.Host),
|
||||
Since: time.Now().UTC().Add(-time.Duration(normalized.Hours) * time.Hour),
|
||||
BucketMinutes: normalized.BucketMinutes,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pointMap := make(map[int64]int64, len(points))
|
||||
for _, item := range points {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
pointMap[item.BucketEpoch] = item.RequestCount
|
||||
}
|
||||
bucketDuration := time.Duration(normalized.BucketMinutes) * time.Minute
|
||||
start := time.Now().UTC().Add(-time.Duration(normalized.Hours) * time.Hour).Truncate(bucketDuration)
|
||||
end := time.Now().UTC().Truncate(bucketDuration)
|
||||
views := make([]AccessLogIPTrendPoint, 0, int(end.Sub(start)/bucketDuration)+1)
|
||||
for cursor := start; !cursor.After(end); cursor = cursor.Add(bucketDuration) {
|
||||
views = append(views, AccessLogIPTrendPoint{
|
||||
BucketStartedAt: cursor,
|
||||
RequestCount: pointMap[cursor.Unix()],
|
||||
})
|
||||
}
|
||||
return &AccessLogIPTrendView{
|
||||
RemoteAddr: normalized.RemoteAddr,
|
||||
Hours: normalized.Hours,
|
||||
BucketMinutes: normalized.BucketMinutes,
|
||||
Points: views,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func CleanupAccessLogs(input AccessLogCleanupInput) (*AccessLogCleanupResult, error) {
|
||||
if input.RetentionDays <= 0 || input.RetentionDays > nodeAccessLogRetentionDays {
|
||||
return nil, errors.New("retention_days 必须在 1 到 90 之间")
|
||||
}
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
|
||||
deleted, err := model.DeleteNodeAccessLogsBefore(cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &AccessLogCleanupResult{
|
||||
RetentionDays: input.RetentionDays,
|
||||
DeletedCount: deleted,
|
||||
Cutoff: cutoff,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildModelAccessLogQuery(input AccessLogQuery) model.NodeAccessLogQuery {
|
||||
return model.NodeAccessLogQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
Since: time.Now().UTC().Add(-nodeAccessLogRetentionWindow),
|
||||
Page: input.Page,
|
||||
PageSize: input.PageSize,
|
||||
SortBy: input.SortBy,
|
||||
SortOrder: input.SortOrder,
|
||||
}
|
||||
}
|
||||
|
||||
func listNodeNameMap(logs []*model.NodeAccessLog) (map[string]string, error) {
|
||||
nodeIDs := make([]string, 0, len(logs))
|
||||
seen := make(map[string]struct{}, len(logs))
|
||||
for _, item := range logs {
|
||||
if item == nil || item.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[item.NodeID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[item.NodeID] = struct{}{}
|
||||
nodeIDs = append(nodeIDs, item.NodeID)
|
||||
}
|
||||
nodes, err := model.ListNodesByNodeIDs(nodeIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[string]string, len(nodes))
|
||||
for _, node := range nodes {
|
||||
if node == nil {
|
||||
continue
|
||||
}
|
||||
result[node.NodeID] = node.Name
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func normalizeAccessLogQuery(input AccessLogQuery) AccessLogQuery {
|
||||
return AccessLogQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
Page: normalizeAccessLogPage(input.Page),
|
||||
PageSize: normalizeAccessLogPageSize(input.PageSize),
|
||||
SortBy: normalizeAccessLogSortBy(input.SortBy),
|
||||
SortOrder: normalizeAccessLogSortOrder(input.SortOrder),
|
||||
FoldMinutes: input.FoldMinutes,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeAccessLogIPSummaryQuery(input AccessLogIPSummaryQuery) AccessLogIPSummaryQuery {
|
||||
return AccessLogIPSummaryQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Page: normalizeAccessLogPage(input.Page),
|
||||
PageSize: normalizeAccessLogPageSize(input.PageSize),
|
||||
SortBy: normalizeIPSummarySortBy(input.SortBy),
|
||||
SortOrder: normalizeAccessLogSortOrder(input.SortOrder),
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeFoldedAccessLogIPQuery(input FoldedAccessLogIPQuery) (FoldedAccessLogIPQuery, time.Time, error) {
|
||||
foldMinutes, err := normalizeFoldMinutes(input.FoldMinutes)
|
||||
if err != nil {
|
||||
return FoldedAccessLogIPQuery{}, time.Time{}, err
|
||||
}
|
||||
bucketStartedAt, err := time.Parse(time.RFC3339, strings.TrimSpace(input.BucketStartedAt))
|
||||
if err != nil {
|
||||
return FoldedAccessLogIPQuery{}, time.Time{}, errors.New("bucket_started_at 必须为 RFC3339 时间")
|
||||
}
|
||||
normalizedSortBy := strings.TrimSpace(input.SortBy)
|
||||
switch normalizedSortBy {
|
||||
case "last_seen_at", "remote_addr":
|
||||
default:
|
||||
normalizedSortBy = "request_count"
|
||||
}
|
||||
return FoldedAccessLogIPQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
BucketStartedAt: strings.TrimSpace(input.BucketStartedAt),
|
||||
FoldMinutes: foldMinutes,
|
||||
Page: normalizeAccessLogPage(input.Page),
|
||||
PageSize: normalizeAccessLogPageSize(input.PageSize),
|
||||
SortBy: normalizedSortBy,
|
||||
SortOrder: normalizeAccessLogSortOrder(input.SortOrder),
|
||||
}, bucketStartedAt.UTC(), nil
|
||||
}
|
||||
|
||||
func normalizeAccessLogIPTrendQuery(input AccessLogIPTrendQuery) (AccessLogIPTrendQuery, error) {
|
||||
remoteAddr := strings.TrimSpace(input.RemoteAddr)
|
||||
if remoteAddr == "" {
|
||||
return AccessLogIPTrendQuery{}, errors.New("remote_addr 不能为空")
|
||||
}
|
||||
hours := input.Hours
|
||||
if hours <= 0 {
|
||||
hours = defaultIPTrendHours
|
||||
}
|
||||
if hours > maxIPTrendHours {
|
||||
hours = maxIPTrendHours
|
||||
}
|
||||
bucketMinutes := input.BucketMinutes
|
||||
if bucketMinutes <= 0 {
|
||||
bucketMinutes = defaultIPTrendBucketMinute
|
||||
}
|
||||
switch bucketMinutes {
|
||||
case 5, 10, 15, 30, 60:
|
||||
default:
|
||||
return AccessLogIPTrendQuery{}, errors.New("bucket_minutes 仅支持 5、10、15、30、60")
|
||||
}
|
||||
return AccessLogIPTrendQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: remoteAddr,
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Hours: hours,
|
||||
BucketMinutes: bucketMinutes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeAccessLogPage(page int) int {
|
||||
if page < 0 {
|
||||
return 0
|
||||
}
|
||||
return page
|
||||
}
|
||||
|
||||
func normalizeAccessLogPageSize(pageSize int) int {
|
||||
if pageSize <= 0 {
|
||||
return defaultAccessLogPageSize
|
||||
}
|
||||
if pageSize > maxAccessLogPageSize {
|
||||
return maxAccessLogPageSize
|
||||
}
|
||||
return pageSize
|
||||
}
|
||||
|
||||
func normalizeAccessLogSortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "status_code", "remote_addr", "host", "path":
|
||||
return strings.TrimSpace(sortBy)
|
||||
default:
|
||||
return defaultAccessLogSortBy
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeAccessLogSortOrder(sortOrder string) string {
|
||||
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
|
||||
return "asc"
|
||||
}
|
||||
return defaultAccessLogSortOrder
|
||||
}
|
||||
|
||||
func normalizeFoldSortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "request_count":
|
||||
return "request_count"
|
||||
default:
|
||||
return "bucket_started_at"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeIPSummarySortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "recent_requests", "last_seen_at", "remote_addr":
|
||||
return strings.TrimSpace(sortBy)
|
||||
default:
|
||||
return "total_requests"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeFoldMinutes(value int) (int, error) {
|
||||
if value <= 0 {
|
||||
return defaultAccessLogFoldMinute, nil
|
||||
}
|
||||
switch value {
|
||||
case 3, 5:
|
||||
return value, nil
|
||||
default:
|
||||
return 0, errors.New("fold_minutes 仅支持 3 或 5")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
|
||||
)
|
||||
|
||||
var accessLogGeoProviderFactory = func() (geoip.GeoIPService, error) {
|
||||
return geoip.NewMaxMindGeoIPService()
|
||||
}
|
||||
|
||||
type accessLogRegionResolver struct {
|
||||
provider geoip.GeoIPService
|
||||
cache map[string]string
|
||||
}
|
||||
|
||||
func newAccessLogRegionResolver() (*accessLogRegionResolver, error) {
|
||||
provider, err := accessLogGeoProviderFactory()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &accessLogRegionResolver{
|
||||
provider: provider,
|
||||
cache: make(map[string]string),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Close() {
|
||||
if r == nil || r.provider == nil {
|
||||
return
|
||||
}
|
||||
if err := r.provider.Close(); err != nil {
|
||||
slog.Warn("close access log geo provider failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Resolve(rawIP string) string {
|
||||
if r == nil || r.provider == nil {
|
||||
return ""
|
||||
}
|
||||
normalizedIP := normalizeAccessLogIP(rawIP)
|
||||
if normalizedIP == "" {
|
||||
return ""
|
||||
}
|
||||
if cached, ok := r.cache[normalizedIP]; ok {
|
||||
return cached
|
||||
}
|
||||
|
||||
info, err := r.provider.GetGeoInfo(net.ParseIP(normalizedIP))
|
||||
if err != nil || info == nil {
|
||||
r.cache[normalizedIP] = ""
|
||||
return ""
|
||||
}
|
||||
|
||||
region := strings.TrimSpace(info.Name)
|
||||
if region == "" {
|
||||
region = strings.TrimSpace(info.ISOCode)
|
||||
}
|
||||
r.cache[normalizedIP] = region
|
||||
return region
|
||||
}
|
||||
|
||||
func normalizeAccessLogIP(raw string) string {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
trimmed = strings.TrimPrefix(trimmed, "[")
|
||||
trimmed = strings.TrimSuffix(trimmed, "]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(raw))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
host = strings.TrimPrefix(host, "[")
|
||||
host = strings.TrimSuffix(host, "]")
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,320 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestListAccessLogsIncludesSummaryTotals(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now()
|
||||
if err := model.DB.Create(&model.Node{
|
||||
NodeID: "node-a",
|
||||
Name: "edge-a",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed node-a: %v", err)
|
||||
}
|
||||
if err := model.DB.Create(&model.Node{
|
||||
NodeID: "node-b",
|
||||
Name: "edge-b",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed node-b: %v", err)
|
||||
}
|
||||
|
||||
logs := []*model.NodeAccessLog{
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-5 * time.Minute),
|
||||
RemoteAddr: "1.1.1.1",
|
||||
Region: "United States",
|
||||
Host: "a.example.com",
|
||||
Path: "/alpha",
|
||||
StatusCode: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-4 * time.Minute),
|
||||
RemoteAddr: "2.2.2.2",
|
||||
Region: "China",
|
||||
Host: "a.example.com",
|
||||
Path: "/beta",
|
||||
StatusCode: 404,
|
||||
},
|
||||
{
|
||||
NodeID: "node-b",
|
||||
LoggedAt: now.Add(-3 * time.Minute),
|
||||
RemoteAddr: "1.1.1.1",
|
||||
Region: "United States",
|
||||
Host: "b.example.com",
|
||||
Path: "/gamma",
|
||||
StatusCode: 502,
|
||||
},
|
||||
{
|
||||
NodeID: "node-b",
|
||||
LoggedAt: now.Add(-2 * time.Minute),
|
||||
RemoteAddr: "",
|
||||
Host: "b.example.com",
|
||||
Path: "/delta",
|
||||
StatusCode: 200,
|
||||
},
|
||||
}
|
||||
seedNodeAccessLogs(t, logs)
|
||||
|
||||
result, err := ListAccessLogs(AccessLogQuery{Page: 0, PageSize: 2})
|
||||
if err != nil {
|
||||
t.Fatalf("ListAccessLogs failed: %v", err)
|
||||
}
|
||||
if result.TotalRecord != 4 {
|
||||
t.Fatalf("expected total_record=4, got %d", result.TotalRecord)
|
||||
}
|
||||
if result.TotalIP != 2 {
|
||||
t.Fatalf("expected total_ip=2, got %d", result.TotalIP)
|
||||
}
|
||||
if len(result.Items) != 2 {
|
||||
t.Fatalf("expected current page items=2, got %d", len(result.Items))
|
||||
}
|
||||
if result.Items[1].Region == "" {
|
||||
t.Fatalf("expected region to be returned, got %+v", result.Items[1])
|
||||
}
|
||||
if !result.HasMore {
|
||||
t.Fatal("expected has_more to be true")
|
||||
}
|
||||
|
||||
filtered, err := ListAccessLogs(AccessLogQuery{NodeID: "node-a", Page: 0, PageSize: 50})
|
||||
if err != nil {
|
||||
t.Fatalf("ListAccessLogs filtered failed: %v", err)
|
||||
}
|
||||
if filtered.TotalRecord != 2 {
|
||||
t.Fatalf("expected filtered total_record=2, got %d", filtered.TotalRecord)
|
||||
}
|
||||
if filtered.TotalIP != 2 {
|
||||
t.Fatalf("expected filtered total_ip=2, got %d", filtered.TotalIP)
|
||||
}
|
||||
if len(filtered.Items) != 2 {
|
||||
t.Fatalf("expected filtered items=2, got %d", len(filtered.Items))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListAccessLogsUsesDefaultPageSize(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now()
|
||||
if err := model.DB.Create(&model.Node{
|
||||
NodeID: "node-default-page-size",
|
||||
Name: "edge-default-page-size",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed node: %v", err)
|
||||
}
|
||||
|
||||
logs := make([]*model.NodeAccessLog, 0, 25)
|
||||
for index := range 25 {
|
||||
logs = append(logs, &model.NodeAccessLog{
|
||||
NodeID: "node-default-page-size",
|
||||
LoggedAt: now.Add(-time.Duration(index) * time.Minute),
|
||||
RemoteAddr: "1.1.1.1",
|
||||
Host: "example.com",
|
||||
Path: "/default-page-size",
|
||||
StatusCode: 200,
|
||||
})
|
||||
}
|
||||
seedNodeAccessLogs(t, logs)
|
||||
|
||||
result, err := ListAccessLogs(AccessLogQuery{})
|
||||
if err != nil {
|
||||
t.Fatalf("ListAccessLogs failed: %v", err)
|
||||
}
|
||||
if result.PageSize != 20 {
|
||||
t.Fatalf("expected default page_size=20, got %d", result.PageSize)
|
||||
}
|
||||
if len(result.Items) != 20 {
|
||||
t.Fatalf("expected current page items=20, got %d", len(result.Items))
|
||||
}
|
||||
if !result.HasMore {
|
||||
t.Fatal("expected has_more to be true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListFoldedAccessLogsAndIPSummaries(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Date(2026, 3, 19, 8, 12, 30, 0, time.UTC)
|
||||
if err := model.DB.Create(&model.Node{
|
||||
NodeID: "node-folded",
|
||||
Name: "edge-folded",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed node: %v", err)
|
||||
}
|
||||
logs := []*model.NodeAccessLog{
|
||||
{
|
||||
NodeID: "node-folded",
|
||||
LoggedAt: now.Add(-4 * time.Minute),
|
||||
RemoteAddr: "203.0.113.1",
|
||||
Host: "alpha.example.com",
|
||||
Path: "/first",
|
||||
StatusCode: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-folded",
|
||||
LoggedAt: now.Add(-3 * time.Minute),
|
||||
RemoteAddr: "203.0.113.1",
|
||||
Host: "alpha.example.com",
|
||||
Path: "/second",
|
||||
StatusCode: 502,
|
||||
},
|
||||
{
|
||||
NodeID: "node-folded",
|
||||
LoggedAt: now.Add(-2 * time.Minute),
|
||||
RemoteAddr: "203.0.113.2",
|
||||
Host: "beta.example.com",
|
||||
Path: "/third",
|
||||
StatusCode: 404,
|
||||
},
|
||||
}
|
||||
seedNodeAccessLogs(t, logs)
|
||||
|
||||
folded, err := ListFoldedAccessLogs(AccessLogQuery{
|
||||
NodeID: "node-folded",
|
||||
Page: 0,
|
||||
PageSize: 10,
|
||||
SortBy: "request_count",
|
||||
SortOrder: "desc",
|
||||
FoldMinutes: 5,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ListFoldedAccessLogs failed: %v", err)
|
||||
}
|
||||
if len(folded.Items) != 2 {
|
||||
t.Fatalf("expected two folded buckets, got %+v", folded.Items)
|
||||
}
|
||||
if folded.TotalRecord != 3 || folded.TotalBucket != 2 {
|
||||
t.Fatalf("unexpected folded totals: %+v", folded)
|
||||
}
|
||||
if folded.Items[0].RequestCount+folded.Items[1].RequestCount != 3 {
|
||||
t.Fatalf("unexpected folded request count sum: %+v", folded.Items)
|
||||
}
|
||||
if folded.Items[0].RequestCount != 2 {
|
||||
t.Fatalf("expected folded buckets to sort by request_count desc, got %+v", folded.Items)
|
||||
}
|
||||
|
||||
bucketIPs, err := ListFoldedAccessLogIPs(FoldedAccessLogIPQuery{
|
||||
NodeID: "node-folded",
|
||||
BucketStartedAt: folded.Items[0].BucketStartedAt.Format(time.RFC3339),
|
||||
FoldMinutes: 5,
|
||||
Page: 0,
|
||||
PageSize: 10,
|
||||
SortBy: "request_count",
|
||||
SortOrder: "desc",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ListFoldedAccessLogIPs failed: %v", err)
|
||||
}
|
||||
if bucketIPs.TotalIP != 1 || len(bucketIPs.Items) != 1 {
|
||||
t.Fatalf("expected one folded bucket IP row, got %+v", bucketIPs)
|
||||
}
|
||||
if bucketIPs.Items[0].RemoteAddr != "203.0.113.1" || bucketIPs.Items[0].RequestCount != 2 {
|
||||
t.Fatalf("unexpected top folded bucket IP row: %+v", bucketIPs.Items[0])
|
||||
}
|
||||
|
||||
ipSummaries, err := ListAccessLogIPSummaries(AccessLogIPSummaryQuery{
|
||||
NodeID: "node-folded",
|
||||
Page: 0,
|
||||
PageSize: 10,
|
||||
SortBy: "total_requests",
|
||||
SortOrder: "desc",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ListAccessLogIPSummaries failed: %v", err)
|
||||
}
|
||||
if len(ipSummaries.Items) != 2 {
|
||||
t.Fatalf("expected two ip summary rows, got %+v", ipSummaries.Items)
|
||||
}
|
||||
if ipSummaries.Items[0].RemoteAddr != "203.0.113.1" || ipSummaries.Items[0].TotalRequests != 2 {
|
||||
t.Fatalf("unexpected top ip summary row: %+v", ipSummaries.Items[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupAccessLogsDeletesExpiredData(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
seedNodeAccessLogs(t, []*model.NodeAccessLog{
|
||||
{
|
||||
NodeID: "node-cleanup",
|
||||
LoggedAt: now.Add(-10 * 24 * time.Hour),
|
||||
RemoteAddr: "203.0.113.9",
|
||||
Host: "cleanup.example.com",
|
||||
Path: "/old",
|
||||
StatusCode: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-cleanup",
|
||||
LoggedAt: now.Add(-2 * 24 * time.Hour),
|
||||
RemoteAddr: "203.0.113.10",
|
||||
Host: "cleanup.example.com",
|
||||
Path: "/recent",
|
||||
StatusCode: 200,
|
||||
},
|
||||
})
|
||||
|
||||
result, err := CleanupAccessLogs(AccessLogCleanupInput{RetentionDays: 7})
|
||||
if err != nil {
|
||||
t.Fatalf("CleanupAccessLogs failed: %v", err)
|
||||
}
|
||||
if result.DeletedCount != 1 {
|
||||
t.Fatalf("expected 1 deleted record, got %+v", result)
|
||||
}
|
||||
|
||||
remaining, err := ListAccessLogs(AccessLogQuery{Page: 0, PageSize: 10, NodeID: "node-cleanup"})
|
||||
if err != nil {
|
||||
t.Fatalf("ListAccessLogs failed after cleanup: %v", err)
|
||||
}
|
||||
if len(remaining.Items) != 1 || remaining.Items[0].Path != "/recent" {
|
||||
t.Fatalf("unexpected remaining logs after cleanup: %+v", remaining.Items)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersistNodeAccessLogsTruncatesLongPath(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
longPath := "/" + strings.Repeat("a", 140)
|
||||
reportedAt := time.Now().UTC()
|
||||
if err := persistNodeAccessLogs(model.DB, "node-truncate", []AgentNodeAccessLog{
|
||||
{
|
||||
LoggedAtUnix: reportedAt.Unix(),
|
||||
RemoteAddr: "203.0.113.10",
|
||||
Host: "truncate.example.com",
|
||||
Path: longPath,
|
||||
StatusCode: 200,
|
||||
},
|
||||
}, reportedAt); err != nil {
|
||||
t.Fatalf("persistNodeAccessLogs failed: %v", err)
|
||||
}
|
||||
|
||||
logs, err := model.ListNodeAccessLogs(model.NodeAccessLogQuery{
|
||||
NodeID: "node-truncate",
|
||||
Page: 0,
|
||||
PageSize: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodeAccessLogs failed: %v", err)
|
||||
}
|
||||
if len(logs) != 1 {
|
||||
t.Fatalf("expected one stored log, got %+v", logs)
|
||||
}
|
||||
if got := len([]rune(logs[0].Path)); got != nodeAccessLogPathMaxLength {
|
||||
t.Fatalf("expected truncated path length %d, got %d (%q)", nodeAccessLogPathMaxLength, got, logs[0].Path)
|
||||
}
|
||||
}
|
||||
|
||||
func seedNodeAccessLogs(t *testing.T, logs []*model.NodeAccessLog) {
|
||||
t.Helper()
|
||||
for _, item := range logs {
|
||||
if err := model.DB.Create(item).Error; err != nil {
|
||||
t.Fatalf("failed to seed access log: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,529 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
NodeStatusOnline = "online"
|
||||
NodeStatusOffline = "offline"
|
||||
NodeStatusPending = "pending"
|
||||
ApplyResultOK = "success"
|
||||
ApplyResultWarning = "warning"
|
||||
ApplyResultFailed = "failed"
|
||||
OpenrestyStatusHealthy = "healthy"
|
||||
OpenrestyStatusUnhealthy = "unhealthy"
|
||||
OpenrestyStatusUnknown = "unknown"
|
||||
)
|
||||
|
||||
type AgentNodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"ext_version"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastError string `json:"last_error"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
Profile *AgentNodeSystemProfile `json:"profile,omitempty"`
|
||||
Snapshot *AgentNodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
OpenrestyObservation *AgentNodeOpenrestyObservation `json:"openresty_observation,omitempty"`
|
||||
TrafficReport *AgentNodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []AgentNodeAccessLog `json:"access_logs,omitempty"`
|
||||
BufferedObservability []AgentBufferedObservabilityRecord `json:"buffered_observability,omitempty"`
|
||||
HealthEvents []AgentNodeHealthEvent `json:"health_events"`
|
||||
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
|
||||
}
|
||||
|
||||
type ApplyLogPayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Version string `json:"version"`
|
||||
Result string `json:"result"`
|
||||
Message string `json:"message"`
|
||||
Checksum string `json:"checksum"`
|
||||
MainConfigChecksum string `json:"main_config_checksum"`
|
||||
RouteConfigChecksum string `json:"route_config_checksum"`
|
||||
SupportFileCount int `json:"support_file_count"`
|
||||
}
|
||||
|
||||
type ApplyLogListQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
PageNo int `json:"pageNo"`
|
||||
PageSize int `json:"pageSize"`
|
||||
}
|
||||
|
||||
type ApplyLogListResult struct {
|
||||
Rows []*model.ApplyLog `json:"rows"`
|
||||
Current int `json:"current"`
|
||||
Total int `json:"total"`
|
||||
TotalPage int `json:"totalPage"`
|
||||
}
|
||||
|
||||
type ApplyLogCleanupInput struct {
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
}
|
||||
|
||||
type ApplyLogCleanupResult struct {
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
Cutoff *time.Time `json:"cutoff,omitempty"`
|
||||
}
|
||||
|
||||
type AgentConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
SourceConfigJSON string `json:"source_config_json"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type AgentSettings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
UpdateRepo string `json:"update_repo"`
|
||||
UpdateNow bool `json:"update_now"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
RestartOpenrestyNow bool `json:"restart_openresty_now"`
|
||||
}
|
||||
|
||||
type ActiveConfigMeta struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
type HeartbeatResponse struct {
|
||||
Node *model.Node `json:"node"`
|
||||
AgentSettings *AgentSettings `json:"agent_settings"`
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
WAFIPGroups []AgentWAFIPGroup `json:"waf_ip_groups,omitempty"`
|
||||
}
|
||||
|
||||
type AgentWAFIPGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IPList []string `json:"ip_list"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
type AgentWAFIPGroupSyncInput struct {
|
||||
IDs []uint `json:"ids"`
|
||||
Checksums map[string]string `json:"checksums"`
|
||||
}
|
||||
|
||||
type AgentWAFIPGroupSyncResult struct {
|
||||
Groups []AgentWAFIPGroup `json:"groups"`
|
||||
}
|
||||
|
||||
type NodeView struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
IPManualOverride bool `json:"ip_manual_override"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
GeoManualOverride bool `json:"geo_manual_override"`
|
||||
AccessToken string `json:"access_token"`
|
||||
AutoUpdateEnabled bool `json:"auto_update_enabled"`
|
||||
UpdateRequested bool `json:"update_requested"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
RestartOpenrestyRequested bool `json:"restart_openresty_requested"`
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"ext_version"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
Status string `json:"status"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastSeenAt any `json:"last_seen_at"`
|
||||
LastError string `json:"last_error"`
|
||||
LatestApplyResult string `json:"latest_apply_result"`
|
||||
LatestApplyMessage string `json:"latest_apply_message"`
|
||||
LatestApplyChecksum string `json:"latest_apply_checksum"`
|
||||
LatestMainConfigChecksum string `json:"latest_main_config_checksum"`
|
||||
LatestRouteConfigChecksum string `json:"latest_route_config_checksum"`
|
||||
LatestSupportFileCount int `json:"latest_support_file_count"`
|
||||
LatestApplyAt *time.Time `json:"latest_apply_at"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
// TunnelRelay fields
|
||||
NodeType string `json:"node_type"`
|
||||
RelayBindPort int `json:"relay_bind_port"`
|
||||
RelayVhostHTTPPort int `json:"relay_vhost_http_port"`
|
||||
RelayAgentAccessAddr string `json:"relay_agent_access_addr"`
|
||||
RelayClientAccessAddr string `json:"relay_client_access_addr"`
|
||||
RelayClientProxyURL string `json:"relay_client_proxy_url"`
|
||||
RelayStatus string `json:"relay_status"`
|
||||
RelayWebServerEnabled bool `json:"relay_web_server_enabled"`
|
||||
}
|
||||
|
||||
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
|
||||
payload = normalizeAgentNodePayload(payload)
|
||||
if err := validateAgentNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
previous := *node
|
||||
updateNow := node.UpdateRequested
|
||||
restartOpenrestyNow := node.RestartOpenrestyRequested
|
||||
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(node.UpdateTag)
|
||||
applyNodeRuntime(node, payload, true)
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = ReleaseChannelStable.String()
|
||||
node.UpdateTag = ""
|
||||
node.RestartOpenrestyRequested = false
|
||||
changes := collectNodeHeartbeatChanges(&previous, node)
|
||||
if len(changes) > 0 {
|
||||
if err := model.DB.Model(node).Updates(changes).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
persistHeartbeatObservability(node.NodeID, payload, node.LastSeenAt)
|
||||
activeConfig, err := GetActiveConfigMetaForAgent()
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
wafIPGroups, err := ChangedWAFIPGroupsForAgent(nil, payload.WAFIPGroupChecksums)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &HeartbeatResponse{
|
||||
Node: node,
|
||||
AgentSettings: buildAgentSettings(node, updateNow, updateChannel.String(), updateTag, restartOpenrestyNow),
|
||||
ActiveConfig: activeConfig,
|
||||
WAFIPGroups: wafIPGroups,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildAgentSettings(node *model.Node, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *AgentSettings {
|
||||
autoUpdate := false
|
||||
if node != nil {
|
||||
autoUpdate = node.AutoUpdateEnabled
|
||||
}
|
||||
if strings.TrimSpace(updateChannel) == "" {
|
||||
updateChannel = ReleaseChannelStable.String()
|
||||
}
|
||||
return &AgentSettings{
|
||||
HeartbeatInterval: common.AgentHeartbeatInterval,
|
||||
WebsocketUpgradeEnabled: common.AgentWebsocketUpgradeEnabled,
|
||||
AutoUpdate: autoUpdate,
|
||||
UpdateRepo: common.AgentUpdateRepo,
|
||||
UpdateNow: updateNow,
|
||||
UpdateChannel: updateChannel,
|
||||
UpdateTag: strings.TrimSpace(updateTag),
|
||||
RestartOpenrestyNow: restartOpenrestyNow,
|
||||
}
|
||||
}
|
||||
|
||||
func GetActiveConfigMetaForAgent() (*ActiveConfigMeta, error) {
|
||||
version, err := model.GetActiveConfigVersion()
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &ActiveConfigMeta{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func GetActiveConfigForAgent() (*AgentConfigResponse, error) {
|
||||
version, err := model.GetActiveConfigVersion()
|
||||
if err != nil {
|
||||
slog.Error("agent requested active config but no active version is available")
|
||||
return nil, err
|
||||
}
|
||||
var supportFiles []SupportFile
|
||||
if version.SupportFilesJSON != "" {
|
||||
if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
slog.Debug("agent fetched active config", "version", version.Version, "checksum", version.Checksum)
|
||||
return &AgentConfigResponse{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
SourceConfigJSON: version.SnapshotJSON,
|
||||
SupportFiles: sourceSupportFiles(supportFiles),
|
||||
CreatedAt: version.CreatedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
|
||||
payload.Result = strings.ToLower(payload.Result)
|
||||
utils.TrimStringFields(
|
||||
&payload.NodeID,
|
||||
&payload.Version,
|
||||
&payload.Result,
|
||||
&payload.Message,
|
||||
&payload.Checksum,
|
||||
&payload.MainConfigChecksum,
|
||||
&payload.RouteConfigChecksum,
|
||||
)
|
||||
payload.Message = truncateForDatabase(payload.Message, 16000)
|
||||
return payload
|
||||
}
|
||||
|
||||
func ReportApplyLog(payload ApplyLogPayload) (*model.ApplyLog, error) {
|
||||
now := time.Now()
|
||||
payload = normalizeApplyLogPayload(payload)
|
||||
if payload.NodeID == "" {
|
||||
return nil, errors.New("node_id 不能为空")
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return nil, errors.New("version 不能为空")
|
||||
}
|
||||
if payload.Result != ApplyResultOK && payload.Result != ApplyResultWarning && payload.Result != ApplyResultFailed {
|
||||
return nil, errors.New("result 仅支持 success、warning 或 failed")
|
||||
}
|
||||
slog.Debug("agent apply log received", "node_id", payload.NodeID, "version", payload.Version, "result", payload.Result)
|
||||
|
||||
log := &model.ApplyLog{
|
||||
NodeID: payload.NodeID,
|
||||
Version: payload.Version,
|
||||
Result: payload.Result,
|
||||
Message: payload.Message,
|
||||
Checksum: payload.Checksum,
|
||||
MainConfigChecksum: payload.MainConfigChecksum,
|
||||
RouteConfigChecksum: payload.RouteConfigChecksum,
|
||||
SupportFileCount: payload.SupportFileCount,
|
||||
CreatedAt: now,
|
||||
}
|
||||
err := model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
node := &model.Node{}
|
||||
if err := tx.Where("node_id = ?", payload.NodeID).First(node).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
node.Status = NodeStatusOnline
|
||||
node.LastSeenAt = now
|
||||
if payload.Result == ApplyResultOK {
|
||||
node.CurrentVersion = payload.Version
|
||||
node.LastError = ""
|
||||
} else {
|
||||
node.LastError = payload.Message
|
||||
}
|
||||
if err := tx.Create(log).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(node).Select("status", "last_seen_at", "current_version", "last_error").Updates(node).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if payload.Result == ApplyResultOK {
|
||||
slog.Debug("agent apply reported success", "node_id", payload.NodeID, "version", payload.Version)
|
||||
} else if payload.Result == ApplyResultWarning {
|
||||
slog.Warn("agent apply reported warning", "node_id", payload.NodeID, "version", payload.Version, "message", payload.Message)
|
||||
} else {
|
||||
slog.Error("agent apply reported failure", "node_id", payload.NodeID, "version", payload.Version, "message", payload.Message)
|
||||
}
|
||||
return log, nil
|
||||
}
|
||||
|
||||
func ListNodeViews() ([]*NodeView, error) {
|
||||
nodes, err := model.ListNodes()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeIDs := make([]string, 0, len(nodes))
|
||||
for _, node := range nodes {
|
||||
nodeIDs = append(nodeIDs, node.NodeID)
|
||||
}
|
||||
latestLogs, err := model.GetLatestApplyLogsByNodeIDs(nodeIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]*NodeView, 0, len(nodes))
|
||||
for _, node := range nodes {
|
||||
computedStatus := computeNodeStatus(node)
|
||||
view := buildNodeView(node)
|
||||
view.Status = computedStatus
|
||||
if log, ok := latestLogs[node.NodeID]; ok {
|
||||
view.LatestApplyResult = log.Result
|
||||
view.LatestApplyMessage = log.Message
|
||||
view.LatestApplyChecksum = log.Checksum
|
||||
view.LatestMainConfigChecksum = log.MainConfigChecksum
|
||||
view.LatestRouteConfigChecksum = log.RouteConfigChecksum
|
||||
view.LatestSupportFileCount = log.SupportFileCount
|
||||
view.LatestApplyAt = &log.CreatedAt
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func truncateForDatabase(value string, max int) string {
|
||||
if max <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(value))
|
||||
if len(runes) <= max {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:max])
|
||||
}
|
||||
|
||||
const (
|
||||
defaultApplyLogPageSize = 20
|
||||
maxApplyLogPageSize = 200
|
||||
maxApplyLogRetentionDays = 3650
|
||||
)
|
||||
|
||||
func ListApplyLogsPage(input ApplyLogListQuery) (*ApplyLogListResult, error) {
|
||||
pageNo := normalizeApplyLogPageNo(input.PageNo)
|
||||
pageSize := normalizeApplyLogPageSize(input.PageSize)
|
||||
nodeID := strings.TrimSpace(input.NodeID)
|
||||
rows, err := model.ListApplyLogs(model.ApplyLogQuery{
|
||||
NodeID: nodeID,
|
||||
PageNo: pageNo,
|
||||
PageSize: pageSize,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total, err := model.CountApplyLogs(nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalPage := 0
|
||||
if total > 0 {
|
||||
totalPage = int((total + int64(pageSize) - 1) / int64(pageSize))
|
||||
}
|
||||
return &ApplyLogListResult{
|
||||
Rows: rows,
|
||||
Current: pageNo,
|
||||
Total: int(total),
|
||||
TotalPage: totalPage,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func CleanupApplyLogs(input ApplyLogCleanupInput) (*ApplyLogCleanupResult, error) {
|
||||
if input.DeleteAll {
|
||||
deleted, err := model.DeleteAllApplyLogs()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ApplyLogCleanupResult{
|
||||
DeleteAll: true,
|
||||
DeletedCount: deleted,
|
||||
}, nil
|
||||
}
|
||||
if input.RetentionDays <= 0 || input.RetentionDays > maxApplyLogRetentionDays {
|
||||
return nil, errors.New("retention_days 必须在 1 到 3650 之间")
|
||||
}
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
|
||||
deleted, err := model.DeleteApplyLogsBefore(cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ApplyLogCleanupResult{
|
||||
RetentionDays: input.RetentionDays,
|
||||
DeletedCount: deleted,
|
||||
Cutoff: &cutoff,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeApplyLogPageNo(pageNo int) int {
|
||||
if pageNo <= 0 {
|
||||
return 1
|
||||
}
|
||||
return pageNo
|
||||
}
|
||||
|
||||
func normalizeApplyLogPageSize(pageSize int) int {
|
||||
if pageSize <= 0 {
|
||||
return defaultApplyLogPageSize
|
||||
}
|
||||
if pageSize > maxApplyLogPageSize {
|
||||
return maxApplyLogPageSize
|
||||
}
|
||||
return pageSize
|
||||
}
|
||||
|
||||
func computeNodeStatus(node *model.Node) string {
|
||||
if node == nil {
|
||||
return NodeStatusOffline
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" && IsRelayWSConnected(node.NodeID) {
|
||||
return NodeStatusOnline
|
||||
}
|
||||
if node.NodeType == "tunnel_client" && IsFlaredWSConnected(node.NodeID) {
|
||||
return NodeStatusOnline
|
||||
}
|
||||
if IsAgentWSConnected(node.NodeID) {
|
||||
return NodeStatusOnline
|
||||
}
|
||||
if node.LastSeenAt.IsZero() {
|
||||
return NodeStatusPending
|
||||
}
|
||||
if time.Since(node.LastSeenAt) > common.NodeOfflineThreshold {
|
||||
return NodeStatusOffline
|
||||
}
|
||||
return NodeStatusOnline
|
||||
}
|
||||
|
||||
func collectNodeHeartbeatChanges(previous *model.Node, current *model.Node) map[string]any {
|
||||
if previous == nil || current == nil {
|
||||
return map[string]any{}
|
||||
}
|
||||
changes := make(map[string]any)
|
||||
appendIfChanged := func(key string, before any, after any) {
|
||||
if before != after {
|
||||
changes[key] = after
|
||||
}
|
||||
}
|
||||
appendIfChanged("name", previous.Name, current.Name)
|
||||
appendIfChanged("ip", previous.IP, current.IP)
|
||||
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
|
||||
appendIfChanged("version", previous.Version, current.Version)
|
||||
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
|
||||
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
|
||||
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
|
||||
appendIfChanged("status", previous.Status, current.Status)
|
||||
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
|
||||
appendIfChanged("last_error", previous.LastError, current.LastError)
|
||||
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
|
||||
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
|
||||
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
|
||||
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
|
||||
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
|
||||
changes["geo_latitude"] = current.GeoLatitude
|
||||
}
|
||||
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
|
||||
changes["geo_longitude"] = current.GeoLongitude
|
||||
}
|
||||
if !previous.LastSeenAt.Equal(current.LastSeenAt) {
|
||||
changes["last_seen_at"] = current.LastSeenAt
|
||||
}
|
||||
return changes
|
||||
}
|
||||
|
||||
func coordinatesEqual(before *float64, after *float64) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return *before == *after
|
||||
}
|
||||
@@ -0,0 +1,522 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestGetActiveConfigForAgentIncludesWAFConfig(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "waf-agent.example.com",
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
if _, err := PublishConfigVersion("root", false); err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
|
||||
activeConfig, err := GetActiveConfigForAgent()
|
||||
if err != nil {
|
||||
t.Fatalf("GetActiveConfigForAgent failed: %v", err)
|
||||
}
|
||||
|
||||
for _, file := range activeConfig.SupportFiles {
|
||||
if file.Path == "waf_config.json" {
|
||||
t.Fatal("agent config should not receive rendered waf_config.json")
|
||||
}
|
||||
}
|
||||
if !strings.Contains(activeConfig.SourceConfigJSON, `"waf"`) {
|
||||
t.Fatal("expected agent config source json to include WAF source configuration")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "agent-waf-ip-group",
|
||||
Domains: []string{"agent-waf-ip-group.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
ipGroup, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "agent runtime group",
|
||||
Type: WAFIPGroupTypeManual,
|
||||
Enabled: true,
|
||||
IPList: []string{"203.0.113.44"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
ruleGroup, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "agent refs",
|
||||
Enabled: true,
|
||||
IPBlacklistGroups: []uint{ipGroup.ID},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if _, err = ReplaceWAFSiteRuleGroups(route.ID, []uint{ruleGroup.ID}); err != nil {
|
||||
t.Fatalf("ReplaceWAFSiteRuleGroups failed: %v", err)
|
||||
}
|
||||
if _, err = PublishConfigVersion("root", false); err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
|
||||
groups, err := ChangedWAFIPGroupsForAgent(nil, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("ChangedWAFIPGroupsForAgent failed: %v", err)
|
||||
}
|
||||
if len(groups) != 1 || groups[0].ID != ipGroup.ID || groups[0].IPList[0] != "203.0.113.44" || groups[0].Checksum == "" {
|
||||
t.Fatalf("unexpected changed groups: %#v", groups)
|
||||
}
|
||||
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
|
||||
same, err := ChangedWAFIPGroupsForAgent(nil, map[string]string{groupKey: groups[0].Checksum})
|
||||
if err != nil {
|
||||
t.Fatalf("ChangedWAFIPGroupsForAgent with checksum failed: %v", err)
|
||||
}
|
||||
if len(same) != 0 {
|
||||
t.Fatalf("expected no delta for matching checksum, got %#v", same)
|
||||
}
|
||||
updated, err := UpdateWAFIPGroup(ipGroup.ID, WAFIPGroupInput{
|
||||
Name: "agent runtime group",
|
||||
Type: WAFIPGroupTypeManual,
|
||||
Enabled: true,
|
||||
IPList: []string{"203.0.113.45"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
delta, err := ChangedWAFIPGroupsForAgent(nil, map[string]string{groupKey: groups[0].Checksum})
|
||||
if err != nil {
|
||||
t.Fatalf("ChangedWAFIPGroupsForAgent after update failed: %v", err)
|
||||
}
|
||||
if len(delta) != 1 || delta[0].ID != updated.ID || delta[0].IPList[0] != "203.0.113.45" || delta[0].Checksum == groups[0].Checksum {
|
||||
t.Fatalf("expected updated group delta, got %#v", delta)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterNodeWithAccessToken(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// 1. Success path
|
||||
latitude := 31.2304
|
||||
longitude := 121.4737
|
||||
node, err := CreateNode(NodeInput{
|
||||
Name: "reserved-node-1",
|
||||
IP: "192.168.1.10",
|
||||
GeoManualOverride: true,
|
||||
GeoName: "Shanghai",
|
||||
GeoLatitude: &latitude,
|
||||
GeoLongitude: &longitude,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create node: %v", err)
|
||||
}
|
||||
|
||||
stored, err := model.GetNodeByID(node.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to fetch stored node: %v", err)
|
||||
}
|
||||
|
||||
payload := AgentNodePayload{
|
||||
Name: "payload-name-should-be-ignored",
|
||||
IP: "192.168.1.20",
|
||||
Version: "v1.0.1",
|
||||
ExtVersion: "1.27.1.3",
|
||||
OpenrestyStatus: "healthy",
|
||||
}
|
||||
|
||||
resp, err := RegisterNodeWithAccessToken(stored, payload)
|
||||
if err != nil {
|
||||
t.Fatalf("RegisterNodeWithAccessToken failed: %v", err)
|
||||
}
|
||||
|
||||
if resp.NodeID != stored.NodeID || resp.AccessToken != stored.AccessToken || resp.Name != "reserved-node-1" {
|
||||
t.Errorf("unexpected response: %+v", resp)
|
||||
}
|
||||
|
||||
// Verify that the node was updated in the DB
|
||||
updated, err := model.GetNodeByID(node.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to fetch updated node: %v", err)
|
||||
}
|
||||
if updated.Version != "v1.0.1" || updated.ExtVersion != "1.27.1.3" || updated.OpenrestyStatus != "healthy" {
|
||||
t.Errorf("node attributes were not updated: %+v", updated)
|
||||
}
|
||||
// Name should be preserved since preserveName is true
|
||||
if updated.Name != "reserved-node-1" {
|
||||
t.Errorf("expected name to be preserved, got %s", updated.Name)
|
||||
}
|
||||
|
||||
// 2. Fail path - Nil Node
|
||||
_, err = RegisterNodeWithAccessToken(nil, payload)
|
||||
if err == nil || !strings.Contains(err.Error(), "节点不存在") {
|
||||
t.Errorf("expected error '节点不存在', got %v", err)
|
||||
}
|
||||
|
||||
// 3. Fail path - Invalid Payload (empty IP)
|
||||
badPayload := payload
|
||||
badPayload.IP = ""
|
||||
_, err = RegisterNodeWithAccessToken(stored, badPayload)
|
||||
if err == nil || !strings.Contains(err.Error(), "ip 不能为空") {
|
||||
t.Errorf("expected error 'ip 不能为空', got %v", err)
|
||||
}
|
||||
|
||||
// 4. Name update if empty
|
||||
emptyNameNode := &model.Node{
|
||||
NodeID: "node-empty-name",
|
||||
Name: "",
|
||||
AccessToken: "empty-name-token",
|
||||
}
|
||||
if err := emptyNameNode.Insert(); err != nil {
|
||||
t.Fatalf("failed to insert emptyNameNode: %v", err)
|
||||
}
|
||||
payloadWithName := payload
|
||||
payloadWithName.Name = "filled-name"
|
||||
payloadWithName.IP = "192.168.1.30"
|
||||
_, err = RegisterNodeWithAccessToken(emptyNameNode, payloadWithName)
|
||||
if err != nil {
|
||||
t.Fatalf("RegisterNodeWithAccessToken empty name node failed: %v", err)
|
||||
}
|
||||
updatedEmptyName, err := model.GetNodeByNodeID("node-empty-name")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to fetch updatedEmptyName: %v", err)
|
||||
}
|
||||
if updatedEmptyName.Name != "filled-name" {
|
||||
t.Errorf("expected name to be filled, got %s", updatedEmptyName.Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterNodeWithDiscovery(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// 1. Success path
|
||||
payload := AgentNodePayload{
|
||||
Name: "discovery-node",
|
||||
IP: "192.168.2.10",
|
||||
Version: "v1.0.0",
|
||||
ExtVersion: "1.27.1.3",
|
||||
OpenrestyStatus: "healthy",
|
||||
}
|
||||
|
||||
resp, err := RegisterNodeWithDiscovery(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("RegisterNodeWithDiscovery failed: %v", err)
|
||||
}
|
||||
|
||||
if resp.NodeID == "" || resp.AccessToken == "" || resp.Name != "discovery-node" {
|
||||
t.Errorf("unexpected response: %+v", resp)
|
||||
}
|
||||
|
||||
// Verify database persistence
|
||||
node, err := model.GetNodeByNodeID(resp.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to fetch node: %v", err)
|
||||
}
|
||||
if node.IP != "192.168.2.10" || node.Version != "v1.0.0" || node.Name != "discovery-node" {
|
||||
t.Errorf("unexpected stored node data: %+v", node)
|
||||
}
|
||||
|
||||
// 2. Name fallback if payload name is empty
|
||||
payloadNoName := payload
|
||||
payloadNoName.Name = ""
|
||||
payloadNoName.IP = "192.168.2.20"
|
||||
respNoName, err := RegisterNodeWithDiscovery(payloadNoName)
|
||||
if err != nil {
|
||||
t.Fatalf("RegisterNodeWithDiscovery no name failed: %v", err)
|
||||
}
|
||||
nodeNoName, err := model.GetNodeByNodeID(respNoName.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to fetch no-name node: %v", err)
|
||||
}
|
||||
if nodeNoName.Name != respNoName.NodeID {
|
||||
t.Errorf("expected name fallback to NodeID, got %s", nodeNoName.Name)
|
||||
}
|
||||
|
||||
// 3. Fail path - Invalid Payload (empty AgentVersion)
|
||||
badPayload := payload
|
||||
badPayload.Version = ""
|
||||
_, err = RegisterNodeWithDiscovery(badPayload)
|
||||
if err == nil || !strings.Contains(err.Error(), "version 不能为空") {
|
||||
t.Errorf("expected error 'version 不能为空', got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReportApplyLog_Success(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// Seed node
|
||||
node := &model.Node{
|
||||
NodeID: "node-apply-1",
|
||||
Name: "apply-edge",
|
||||
IP: "192.168.3.10",
|
||||
AccessToken: "apply-token",
|
||||
Version: "v1.0.0",
|
||||
Status: NodeStatusOffline,
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to insert node: %v", err)
|
||||
}
|
||||
|
||||
payload := ApplyLogPayload{
|
||||
NodeID: "node-apply-1",
|
||||
Version: "20260531-001",
|
||||
Result: "success",
|
||||
Message: "Configuration applied successfully",
|
||||
Checksum: "chk-1",
|
||||
MainConfigChecksum: "m-chk-1",
|
||||
RouteConfigChecksum: "r-chk-1",
|
||||
SupportFileCount: 3,
|
||||
}
|
||||
|
||||
log, err := ReportApplyLog(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("ReportApplyLog failed: %v", err)
|
||||
}
|
||||
|
||||
if log.NodeID != "node-apply-1" || log.Result != "success" || log.Message != "Configuration applied successfully" {
|
||||
t.Errorf("unexpected returned log: %+v", log)
|
||||
}
|
||||
|
||||
// Verify that the node status and current version are updated in the DB
|
||||
updatedNode, err := model.GetNodeByNodeID("node-apply-1")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload node: %v", err)
|
||||
}
|
||||
if updatedNode.CurrentVersion != "20260531-001" || updatedNode.Status != NodeStatusOnline || updatedNode.LastError != "" {
|
||||
t.Errorf("node was not updated correctly: %+v", updatedNode)
|
||||
}
|
||||
|
||||
// Verify apply log is stored
|
||||
storedLogs, err := model.ListApplyLogs(model.ApplyLogQuery{NodeID: "node-apply-1", PageNo: 1, PageSize: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("ListApplyLogs failed: %v", err)
|
||||
}
|
||||
if len(storedLogs) != 1 || storedLogs[0].Checksum != "chk-1" {
|
||||
t.Errorf("expected 1 log, got: %d", len(storedLogs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestReportApplyLog_WarningAndFailure(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// Seed node
|
||||
node := &model.Node{
|
||||
NodeID: "node-apply-2",
|
||||
Name: "apply-edge-2",
|
||||
IP: "192.168.3.20",
|
||||
AccessToken: "apply-token-2",
|
||||
Version: "v1.0.0",
|
||||
CurrentVersion: "20260531-001", // Old version
|
||||
Status: NodeStatusOnline,
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to insert node: %v", err)
|
||||
}
|
||||
|
||||
// 1. Report Failure
|
||||
failPayload := ApplyLogPayload{
|
||||
NodeID: "node-apply-2",
|
||||
Version: "20260531-002", // Target failed version
|
||||
Result: "failed",
|
||||
Message: "reload process exited with code 1",
|
||||
}
|
||||
|
||||
_, err := ReportApplyLog(failPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("ReportApplyLog failed: %v", err)
|
||||
}
|
||||
|
||||
// Node CurrentVersion should NOT be updated. Node LastError should be updated.
|
||||
updatedNode, err := model.GetNodeByNodeID("node-apply-2")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload node: %v", err)
|
||||
}
|
||||
if updatedNode.CurrentVersion != "20260531-001" {
|
||||
t.Errorf("expected CurrentVersion to remain unchanged, got %s", updatedNode.CurrentVersion)
|
||||
}
|
||||
if updatedNode.LastError != "reload process exited with code 1" {
|
||||
t.Errorf("expected LastError to be set, got %s", updatedNode.LastError)
|
||||
}
|
||||
|
||||
// 2. Report Warning (e.g. rolled back to old version successfully)
|
||||
warningPayload := ApplyLogPayload{
|
||||
NodeID: "node-apply-2",
|
||||
Version: "20260531-002",
|
||||
Result: "warning",
|
||||
Message: "reload failed, rolled back to 20260531-001 successfully",
|
||||
}
|
||||
|
||||
_, err = ReportApplyLog(warningPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("ReportApplyLog warning failed: %v", err)
|
||||
}
|
||||
|
||||
updatedNodeWarning, err := model.GetNodeByNodeID("node-apply-2")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload node: %v", err)
|
||||
}
|
||||
// CurrentVersion remains 20260531-001. LastError is the warning message.
|
||||
if updatedNodeWarning.CurrentVersion != "20260531-001" {
|
||||
t.Errorf("expected CurrentVersion to remain unchanged, got %s", updatedNodeWarning.CurrentVersion)
|
||||
}
|
||||
if updatedNodeWarning.LastError != "reload failed, rolled back to 20260531-001 successfully" {
|
||||
t.Errorf("expected LastError to be warning message, got %s", updatedNodeWarning.LastError)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReportApplyLog_Failures(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// Seed node
|
||||
node := &model.Node{
|
||||
NodeID: "node-apply-3",
|
||||
Name: "apply-edge-3",
|
||||
IP: "192.168.3.30",
|
||||
AccessToken: "apply-token-3",
|
||||
Version: "v1.0.0",
|
||||
Status: NodeStatusOnline,
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to insert node: %v", err)
|
||||
}
|
||||
|
||||
// 1. Missing NodeID
|
||||
_, err := ReportApplyLog(ApplyLogPayload{Version: "v1", Result: "success"})
|
||||
if err == nil || !strings.Contains(err.Error(), "node_id 不能为空") {
|
||||
t.Errorf("expected empty node_id error, got %v", err)
|
||||
}
|
||||
|
||||
// 2. Missing Version
|
||||
_, err = ReportApplyLog(ApplyLogPayload{NodeID: "node-apply-3", Result: "success"})
|
||||
if err == nil || !strings.Contains(err.Error(), "version 不能为空") {
|
||||
t.Errorf("expected empty version error, got %v", err)
|
||||
}
|
||||
|
||||
// 3. Invalid Result
|
||||
_, err = ReportApplyLog(ApplyLogPayload{NodeID: "node-apply-3", Version: "v1", Result: "corrupted"})
|
||||
if err == nil || !strings.Contains(err.Error(), "result 仅支持 success、warning 或 failed") {
|
||||
t.Errorf("expected invalid result error, got %v", err)
|
||||
}
|
||||
|
||||
// 4. Non-existent NodeID
|
||||
_, err = ReportApplyLog(ApplyLogPayload{NodeID: "non-existent-node-xyz", Version: "v1", Result: "success"})
|
||||
if err == nil || !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Errorf("expected record not found error, got %v", err)
|
||||
}
|
||||
|
||||
// 5. Truncate excessively long message
|
||||
veryLongMsg := strings.Repeat("A", 20000)
|
||||
log, err := ReportApplyLog(ApplyLogPayload{
|
||||
NodeID: "node-apply-3",
|
||||
Version: "v1",
|
||||
Result: "success",
|
||||
Message: veryLongMsg,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ReportApplyLog with very long message failed: %v", err)
|
||||
}
|
||||
if len(log.Message) != 16000 {
|
||||
t.Errorf("expected message to be truncated to 16000, got %d", len(log.Message))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListAndCleanupApplyLogs(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// Seed node
|
||||
node := &model.Node{
|
||||
NodeID: "node-logs",
|
||||
Name: "logs-edge",
|
||||
IP: "192.168.4.10",
|
||||
AccessToken: "logs-token",
|
||||
Status: NodeStatusOnline,
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to insert node: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
// Seed logs of different ages
|
||||
logs := []model.ApplyLog{
|
||||
{NodeID: "node-logs", Version: "v1", Result: "success", Message: "1", CreatedAt: now.Add(-10 * 24 * time.Hour)}, // 10 days ago
|
||||
{NodeID: "node-logs", Version: "v2", Result: "success", Message: "2", CreatedAt: now.Add(-5 * 24 * time.Hour)}, // 5 days ago
|
||||
{NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now}, // Now
|
||||
}
|
||||
for i := range logs {
|
||||
if err := model.DB.Create(&logs[i]).Error; err != nil {
|
||||
t.Fatalf("failed to seed log: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 1. Test pagination using ListApplyLogsPage
|
||||
pageResult, err := ListApplyLogsPage(ApplyLogListQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 2,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ListApplyLogsPage failed: %v", err)
|
||||
}
|
||||
if pageResult.Total != 3 || len(pageResult.Rows) != 2 || pageResult.TotalPage != 2 {
|
||||
t.Errorf("unexpected pagination result: %+v", pageResult)
|
||||
}
|
||||
|
||||
// 2. Test Cleanup with RetentionDays = 7
|
||||
cleanupResult, err := CleanupApplyLogs(ApplyLogCleanupInput{
|
||||
DeleteAll: false,
|
||||
RetentionDays: 7,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CleanupApplyLogs failed: %v", err)
|
||||
}
|
||||
if cleanupResult.DeletedCount != 1 {
|
||||
t.Errorf("expected 1 log to be deleted, got %d", cleanupResult.DeletedCount)
|
||||
}
|
||||
|
||||
// Verify remaining logs: newer logs (v2 and v3) should still be in the DB
|
||||
remainingLogs, err := model.ListApplyLogs(model.ApplyLogQuery{NodeID: "node-logs", PageNo: 1, PageSize: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("ListApplyLogs failed: %v", err)
|
||||
}
|
||||
if len(remainingLogs) != 2 {
|
||||
t.Errorf("expected 2 remaining logs, got %d", len(remainingLogs))
|
||||
}
|
||||
|
||||
// 3. Test Cleanup with DeleteAll = true
|
||||
cleanupAll, err := CleanupApplyLogs(ApplyLogCleanupInput{
|
||||
DeleteAll: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CleanupApplyLogs deleteAll failed: %v", err)
|
||||
}
|
||||
if cleanupAll.DeletedCount != 2 {
|
||||
t.Errorf("expected 2 remaining logs to be deleted, got %d", cleanupAll.DeletedCount)
|
||||
}
|
||||
|
||||
// Verify DB is empty of apply logs
|
||||
finalLogs, err := model.ListApplyLogs(model.ApplyLogQuery{NodeID: "node-logs", PageNo: 1, PageSize: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("ListApplyLogs failed: %v", err)
|
||||
}
|
||||
if len(finalLogs) != 0 {
|
||||
t.Errorf("expected 0 remaining logs, got %d", len(finalLogs))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
const (
|
||||
AgentWSMessageTypeStatus = "status"
|
||||
AgentWSMessageTypeSettings = "settings"
|
||||
AgentWSMessageTypeActiveConfig = "active_config"
|
||||
AgentWSMessageTypeForceSyncConfig = "force_sync_config"
|
||||
AgentWSMessageTypeWAFIPGroups = "waf_ip_groups"
|
||||
AgentWSMessageTypePing = "ping"
|
||||
AgentWSMessageTypePong = "pong"
|
||||
|
||||
AgentWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__"
|
||||
)
|
||||
|
||||
type AgentWSInboundMessage struct {
|
||||
Type string `json:"type"`
|
||||
Payload json.RawMessage `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
type AgentWSBroadcastResult struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
ClientCount int `json:"client_count"`
|
||||
SuccessCount int `json:"success_count"`
|
||||
FailedNodes []string `json:"failed_nodes"`
|
||||
}
|
||||
|
||||
var DefaultAgentWSHub = NewWSHub("agent")
|
||||
|
||||
func RegisterAgentWSClient(nodeID string) *WSClient {
|
||||
return DefaultAgentWSHub.Register(nodeID)
|
||||
}
|
||||
|
||||
func UnregisterAgentWSClient(client *WSClient) {
|
||||
DefaultAgentWSHub.Unregister(client)
|
||||
}
|
||||
|
||||
func DisconnectAgentWSClient(nodeID string) {
|
||||
DefaultAgentWSHub.Disconnect(nodeID)
|
||||
}
|
||||
|
||||
func IsAgentWSConnected(nodeID string) bool {
|
||||
return DefaultAgentWSHub.IsConnected(nodeID)
|
||||
}
|
||||
|
||||
func SendAgentWSSettings(nodeID string, settings *AgentSettings) bool {
|
||||
if settings == nil {
|
||||
return false
|
||||
}
|
||||
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: AgentWSMessageTypeSettings,
|
||||
Payload: settings,
|
||||
})
|
||||
}
|
||||
|
||||
func SendAgentWSActiveConfig(nodeID string, activeConfig *ActiveConfigMeta) bool {
|
||||
if activeConfig == nil {
|
||||
return false
|
||||
}
|
||||
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: AgentWSMessageTypeActiveConfig,
|
||||
Payload: activeConfig,
|
||||
})
|
||||
}
|
||||
|
||||
func SendAgentWSForceSyncConfig(nodeID string, activeConfig *ActiveConfigMeta) bool {
|
||||
if activeConfig == nil {
|
||||
return false
|
||||
}
|
||||
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: AgentWSMessageTypeForceSyncConfig,
|
||||
Payload: activeConfig,
|
||||
})
|
||||
}
|
||||
|
||||
func SendAgentWSWAFIPGroups(nodeID string, groups []AgentWAFIPGroup) bool {
|
||||
if len(groups) == 0 {
|
||||
return false
|
||||
}
|
||||
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: AgentWSMessageTypeWAFIPGroups,
|
||||
Payload: groups,
|
||||
})
|
||||
}
|
||||
|
||||
func SendAgentWSPong(nodeID string) bool {
|
||||
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: AgentWSMessageTypePong,
|
||||
})
|
||||
}
|
||||
|
||||
func BroadcastAgentWSActiveConfig(activeConfig *ActiveConfigMeta) AgentWSBroadcastResult {
|
||||
if activeConfig == nil {
|
||||
slog.Debug("agent ws broadcast skipped because active config is nil")
|
||||
return AgentWSBroadcastResult{}
|
||||
}
|
||||
|
||||
res := DefaultAgentWSHub.Broadcast(WSMessage{
|
||||
Type: AgentWSMessageTypeActiveConfig,
|
||||
Payload: activeConfig,
|
||||
})
|
||||
|
||||
result := AgentWSBroadcastResult{
|
||||
Version: activeConfig.Version,
|
||||
Checksum: activeConfig.Checksum,
|
||||
ClientCount: res.ClientCount,
|
||||
SuccessCount: res.SuccessCount,
|
||||
FailedNodes: res.FailedIDs,
|
||||
}
|
||||
|
||||
slog.Debug("agent ws broadcast active config",
|
||||
"version", result.Version,
|
||||
"checksum", result.Checksum,
|
||||
"client_count", result.ClientCount,
|
||||
"success_count", result.SuccessCount,
|
||||
"failed_nodes", result.FailedNodes,
|
||||
)
|
||||
return result
|
||||
}
|
||||
|
||||
func BroadcastAgentWSWAFIPGroups(groups []AgentWAFIPGroup) WSBroadcastResult {
|
||||
if len(groups) == 0 {
|
||||
return WSBroadcastResult{}
|
||||
}
|
||||
result := DefaultAgentWSHub.Broadcast(WSMessage{
|
||||
Type: AgentWSMessageTypeWAFIPGroups,
|
||||
Payload: groups,
|
||||
})
|
||||
slog.Debug("agent ws broadcast waf ip groups",
|
||||
"group_count", len(groups),
|
||||
"client_count", result.ClientCount,
|
||||
"success_count", result.SuccessCount,
|
||||
"failed_nodes", result.FailedIDs,
|
||||
)
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,500 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type PublicAuthSource struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
AuthorizeURL string `json:"authorize_url"`
|
||||
IconURL string `json:"icon_url"`
|
||||
}
|
||||
|
||||
type OAuthProfile struct {
|
||||
ExternalID string
|
||||
ExternalUsername string
|
||||
DisplayName string
|
||||
Email string
|
||||
}
|
||||
|
||||
type OAuthCallbackResult struct {
|
||||
Status string `json:"status"`
|
||||
User *model.User `json:"user,omitempty"`
|
||||
}
|
||||
|
||||
type LinkExistingRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type PendingExternalAccount struct {
|
||||
AuthSourceID uint `json:"auth_source_id"`
|
||||
ExternalID string `json:"external_id"`
|
||||
ExternalUsername string `json:"external_username"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Email string `json:"email"`
|
||||
}
|
||||
|
||||
type oidcDiscovery struct {
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
UserInfoEndpoint string `json:"userinfo_endpoint"`
|
||||
JWKSURI string `json:"jwks_uri"`
|
||||
Issuer string `json:"issuer"`
|
||||
}
|
||||
|
||||
type oauthTokenResponse struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
IDToken string `json:"id_token"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
|
||||
var oauthHTTPClient = &http.Client{Timeout: 8 * time.Second}
|
||||
|
||||
func GenerateOAuthState() (string, error) {
|
||||
buffer := make([]byte, 24)
|
||||
if _, err := rand.Read(buffer); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(buffer), nil
|
||||
}
|
||||
|
||||
func PublicAuthSources(baseAPIPath string) ([]PublicAuthSource, error) {
|
||||
sources, err := model.GetActiveAuthSources()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]PublicAuthSource, 0, len(sources))
|
||||
for _, source := range sources {
|
||||
result = append(result, PublicAuthSource{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
AuthorizeURL: fmt.Sprintf("%s/oauth/%s/authorize", strings.TrimRight(baseAPIPath, "/"), url.PathEscape(source.Name)),
|
||||
IconURL: source.IconURL,
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func BuildAuthorizeURL(ctx context.Context, source *model.AuthSource, redirectURL string, state string) (string, error) {
|
||||
source.Normalize()
|
||||
switch source.Type {
|
||||
case model.AuthSourceTypeGitHub:
|
||||
authorizeURL, err := url.Parse("https://github.com/login/oauth/authorize")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
values := authorizeURL.Query()
|
||||
values.Set("client_id", source.ClientID)
|
||||
values.Set("redirect_uri", redirectURL)
|
||||
values.Set("scope", source.Scopes)
|
||||
values.Set("state", state)
|
||||
authorizeURL.RawQuery = values.Encode()
|
||||
return authorizeURL.String(), nil
|
||||
case model.AuthSourceTypeOIDC:
|
||||
discovery, err := fetchOIDCDiscovery(ctx, source.OpenIDDiscoveryURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
authorizeURL, err := url.Parse(discovery.AuthorizationEndpoint)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
values := authorizeURL.Query()
|
||||
values.Set("client_id", source.ClientID)
|
||||
values.Set("redirect_uri", redirectURL)
|
||||
values.Set("response_type", "code")
|
||||
values.Set("scope", source.Scopes)
|
||||
values.Set("state", state)
|
||||
authorizeURL.RawQuery = values.Encode()
|
||||
return authorizeURL.String(), nil
|
||||
default:
|
||||
return "", errors.New("不支持的认证源类型")
|
||||
}
|
||||
}
|
||||
|
||||
func ExchangeOAuthProfile(ctx context.Context, source *model.AuthSource, code string, redirectURL string) (*OAuthProfile, error) {
|
||||
if strings.TrimSpace(code) == "" {
|
||||
return nil, errors.New("授权 code 不能为空")
|
||||
}
|
||||
source.Normalize()
|
||||
switch source.Type {
|
||||
case model.AuthSourceTypeGitHub:
|
||||
return exchangeGitHubProfile(ctx, source, code, redirectURL)
|
||||
case model.AuthSourceTypeOIDC:
|
||||
return exchangeOIDCProfile(ctx, source, code, redirectURL)
|
||||
default:
|
||||
return nil, errors.New("不支持的认证源类型")
|
||||
}
|
||||
}
|
||||
|
||||
func CompleteOAuthLogin(source *model.AuthSource, profile *OAuthProfile, currentUserID *int) (*OAuthCallbackResult, *PendingExternalAccount, error) {
|
||||
if source == nil || profile == nil || strings.TrimSpace(profile.ExternalID) == "" {
|
||||
return nil, nil, errors.New("第三方账号资料不完整")
|
||||
}
|
||||
|
||||
account, err := model.FindExternalAccount(source.ID, profile.ExternalID)
|
||||
if err == nil {
|
||||
user, err := model.GetUserById(account.UserID, false)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if user.Status != common.UserStatusEnabled {
|
||||
return nil, nil, errors.New("用户已被封禁")
|
||||
}
|
||||
return &OAuthCallbackResult{Status: "logged_in", User: user}, nil, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
if currentUserID != nil && *currentUserID > 0 {
|
||||
user, err := model.GetUserById(*currentUserID, false)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if user.Status != common.UserStatusEnabled {
|
||||
return nil, nil, errors.New("用户已被封禁")
|
||||
}
|
||||
if err := model.LinkExternalAccount(&model.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.Id,
|
||||
ExternalID: profile.ExternalID,
|
||||
ExternalUsername: profile.ExternalUsername,
|
||||
Email: profile.Email,
|
||||
}); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return &OAuthCallbackResult{Status: "linked", User: user}, nil, nil
|
||||
}
|
||||
|
||||
pending := &PendingExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
ExternalID: profile.ExternalID,
|
||||
ExternalUsername: profile.ExternalUsername,
|
||||
DisplayName: profile.DisplayName,
|
||||
Email: profile.Email,
|
||||
}
|
||||
return &OAuthCallbackResult{Status: "link_required"}, pending, nil
|
||||
}
|
||||
|
||||
func LinkPendingExternalAccount(pending *PendingExternalAccount, input LinkExistingRequest) (*model.User, error) {
|
||||
if pending == nil || pending.AuthSourceID == 0 || pending.ExternalID == "" {
|
||||
return nil, errors.New("待绑定第三方账号已失效,请重新登录")
|
||||
}
|
||||
user := model.User{
|
||||
Username: strings.TrimSpace(input.Username),
|
||||
Password: input.Password,
|
||||
}
|
||||
if err := user.ValidateAndFill(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if user.Status != common.UserStatusEnabled {
|
||||
return nil, errors.New("用户已被封禁")
|
||||
}
|
||||
|
||||
if existing, err := model.FindExternalAccount(pending.AuthSourceID, pending.ExternalID); err == nil {
|
||||
if existing.UserID != user.Id {
|
||||
return nil, errors.New("该第三方账号已绑定其他用户")
|
||||
}
|
||||
return &user, nil
|
||||
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := model.LinkExternalAccount(&model.ExternalAccount{
|
||||
AuthSourceID: pending.AuthSourceID,
|
||||
UserID: user.Id,
|
||||
ExternalID: pending.ExternalID,
|
||||
ExternalUsername: pending.ExternalUsername,
|
||||
Email: pending.Email,
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// CreateUserFromOAuthProfile 根据 OAuth 资料创建新用户
|
||||
func createUserFromOAuthProfile(source *model.AuthSource, profile *OAuthProfile) (*model.User, error) {
|
||||
displayName := strings.TrimSpace(profile.DisplayName)
|
||||
if displayName == "" {
|
||||
displayName = strings.TrimSpace(profile.ExternalUsername)
|
||||
}
|
||||
if displayName == "" {
|
||||
displayName = source.DisplayName + " User"
|
||||
}
|
||||
if len([]rune(displayName)) > 20 {
|
||||
displayName = string([]rune(displayName)[:20])
|
||||
}
|
||||
|
||||
prefix := source.Type
|
||||
if prefix == "" {
|
||||
prefix = "oauth"
|
||||
}
|
||||
var username string
|
||||
for index := 0; index < 20; index++ {
|
||||
username = fmt.Sprintf("%s_%d", prefix, model.GetMaxUserId()+1+index)
|
||||
if !model.IsUsernameAlreadyTaken(username) {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
user := &model.User{
|
||||
Username: username,
|
||||
DisplayName: displayName,
|
||||
Email: profile.Email,
|
||||
Role: common.RoleCommonUser,
|
||||
Status: common.UserStatusEnabled,
|
||||
}
|
||||
if err := user.Insert(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := model.LinkExternalAccount(&model.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.Id,
|
||||
ExternalID: profile.ExternalID,
|
||||
ExternalUsername: profile.ExternalUsername,
|
||||
Email: profile.Email,
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func exchangeGitHubProfile(ctx context.Context, source *model.AuthSource, code string, redirectURL string) (*OAuthProfile, error) {
|
||||
values := map[string]string{
|
||||
"client_id": source.ClientID,
|
||||
"client_secret": source.ClientSecret,
|
||||
"code": code,
|
||||
"redirect_uri": redirectURL,
|
||||
}
|
||||
body, err := json.Marshal(values)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, "https://github.com/login/oauth/access_token", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := oauthHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
slog.Error("github oauth access token request failed", "error", err)
|
||||
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试")
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("GitHub token 接口返回异常状态: %s", resp.Status)
|
||||
}
|
||||
var token oauthTokenResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&token); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if token.AccessToken == "" {
|
||||
return nil, errors.New("GitHub 未返回 access token")
|
||||
}
|
||||
|
||||
req, err = http.NewRequestWithContext(ctx, http.MethodGet, "https://api.github.com/user", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+token.AccessToken)
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
resp, err = oauthHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
slog.Error("github user info request failed", "error", err)
|
||||
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试")
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("GitHub 用户接口返回异常状态: %s", resp.Status)
|
||||
}
|
||||
var githubUser struct {
|
||||
ID int64 `json:"id"`
|
||||
Login string `json:"login"`
|
||||
Name string `json:"name"`
|
||||
Email string `json:"email"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&githubUser); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if githubUser.ID == 0 && githubUser.Login == "" {
|
||||
return nil, errors.New("GitHub 用户资料缺少唯一标识")
|
||||
}
|
||||
return &OAuthProfile{
|
||||
ExternalID: githubUser.Login,
|
||||
ExternalUsername: githubUser.Login,
|
||||
DisplayName: firstNonEmpty(githubUser.Name, githubUser.Login),
|
||||
Email: githubUser.Email,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func exchangeOIDCProfile(ctx context.Context, source *model.AuthSource, code string, redirectURL string) (*OAuthProfile, error) {
|
||||
discovery, err := fetchOIDCDiscovery(ctx, source.OpenIDDiscoveryURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
token, err := exchangeOIDCToken(ctx, discovery.TokenEndpoint, source, code, redirectURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if token.AccessToken == "" {
|
||||
return nil, errors.New("OIDC 未返回 access token")
|
||||
}
|
||||
claims, err := fetchOIDCUserInfo(ctx, discovery.UserInfoEndpoint, token.AccessToken)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(claims) == 0 && token.IDToken != "" {
|
||||
claims = decodeJWTClaims(token.IDToken)
|
||||
}
|
||||
profile := profileFromClaims(claims)
|
||||
if profile.ExternalID == "" {
|
||||
return nil, errors.New("OIDC 用户资料缺少 sub")
|
||||
}
|
||||
return profile, nil
|
||||
}
|
||||
|
||||
func fetchOIDCDiscovery(ctx context.Context, discoveryURL string) (*oidcDiscovery, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, discoveryURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := oauthHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("无法获取 OIDC discovery 配置: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("OIDC discovery 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
var discovery oidcDiscovery
|
||||
if err := json.NewDecoder(resp.Body).Decode(&discovery); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if discovery.AuthorizationEndpoint == "" || discovery.TokenEndpoint == "" {
|
||||
return nil, errors.New("OIDC discovery 缺少授权或 token 端点")
|
||||
}
|
||||
return &discovery, nil
|
||||
}
|
||||
|
||||
func exchangeOIDCToken(ctx context.Context, tokenEndpoint string, source *model.AuthSource, code string, redirectURL string) (*oauthTokenResponse, error) {
|
||||
form := url.Values{}
|
||||
form.Set("grant_type", "authorization_code")
|
||||
form.Set("client_id", source.ClientID)
|
||||
form.Set("client_secret", source.ClientSecret)
|
||||
form.Set("code", code)
|
||||
form.Set("redirect_uri", redirectURL)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, tokenEndpoint, strings.NewReader(form.Encode()))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := oauthHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("OIDC token 请求失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||
return nil, fmt.Errorf("OIDC token 接口返回异常状态: %s %s", resp.Status, strings.TrimSpace(string(raw)))
|
||||
}
|
||||
var token oauthTokenResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&token); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &token, nil
|
||||
}
|
||||
|
||||
func fetchOIDCUserInfo(ctx context.Context, endpoint string, accessToken string) (map[string]any, error) {
|
||||
if endpoint == "" {
|
||||
return map[string]any{}, nil
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := oauthHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("OIDC userinfo 请求失败: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||
return nil, fmt.Errorf("OIDC userinfo 返回异常状态: %s %s", resp.Status, strings.TrimSpace(string(raw)))
|
||||
}
|
||||
var claims map[string]any
|
||||
if err := json.NewDecoder(resp.Body).Decode(&claims); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func decodeJWTClaims(token string) map[string]any {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) < 2 {
|
||||
return map[string]any{}
|
||||
}
|
||||
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return map[string]any{}
|
||||
}
|
||||
var claims map[string]any
|
||||
if err := json.Unmarshal(payload, &claims); err != nil {
|
||||
return map[string]any{}
|
||||
}
|
||||
return claims
|
||||
}
|
||||
|
||||
func profileFromClaims(claims map[string]any) *OAuthProfile {
|
||||
stringClaim := func(keys ...string) string {
|
||||
for _, key := range keys {
|
||||
if value, ok := claims[key].(string); ok && strings.TrimSpace(value) != "" {
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
return &OAuthProfile{
|
||||
ExternalID: stringClaim("sub"),
|
||||
ExternalUsername: stringClaim("preferred_username", "nickname", "name", "email"),
|
||||
DisplayName: stringClaim("name", "preferred_username", "nickname", "email"),
|
||||
Email: stringClaim("email"),
|
||||
}
|
||||
}
|
||||
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, value := range values {
|
||||
if strings.TrimSpace(value) != "" {
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestCompleteOAuthLoginRequiresLinkWhenRegistrationDisabled(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
source := createTestAuthSource(t)
|
||||
result, pending, err := CompleteOAuthLogin(source, &OAuthProfile{
|
||||
ExternalID: "external-1",
|
||||
ExternalUsername: "external-user",
|
||||
DisplayName: "External User",
|
||||
Email: "external@example.com",
|
||||
}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("CompleteOAuthLogin failed: %v", err)
|
||||
}
|
||||
if result.Status != "link_required" || pending == nil {
|
||||
t.Fatalf("expected link_required with pending account, got %#v pending=%#v", result, pending)
|
||||
}
|
||||
|
||||
user, err := LinkPendingExternalAccount(pending, LinkExistingRequest{
|
||||
Username: "root",
|
||||
Password: "123456",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("LinkPendingExternalAccount failed: %v", err)
|
||||
}
|
||||
if user.Username != "root" {
|
||||
t.Fatalf("expected root user, got %s", user.Username)
|
||||
}
|
||||
account, err := model.FindExternalAccount(source.ID, "external-1")
|
||||
if err != nil {
|
||||
t.Fatalf("expected external account to be linked: %v", err)
|
||||
}
|
||||
if account.UserID != user.Id {
|
||||
t.Fatalf("expected external account user %d, got %d", user.Id, account.UserID)
|
||||
}
|
||||
}
|
||||
|
||||
func createTestAuthSource(t *testing.T) *model.AuthSource {
|
||||
t.Helper()
|
||||
source := &model.AuthSource{
|
||||
Name: "test-oidc",
|
||||
Type: model.AuthSourceTypeOIDC,
|
||||
DisplayName: "Test OIDC",
|
||||
ClientID: "client-id",
|
||||
ClientSecret: "client-secret",
|
||||
Scopes: "openid profile email",
|
||||
OpenIDDiscoveryURL: "https://idp.example.com/.well-known/openid-configuration",
|
||||
}
|
||||
if err := model.CreateAuthSource(source); err != nil {
|
||||
t.Fatalf("CreateAuthSource failed: %v", err)
|
||||
}
|
||||
return source
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,254 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
type DashboardOverviewView struct {
|
||||
GeneratedAt time.Time `json:"generated_at"`
|
||||
Summary DashboardSummary `json:"summary"`
|
||||
Traffic DashboardTraffic `json:"traffic"`
|
||||
Capacity DashboardCapacity `json:"capacity"`
|
||||
Distributions TrafficDistributions `json:"distributions"`
|
||||
Trends DashboardTrends `json:"trends"`
|
||||
Nodes []DashboardNodeHealth `json:"nodes"`
|
||||
}
|
||||
|
||||
type DashboardSummary struct {
|
||||
TotalNodes int `json:"total_nodes"`
|
||||
OnlineNodes int `json:"online_nodes"`
|
||||
OfflineNodes int `json:"offline_nodes"`
|
||||
PendingNodes int `json:"pending_nodes"`
|
||||
UnhealthyNodes int `json:"unhealthy_nodes"`
|
||||
}
|
||||
|
||||
type DashboardTraffic struct {
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueVisitors int64 `json:"unique_visitors"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
EstimatedQPS float64 `json:"estimated_qps"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
type DashboardCapacity struct {
|
||||
AverageCPUUsagePercent float64 `json:"average_cpu_usage_percent"`
|
||||
AverageMemoryUsagePercent float64 `json:"average_memory_usage_percent"`
|
||||
HighCPUNodes int `json:"high_cpu_nodes"`
|
||||
HighMemoryNodes int `json:"high_memory_nodes"`
|
||||
HighStorageNodes int `json:"high_storage_nodes"`
|
||||
}
|
||||
|
||||
type DashboardTrends struct {
|
||||
Traffic24h []TrafficTrendPoint `json:"traffic_24h"`
|
||||
Capacity24h []CapacityTrendPoint `json:"capacity_24h"`
|
||||
Network24h []NetworkTrendPoint `json:"network_24h"`
|
||||
DiskIO24h []DiskIOTrendPoint `json:"disk_io_24h"`
|
||||
}
|
||||
|
||||
type DashboardNodeHealth struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
Status string `json:"status"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastSeenAt any `json:"last_seen_at"`
|
||||
ActiveEventCount int `json:"active_event_count"`
|
||||
CPUUsagePercent float64 `json:"cpu_usage_percent"`
|
||||
MemoryUsagePercent float64 `json:"memory_usage_percent"`
|
||||
StorageUsagePercent float64 `json:"storage_usage_percent"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
}
|
||||
|
||||
func GetDashboardOverview() (*DashboardOverviewView, error) {
|
||||
now := time.Now()
|
||||
since := now.Add(-24 * time.Hour)
|
||||
|
||||
nodes, err := model.ListNodes()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
snapshots, err := model.ListMetricSnapshotsSince(since)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reports, err := model.ListRequestReportsSince(since)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessLogRegions, err := model.ListNodeAccessLogRegionCounts("", since, 8)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
activeEvents, err := model.ListActiveNodeHealthEvents()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
openrestySnapshots, err := model.ListNodeObservationOpenresty("", since, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
view := &DashboardOverviewView{
|
||||
GeneratedAt: now,
|
||||
Nodes: make([]DashboardNodeHealth, 0, len(nodes)),
|
||||
Distributions: buildTrafficDistributions(reports, accessLogRegions, 8),
|
||||
Trends: DashboardTrends{
|
||||
Traffic24h: buildTrafficTrendPoints(now, reports),
|
||||
Capacity24h: buildCapacityTrendPoints(now, snapshots),
|
||||
Network24h: buildNetworkTrendPoints(now, snapshots, openrestySnapshots),
|
||||
DiskIO24h: buildDiskIOTrendPoints(now, snapshots),
|
||||
},
|
||||
}
|
||||
|
||||
var cpuNodeCount int
|
||||
var memoryNodeCount int
|
||||
latestSnapshots := latestMetricSnapshotsByNode(snapshots)
|
||||
latestTrafficReports := latestTrafficReportsByNode(reports)
|
||||
activeEventsByNode := activeHealthEventsByNode(activeEvents)
|
||||
|
||||
for _, node := range nodes {
|
||||
computedStatus := computeNodeStatus(node)
|
||||
switch computedStatus {
|
||||
case NodeStatusOnline:
|
||||
view.Summary.OnlineNodes++
|
||||
case NodeStatusOffline:
|
||||
view.Summary.OfflineNodes++
|
||||
case NodeStatusPending:
|
||||
view.Summary.PendingNodes++
|
||||
}
|
||||
if node.OpenrestyStatus == OpenrestyStatusUnhealthy {
|
||||
view.Summary.UnhealthyNodes++
|
||||
}
|
||||
|
||||
latestSnapshot := latestSnapshots[node.NodeID]
|
||||
latestTraffic := latestTrafficReports[node.NodeID]
|
||||
nodeActiveEvents := activeEventsByNode[node.NodeID]
|
||||
|
||||
nodeHealth := DashboardNodeHealth{
|
||||
ID: node.ID,
|
||||
NodeID: node.NodeID,
|
||||
Name: node.Name,
|
||||
GeoName: node.GeoName,
|
||||
GeoLatitude: node.GeoLatitude,
|
||||
GeoLongitude: node.GeoLongitude,
|
||||
Status: computedStatus,
|
||||
OpenrestyStatus: node.OpenrestyStatus,
|
||||
CurrentVersion: node.CurrentVersion,
|
||||
LastSeenAt: nodeViewLastSeenAt(node),
|
||||
ActiveEventCount: len(nodeActiveEvents),
|
||||
}
|
||||
|
||||
if latestSnapshot != nil {
|
||||
nodeHealth.CPUUsagePercent = latestSnapshot.CPUUsagePercent
|
||||
nodeHealth.MemoryUsagePercent = percentage(latestSnapshot.MemoryUsedBytes, latestSnapshot.MemoryTotalBytes)
|
||||
nodeHealth.StorageUsagePercent = percentage(latestSnapshot.StorageUsedBytes, latestSnapshot.StorageTotalBytes)
|
||||
if latestSnapshot.CPUUsagePercent > 0 {
|
||||
view.Capacity.AverageCPUUsagePercent += latestSnapshot.CPUUsagePercent
|
||||
cpuNodeCount++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent > 0 {
|
||||
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
|
||||
memoryNodeCount++
|
||||
}
|
||||
if latestSnapshot.CPUUsagePercent >= 80 {
|
||||
view.Capacity.HighCPUNodes++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent >= 85 {
|
||||
view.Capacity.HighMemoryNodes++
|
||||
}
|
||||
if nodeHealth.StorageUsagePercent >= 85 {
|
||||
view.Capacity.HighStorageNodes++
|
||||
}
|
||||
}
|
||||
|
||||
if latestTraffic != nil {
|
||||
nodeHealth.RequestCount = latestTraffic.RequestCount
|
||||
nodeHealth.ErrorCount = latestTraffic.ErrorCount
|
||||
nodeHealth.UniqueVisitorCount = latestTraffic.UniqueVisitorCount
|
||||
view.Traffic.RequestCount += latestTraffic.RequestCount
|
||||
view.Traffic.UniqueVisitors += latestTraffic.UniqueVisitorCount
|
||||
view.Traffic.ErrorCount += latestTraffic.ErrorCount
|
||||
if duration := latestTraffic.WindowEndedAt.Sub(latestTraffic.WindowStartedAt).Seconds(); duration > 0 {
|
||||
view.Traffic.EstimatedQPS += float64(latestTraffic.RequestCount) / duration
|
||||
}
|
||||
view.Traffic.ReportedNodes++
|
||||
}
|
||||
|
||||
view.Nodes = append(view.Nodes, nodeHealth)
|
||||
}
|
||||
|
||||
view.Summary.TotalNodes = len(nodes)
|
||||
|
||||
if cpuNodeCount > 0 {
|
||||
view.Capacity.AverageCPUUsagePercent /= float64(cpuNodeCount)
|
||||
}
|
||||
if memoryNodeCount > 0 {
|
||||
view.Capacity.AverageMemoryUsagePercent /= float64(memoryNodeCount)
|
||||
}
|
||||
|
||||
sort.Slice(view.Nodes, func(i int, j int) bool {
|
||||
if view.Nodes[i].ActiveEventCount == view.Nodes[j].ActiveEventCount {
|
||||
return view.Nodes[i].CPUUsagePercent > view.Nodes[j].CPUUsagePercent
|
||||
}
|
||||
return view.Nodes[i].ActiveEventCount > view.Nodes[j].ActiveEventCount
|
||||
})
|
||||
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func percentage(used int64, total int64) float64 {
|
||||
if used <= 0 || total <= 0 {
|
||||
return 0
|
||||
}
|
||||
return (float64(used) / float64(total)) * 100
|
||||
}
|
||||
|
||||
func latestMetricSnapshotsByNode(snapshots []*model.NodeMetricSnapshot) map[string]*model.NodeMetricSnapshot {
|
||||
result := make(map[string]*model.NodeMetricSnapshot, len(snapshots))
|
||||
for _, snapshot := range snapshots {
|
||||
if snapshot == nil || snapshot.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
if existing, ok := result[snapshot.NodeID]; ok && !snapshot.CapturedAt.After(existing.CapturedAt) {
|
||||
continue
|
||||
}
|
||||
result[snapshot.NodeID] = snapshot
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func latestTrafficReportsByNode(reports []*model.NodeRequestReport) map[string]*model.NodeRequestReport {
|
||||
result := make(map[string]*model.NodeRequestReport, len(reports))
|
||||
for _, report := range reports {
|
||||
if report == nil || report.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
if existing, ok := result[report.NodeID]; ok && !report.WindowEndedAt.After(existing.WindowEndedAt) {
|
||||
continue
|
||||
}
|
||||
result[report.NodeID] = report
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func activeHealthEventsByNode(events []*model.NodeHealthEvent) map[string][]*model.NodeHealthEvent {
|
||||
result := make(map[string][]*model.NodeHealthEvent)
|
||||
for _, event := range events {
|
||||
if event == nil || event.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
result[event.NodeID] = append(result[event.NodeID], event)
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
const (
|
||||
DatabaseCleanupTargetAccessLogs = "node_access_logs"
|
||||
DatabaseCleanupTargetMetricSnapshots = "node_metric_snapshots"
|
||||
DatabaseCleanupTargetRequestReports = "node_request_reports"
|
||||
)
|
||||
|
||||
var databaseCleanupTargets = map[string]string{
|
||||
DatabaseCleanupTargetAccessLogs: "访问日志",
|
||||
DatabaseCleanupTargetMetricSnapshots: "性能快照",
|
||||
DatabaseCleanupTargetRequestReports: "请求聚合",
|
||||
}
|
||||
|
||||
type DatabaseCleanupInput struct {
|
||||
Target string `json:"target"`
|
||||
RetentionDays *int `json:"retention_days"`
|
||||
}
|
||||
|
||||
type DatabaseCleanupResult struct {
|
||||
Target string `json:"target"`
|
||||
TargetLabel string `json:"target_label"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays *int `json:"retention_days,omitempty"`
|
||||
Cutoff *time.Time `json:"cutoff,omitempty"`
|
||||
}
|
||||
|
||||
type DatabaseAutoCleanupSummary struct {
|
||||
RetentionDays int `json:"retention_days"`
|
||||
ExecutedAt time.Time `json:"executed_at"`
|
||||
Results []DatabaseCleanupResult `json:"results"`
|
||||
}
|
||||
|
||||
func CleanupDatabaseObservability(input DatabaseCleanupInput) (*DatabaseCleanupResult, error) {
|
||||
target := strings.TrimSpace(input.Target)
|
||||
targetLabel, ok := databaseCleanupTargets[target]
|
||||
if !ok {
|
||||
return nil, errors.New("unsupported cleanup target")
|
||||
}
|
||||
if input.RetentionDays != nil && *input.RetentionDays <= 0 {
|
||||
return nil, errors.New("retention_days 必须为大于 0 的整数")
|
||||
}
|
||||
|
||||
result := &DatabaseCleanupResult{
|
||||
Target: target,
|
||||
TargetLabel: targetLabel,
|
||||
DeleteAll: input.RetentionDays == nil,
|
||||
}
|
||||
|
||||
if input.RetentionDays == nil {
|
||||
deleted, err := deleteAllObservabilityRows(target)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.DeletedCount = deleted
|
||||
return result, nil
|
||||
}
|
||||
|
||||
retentionDays := *input.RetentionDays
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(retentionDays) * 24 * time.Hour)
|
||||
deleted, err := deleteObservabilityRowsBefore(target, cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.DeletedCount = deleted
|
||||
result.RetentionDays = &retentionDays
|
||||
result.Cutoff = &cutoff
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func RunDatabaseAutoCleanupOnce(now time.Time) (*DatabaseAutoCleanupSummary, error) {
|
||||
if !common.DatabaseAutoCleanupEnabled {
|
||||
return nil, nil
|
||||
}
|
||||
if common.DatabaseAutoCleanupRetentionDays < 1 {
|
||||
return nil, fmt.Errorf("database auto cleanup retention_days must be at least 1")
|
||||
}
|
||||
|
||||
retentionDays := common.DatabaseAutoCleanupRetentionDays
|
||||
results := make([]DatabaseCleanupResult, 0, len(databaseCleanupTargets))
|
||||
for _, target := range []string{
|
||||
DatabaseCleanupTargetAccessLogs,
|
||||
DatabaseCleanupTargetMetricSnapshots,
|
||||
DatabaseCleanupTargetRequestReports,
|
||||
} {
|
||||
result, err := CleanupDatabaseObservability(DatabaseCleanupInput{
|
||||
Target: target,
|
||||
RetentionDays: &retentionDays,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results = append(results, *result)
|
||||
}
|
||||
|
||||
return &DatabaseAutoCleanupSummary{
|
||||
RetentionDays: retentionDays,
|
||||
ExecutedAt: now.UTC(),
|
||||
Results: results,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func StartDatabaseAutoCleanupScheduler(ctx context.Context) {
|
||||
go func() {
|
||||
for {
|
||||
wait := time.Until(nextDatabaseAutoCleanupTime(time.Now()))
|
||||
timer := time.NewTimer(wait)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
timer.Stop()
|
||||
return
|
||||
case <-timer.C:
|
||||
}
|
||||
|
||||
summary, err := RunDatabaseAutoCleanupOnce(time.Now())
|
||||
if err != nil {
|
||||
slog.Error("database auto cleanup failed", "error", err)
|
||||
continue
|
||||
}
|
||||
if summary == nil {
|
||||
continue
|
||||
}
|
||||
totalDeleted := int64(0)
|
||||
for _, item := range summary.Results {
|
||||
totalDeleted += item.DeletedCount
|
||||
}
|
||||
slog.Info(
|
||||
"database auto cleanup completed",
|
||||
"retention_days",
|
||||
summary.RetentionDays,
|
||||
"deleted_count",
|
||||
totalDeleted,
|
||||
)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func nextDatabaseAutoCleanupTime(now time.Time) time.Time {
|
||||
next := time.Date(now.Year(), now.Month(), now.Day(), 3, 0, 0, 0, now.Location())
|
||||
if !next.After(now) {
|
||||
next = next.Add(24 * time.Hour)
|
||||
}
|
||||
return next
|
||||
}
|
||||
|
||||
func deleteAllObservabilityRows(target string) (int64, error) {
|
||||
switch target {
|
||||
case DatabaseCleanupTargetAccessLogs:
|
||||
return model.DeleteAllNodeAccessLogs(nil)
|
||||
case DatabaseCleanupTargetMetricSnapshots:
|
||||
return model.DeleteAllNodeMetricSnapshots(nil)
|
||||
case DatabaseCleanupTargetRequestReports:
|
||||
return model.DeleteAllNodeRequestReports(nil)
|
||||
default:
|
||||
return 0, errors.New("unsupported cleanup target")
|
||||
}
|
||||
}
|
||||
|
||||
func deleteObservabilityRowsBefore(target string, cutoff time.Time) (int64, error) {
|
||||
switch target {
|
||||
case DatabaseCleanupTargetAccessLogs:
|
||||
return model.DeleteNodeAccessLogsBefore(cutoff)
|
||||
case DatabaseCleanupTargetMetricSnapshots:
|
||||
return model.DeleteNodeMetricSnapshotsBefore(nil, cutoff)
|
||||
case DatabaseCleanupTargetRequestReports:
|
||||
return model.DeleteNodeRequestReportsBefore(nil, cutoff)
|
||||
default:
|
||||
return 0, errors.New("unsupported cleanup target")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,166 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestCleanupDatabaseObservabilityDeletesTargetedRows(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
if err := model.DB.Create(&model.NodeMetricSnapshot{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: now.Add(-10 * 24 * time.Hour),
|
||||
CPUUsagePercent: 10,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed old metric snapshot: %v", err)
|
||||
}
|
||||
if err := model.DB.Create(&model.NodeMetricSnapshot{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: now.Add(-12 * time.Hour),
|
||||
CPUUsagePercent: 20,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed recent metric snapshot: %v", err)
|
||||
}
|
||||
|
||||
retentionDays := 7
|
||||
result, err := CleanupDatabaseObservability(DatabaseCleanupInput{
|
||||
Target: DatabaseCleanupTargetMetricSnapshots,
|
||||
RetentionDays: &retentionDays,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CleanupDatabaseObservability failed: %v", err)
|
||||
}
|
||||
if result.DeleteAll {
|
||||
t.Fatal("expected retention cleanup instead of delete_all")
|
||||
}
|
||||
if result.DeletedCount != 1 {
|
||||
t.Fatalf("expected 1 deleted row, got %+v", result)
|
||||
}
|
||||
|
||||
rows, err := model.ListMetricSnapshotsSince(time.Time{})
|
||||
if err != nil {
|
||||
t.Fatalf("ListMetricSnapshotsSince failed: %v", err)
|
||||
}
|
||||
if len(rows) != 1 || rows[0].CPUUsagePercent != 20 {
|
||||
t.Fatalf("unexpected remaining metric snapshots: %+v", rows)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupDatabaseObservabilityDeletesAllRowsWhenRetentionMissing(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
if err := model.DB.Create(&model.NodeAccessLog{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-3 * time.Hour),
|
||||
RemoteAddr: "203.0.113.1",
|
||||
Host: "example.com",
|
||||
Path: "/one",
|
||||
StatusCode: 200,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed first access log: %v", err)
|
||||
}
|
||||
if err := model.DB.Create(&model.NodeAccessLog{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-2 * time.Hour),
|
||||
RemoteAddr: "203.0.113.2",
|
||||
Host: "example.com",
|
||||
Path: "/two",
|
||||
StatusCode: 502,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed second access log: %v", err)
|
||||
}
|
||||
|
||||
result, err := CleanupDatabaseObservability(DatabaseCleanupInput{
|
||||
Target: DatabaseCleanupTargetAccessLogs,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CleanupDatabaseObservability failed: %v", err)
|
||||
}
|
||||
if !result.DeleteAll || result.DeletedCount != 2 {
|
||||
t.Fatalf("unexpected delete-all result: %+v", result)
|
||||
}
|
||||
|
||||
rows, err := model.ListNodeAccessLogs(model.NodeAccessLogQuery{Page: 0, PageSize: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodeAccessLogs failed: %v", err)
|
||||
}
|
||||
if len(rows) != 0 {
|
||||
t.Fatalf("expected all access logs deleted, got %+v", rows)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunDatabaseAutoCleanupOnceDeletesAllObservabilityTargets(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
if err := model.DB.Create(&model.NodeAccessLog{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-48 * time.Hour),
|
||||
RemoteAddr: "203.0.113.10",
|
||||
Host: "example.com",
|
||||
Path: "/access",
|
||||
StatusCode: 200,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed access log: %v", err)
|
||||
}
|
||||
if err := model.DB.Create(&model.NodeMetricSnapshot{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: now.Add(-48 * time.Hour),
|
||||
CPUUsagePercent: 10,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed metric snapshot: %v", err)
|
||||
}
|
||||
if err := model.DB.Create(&model.NodeRequestReport{
|
||||
NodeID: "node-a",
|
||||
WindowStartedAt: now.Add(-49 * time.Hour),
|
||||
WindowEndedAt: now.Add(-48 * time.Hour),
|
||||
RequestCount: 15,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed request report: %v", err)
|
||||
}
|
||||
|
||||
previousEnabled := common.DatabaseAutoCleanupEnabled
|
||||
previousRetentionDays := common.DatabaseAutoCleanupRetentionDays
|
||||
common.DatabaseAutoCleanupEnabled = true
|
||||
common.DatabaseAutoCleanupRetentionDays = 1
|
||||
t.Cleanup(func() {
|
||||
common.DatabaseAutoCleanupEnabled = previousEnabled
|
||||
common.DatabaseAutoCleanupRetentionDays = previousRetentionDays
|
||||
})
|
||||
|
||||
summary, err := RunDatabaseAutoCleanupOnce(now)
|
||||
if err != nil {
|
||||
t.Fatalf("RunDatabaseAutoCleanupOnce failed: %v", err)
|
||||
}
|
||||
if summary == nil || len(summary.Results) != 3 {
|
||||
t.Fatalf("unexpected auto cleanup summary: %+v", summary)
|
||||
}
|
||||
|
||||
accessLogs, err := model.ListNodeAccessLogs(model.NodeAccessLogQuery{Page: 0, PageSize: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodeAccessLogs failed: %v", err)
|
||||
}
|
||||
if len(accessLogs) != 0 {
|
||||
t.Fatalf("expected auto cleanup to delete access logs, got %+v", accessLogs)
|
||||
}
|
||||
metricSnapshots, err := model.ListMetricSnapshotsSince(time.Time{})
|
||||
if err != nil {
|
||||
t.Fatalf("ListMetricSnapshotsSince failed: %v", err)
|
||||
}
|
||||
if len(metricSnapshots) != 0 {
|
||||
t.Fatalf("expected auto cleanup to delete metric snapshots, got %+v", metricSnapshots)
|
||||
}
|
||||
requestReports, err := model.ListRequestReportsSince(time.Time{})
|
||||
if err != nil {
|
||||
t.Fatalf("ListRequestReportsSince failed: %v", err)
|
||||
}
|
||||
if len(requestReports) != 0 {
|
||||
t.Fatalf("expected auto cleanup to delete request reports, got %+v", requestReports)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package service
|
||||
|
||||
const (
|
||||
FlaredWSConnectedLastSeenValue = "__OPENFLARE_FLARED_WS_CONNECTED__"
|
||||
|
||||
FlaredWSMessageTypeActiveConfig = "active_config"
|
||||
FlaredWSMessageTypeForceSync = "force_sync"
|
||||
FlaredWSMessageTypePong = "pong"
|
||||
)
|
||||
|
||||
var DefaultFlaredWSHub = NewWSHub("flared")
|
||||
|
||||
func RegisterFlaredWSClient(nodeID string) *WSClient {
|
||||
return DefaultFlaredWSHub.Register(nodeID)
|
||||
}
|
||||
|
||||
func UnregisterFlaredWSClient(client *WSClient) {
|
||||
DefaultFlaredWSHub.Unregister(client)
|
||||
}
|
||||
|
||||
func DisconnectFlaredWSClient(nodeID string) {
|
||||
DefaultFlaredWSHub.Disconnect(nodeID)
|
||||
}
|
||||
|
||||
func IsFlaredWSConnected(nodeID string) bool {
|
||||
return DefaultFlaredWSHub.IsConnected(nodeID)
|
||||
}
|
||||
|
||||
func SendFlaredWSPong(nodeID string) bool {
|
||||
return DefaultFlaredWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: FlaredWSMessageTypePong,
|
||||
})
|
||||
}
|
||||
|
||||
func SendFlaredWSActiveConfig(nodeID string, activeConfig *ActiveConfigMeta) bool {
|
||||
if activeConfig == nil {
|
||||
return false
|
||||
}
|
||||
return DefaultFlaredWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: FlaredWSMessageTypeActiveConfig,
|
||||
Payload: activeConfig,
|
||||
})
|
||||
}
|
||||
|
||||
func BroadcastFlaredWSActiveConfig(activeConfig *ActiveConfigMeta) WSBroadcastResult {
|
||||
if activeConfig == nil {
|
||||
return WSBroadcastResult{}
|
||||
}
|
||||
result := DefaultFlaredWSHub.Broadcast(WSMessage{
|
||||
Type: FlaredWSMessageTypeActiveConfig,
|
||||
Payload: activeConfig,
|
||||
})
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
|
||||
)
|
||||
|
||||
type GeoIPLookupView struct {
|
||||
Provider string `json:"provider"`
|
||||
IP string `json:"ip"`
|
||||
ISOCode string `json:"iso_code"`
|
||||
Name string `json:"name"`
|
||||
Latitude *float64 `json:"latitude,omitempty"`
|
||||
Longitude *float64 `json:"longitude,omitempty"`
|
||||
}
|
||||
|
||||
func LookupGeoIP(provider string, rawIP string) (*GeoIPLookupView, error) {
|
||||
trimmedProvider := strings.TrimSpace(provider)
|
||||
if !geoip.IsValidProvider(trimmedProvider) {
|
||||
return nil, errors.New("归属方式仅支持 disabled、mmdb、ip-api、geojs、ipinfo")
|
||||
}
|
||||
|
||||
trimmedIP := strings.TrimSpace(rawIP)
|
||||
if trimmedIP == "" {
|
||||
return nil, errors.New("IP 不能为空")
|
||||
}
|
||||
parsedIP := net.ParseIP(trimmedIP)
|
||||
if parsedIP == nil {
|
||||
return nil, errors.New("IP 格式无效")
|
||||
}
|
||||
|
||||
info, err := geoip.LookupGeoInfoWithProvider(trimmedProvider, parsedIP)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if info == nil {
|
||||
return nil, errors.New("未获取到 IP 归属结果")
|
||||
}
|
||||
|
||||
return &GeoIPLookupView{
|
||||
Provider: trimmedProvider,
|
||||
IP: parsedIP.String(),
|
||||
ISOCode: info.ISOCode,
|
||||
Name: info.Name,
|
||||
Latitude: info.Latitude,
|
||||
Longitude: info.Longitude,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
|
||||
)
|
||||
|
||||
type fakeLookupProvider struct{}
|
||||
|
||||
func (f *fakeLookupProvider) Name() string {
|
||||
return "fake-lookup"
|
||||
}
|
||||
|
||||
func (f *fakeLookupProvider) GetGeoInfo(ip net.IP) (*geoip.GeoInfo, error) {
|
||||
return &geoip.GeoInfo{
|
||||
ISOCode: "US",
|
||||
Name: "United States",
|
||||
Latitude: geoipFloat(37.7749),
|
||||
Longitude: geoipFloat(-122.4194),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeLookupProvider) UpdateDatabase() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeLookupProvider) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestLookupGeoIP(t *testing.T) {
|
||||
previousFactory := geoip.ProviderFactoryForTest()
|
||||
geoip.SetProviderFactoryForTest(func(provider string) (geoip.GeoIPService, error) {
|
||||
return &fakeLookupProvider{}, nil
|
||||
})
|
||||
defer geoip.SetProviderFactoryForTest(previousFactory)
|
||||
|
||||
view, err := LookupGeoIP("ipinfo", "8.8.8.8")
|
||||
if err != nil {
|
||||
t.Fatalf("LookupGeoIP failed: %v", err)
|
||||
}
|
||||
if view.Provider != "ipinfo" {
|
||||
t.Fatalf("expected provider ipinfo, got %s", view.Provider)
|
||||
}
|
||||
if view.IP != "8.8.8.8" {
|
||||
t.Fatalf("expected IP 8.8.8.8, got %s", view.IP)
|
||||
}
|
||||
if view.ISOCode != "US" || view.Name != "United States" {
|
||||
t.Fatalf("unexpected lookup view: %+v", view)
|
||||
}
|
||||
if view.Latitude == nil || view.Longitude == nil {
|
||||
t.Fatalf("expected coordinates, got %+v", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupGeoIPRejectsInvalidInput(t *testing.T) {
|
||||
if _, err := LookupGeoIP("invalid", "8.8.8.8"); err == nil {
|
||||
t.Fatal("expected invalid provider to fail")
|
||||
}
|
||||
if _, err := LookupGeoIP("ipinfo", "not-an-ip"); err == nil {
|
||||
t.Fatal("expected invalid IP to fail")
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,101 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/acme"
|
||||
)
|
||||
|
||||
func ObtainSSL(cert *model.TLSCertificate) error {
|
||||
cert.ApplyStatus = "applying"
|
||||
model.DB.Save(cert)
|
||||
|
||||
acmeAccount, err := model.GetAcmeAccountByID(cert.AcmeAccountID)
|
||||
if err != nil {
|
||||
// Fallback to default ACME account if the specified one is not found (e.g. ID 0 during testing)
|
||||
acmeAccount, err = model.GetDefaultAcmeAccount()
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to get ACME account: %v", err))
|
||||
return err
|
||||
}
|
||||
// Self-heal the certificate
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
model.DB.Save(cert)
|
||||
}
|
||||
|
||||
dnsAccount, err := model.GetDnsAccountByID(cert.DnsAccountID)
|
||||
if err != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to get DNS account: %v", err))
|
||||
return err
|
||||
}
|
||||
|
||||
domains := []string{cert.PrimaryDomain}
|
||||
if cert.OtherDomains != "" {
|
||||
for _, d := range strings.Split(cert.OtherDomains, "\n") {
|
||||
d = strings.TrimSpace(d)
|
||||
if d != "" {
|
||||
domains = append(domains, d)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
newAccountURL, newPrivateKeyPEM, result, err := acme.ObtainSSL(
|
||||
acmeAccount.Email,
|
||||
acmeAccount.PrivateKey,
|
||||
acmeAccount.URL,
|
||||
dnsAccount.Type,
|
||||
dnsAccount.Authorization,
|
||||
cert.DNS1,
|
||||
cert.DNS2,
|
||||
cert.DisableCNAME,
|
||||
cert.SkipDNS,
|
||||
cert.KeyAlgorithm,
|
||||
domains,
|
||||
)
|
||||
|
||||
// If new key or URL was generated, save them to the DB
|
||||
if (newPrivateKeyPEM != "" && acmeAccount.PrivateKey != newPrivateKeyPEM) || (newAccountURL != "" && acmeAccount.URL != newAccountURL) {
|
||||
if newPrivateKeyPEM != "" {
|
||||
acmeAccount.PrivateKey = newPrivateKeyPEM
|
||||
}
|
||||
if newAccountURL != "" {
|
||||
acmeAccount.URL = newAccountURL
|
||||
}
|
||||
if acmeAccount.ID == 0 {
|
||||
if dbErr := model.DB.Create(acmeAccount).Error; dbErr != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to create ACME account: %v", dbErr))
|
||||
return dbErr
|
||||
}
|
||||
} else {
|
||||
if dbErr := model.DB.Save(acmeAccount).Error; dbErr != nil {
|
||||
updateCertError(cert, fmt.Sprintf("Failed to save ACME account: %v", dbErr))
|
||||
return dbErr
|
||||
}
|
||||
}
|
||||
// Self-heal the cert
|
||||
cert.AcmeAccountID = acmeAccount.ID
|
||||
model.DB.Save(cert)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
updateCertError(cert, err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
cert.CertPEM = result.CertPEM
|
||||
cert.KeyPEM = result.KeyPEM
|
||||
cert.NotBefore = result.NotBefore
|
||||
cert.NotAfter = result.NotAfter
|
||||
cert.ApplyStatus = "ready"
|
||||
cert.ApplyMessage = ""
|
||||
|
||||
return model.DB.Save(cert).Error
|
||||
}
|
||||
|
||||
func updateCertError(cert *model.TLSCertificate, message string) {
|
||||
cert.ApplyStatus = "error"
|
||||
cert.ApplyMessage = message
|
||||
model.DB.Save(cert)
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
const (
|
||||
ManagedDomainMatchTypeExact = "exact"
|
||||
ManagedDomainMatchTypeWildcard = "wildcard"
|
||||
)
|
||||
|
||||
type ManagedDomainInput struct {
|
||||
Domain string `json:"domain"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
type ManagedDomainMatchCandidate struct {
|
||||
ManagedDomainID uint `json:"managed_domain_id"`
|
||||
Domain string `json:"domain"`
|
||||
MatchType string `json:"match_type"`
|
||||
CertificateID uint `json:"certificate_id"`
|
||||
CertificateName string `json:"certificate_name"`
|
||||
}
|
||||
|
||||
type ManagedDomainMatchResult struct {
|
||||
Domain string `json:"domain"`
|
||||
Matched bool `json:"matched"`
|
||||
Candidate *ManagedDomainMatchCandidate `json:"candidate,omitempty"`
|
||||
Candidates []ManagedDomainMatchCandidate `json:"candidates"`
|
||||
}
|
||||
|
||||
func ListManagedDomains() ([]*model.ManagedDomain, error) {
|
||||
return model.ListManagedDomains()
|
||||
}
|
||||
|
||||
func CreateManagedDomain(input ManagedDomainInput) (*model.ManagedDomain, error) {
|
||||
domain, err := buildManagedDomain(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = domain.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return domain, nil
|
||||
}
|
||||
|
||||
func UpdateManagedDomain(id uint, input ManagedDomainInput) (*model.ManagedDomain, error) {
|
||||
domain, err := model.GetManagedDomainByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domain, err = buildManagedDomain(domain, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = domain.Update(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return domain, nil
|
||||
}
|
||||
|
||||
func DeleteManagedDomain(id uint) error {
|
||||
domain, err := model.GetManagedDomainByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return domain.Delete()
|
||||
}
|
||||
|
||||
func MatchManagedDomainCertificate(rawDomain string) (*ManagedDomainMatchResult, error) {
|
||||
domain := normalizeManagedDomain(rawDomain)
|
||||
if err := validateManagedDomainPattern(domain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
managedDomains, err := model.ListEnabledManagedDomainsWithCertificate()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
candidates := make([]ManagedDomainMatchCandidate, 0)
|
||||
for _, item := range managedDomains {
|
||||
if item.CertID == nil || *item.CertID == 0 {
|
||||
continue
|
||||
}
|
||||
matchType := detectManagedDomainMatchType(item.Domain, domain)
|
||||
if matchType == "" {
|
||||
continue
|
||||
}
|
||||
certificate, err := model.GetTLSCertificateByID(*item.CertID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("托管域名 %s 关联证书不存在", item.Domain)
|
||||
}
|
||||
candidates = append(candidates, ManagedDomainMatchCandidate{
|
||||
ManagedDomainID: item.ID,
|
||||
Domain: item.Domain,
|
||||
MatchType: matchType,
|
||||
CertificateID: certificate.ID,
|
||||
CertificateName: certificate.Name,
|
||||
})
|
||||
}
|
||||
sortManagedDomainCandidates(candidates)
|
||||
result := &ManagedDomainMatchResult{
|
||||
Domain: domain,
|
||||
Matched: len(candidates) > 0,
|
||||
Candidates: candidates,
|
||||
}
|
||||
if len(candidates) > 0 {
|
||||
candidate := candidates[0]
|
||||
result.Candidate = &candidate
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func buildManagedDomain(existing *model.ManagedDomain, input ManagedDomainInput) (*model.ManagedDomain, error) {
|
||||
domain := normalizeManagedDomain(input.Domain)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
if err := validateManagedDomainPattern(domain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if input.CertID != nil && *input.CertID != 0 {
|
||||
if _, err := model.GetTLSCertificateByID(*input.CertID); err != nil {
|
||||
return nil, errors.New("所选证书不存在")
|
||||
}
|
||||
} else {
|
||||
input.CertID = nil
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.ManagedDomain{}
|
||||
}
|
||||
existing.Domain = domain
|
||||
existing.CertID = input.CertID
|
||||
existing.Enabled = input.Enabled
|
||||
existing.Remark = remark
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
func normalizeManagedDomain(domain string) string {
|
||||
return strings.ToLower(strings.TrimSpace(domain))
|
||||
}
|
||||
|
||||
func validateManagedDomainPattern(domain string) error {
|
||||
if domain == "" {
|
||||
return errors.New("域名不能为空")
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
if strings.Contains(domain, "*") {
|
||||
if !strings.HasPrefix(domain, "*.") || strings.Count(domain, "*") != 1 {
|
||||
return errors.New("通配符域名仅支持 *.example.com 格式")
|
||||
}
|
||||
return validateHostname(strings.TrimPrefix(domain, "*."))
|
||||
}
|
||||
return validateHostname(domain)
|
||||
}
|
||||
|
||||
func validateHostname(domain string) error {
|
||||
if domain == "" {
|
||||
return errors.New("域名不能为空")
|
||||
}
|
||||
if len(domain) > 253 {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
labels := strings.Split(domain, ".")
|
||||
if len(labels) < 2 {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
for _, label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
if label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
for _, r := range label {
|
||||
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
|
||||
continue
|
||||
}
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func detectManagedDomainMatchType(pattern string, domain string) string {
|
||||
if pattern == domain {
|
||||
return ManagedDomainMatchTypeExact
|
||||
}
|
||||
if !strings.HasPrefix(pattern, "*.") {
|
||||
return ""
|
||||
}
|
||||
suffix := strings.TrimPrefix(pattern, "*.")
|
||||
if !strings.HasSuffix(domain, "."+suffix) {
|
||||
return ""
|
||||
}
|
||||
prefix := strings.TrimSuffix(domain, "."+suffix)
|
||||
if prefix == "" || strings.Contains(prefix, ".") {
|
||||
return ""
|
||||
}
|
||||
return ManagedDomainMatchTypeWildcard
|
||||
}
|
||||
|
||||
func sortManagedDomainCandidates(candidates []ManagedDomainMatchCandidate) {
|
||||
sort.Slice(candidates, func(i int, j int) bool {
|
||||
left := candidates[i]
|
||||
right := candidates[j]
|
||||
if left.MatchType != right.MatchType {
|
||||
return left.MatchType == ManagedDomainMatchTypeExact
|
||||
}
|
||||
if len(left.Domain) != len(right.Domain) {
|
||||
return len(left.Domain) > len(right.Domain)
|
||||
}
|
||||
return left.ManagedDomainID < right.ManagedDomainID
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package service
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMatchManagedDomainCertificatePrefersExactMatch(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
wildcardCertPEM, wildcardKeyPEM := generateCertificatePair(t, []string{"*.example.com"})
|
||||
wildcardCert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "wildcard-cert",
|
||||
CertPEM: wildcardCertPEM,
|
||||
KeyPEM: wildcardKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create wildcard certificate: %v", err)
|
||||
}
|
||||
exactCertPEM, exactKeyPEM := generateCertificatePair(t, []string{"api.example.com"})
|
||||
exactCert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "exact-cert",
|
||||
CertPEM: exactCertPEM,
|
||||
KeyPEM: exactKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create exact certificate: %v", err)
|
||||
}
|
||||
if _, err = CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "*.example.com",
|
||||
CertID: &wildcardCert.ID,
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to create wildcard managed domain: %v", err)
|
||||
}
|
||||
if _, err = CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "api.example.com",
|
||||
CertID: &exactCert.ID,
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to create exact managed domain: %v", err)
|
||||
}
|
||||
|
||||
result, err := MatchManagedDomainCertificate("api.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("MatchManagedDomainCertificate failed: %v", err)
|
||||
}
|
||||
if !result.Matched || result.Candidate == nil {
|
||||
t.Fatal("expected exact domain to be matched")
|
||||
}
|
||||
if result.Candidate.MatchType != ManagedDomainMatchTypeExact {
|
||||
t.Fatalf("expected exact match first, got %s", result.Candidate.MatchType)
|
||||
}
|
||||
if result.Candidate.CertificateID != exactCert.ID {
|
||||
t.Fatalf("expected exact certificate %d, got %d", exactCert.ID, result.Candidate.CertificateID)
|
||||
}
|
||||
if len(result.Candidates) != 2 {
|
||||
t.Fatalf("expected 2 match candidates, got %d", len(result.Candidates))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatchManagedDomainCertificateSupportsWildcard(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
certPEM, keyPEM := generateCertificatePair(t, []string{"*.example.com"})
|
||||
certificate, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "wildcard-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create certificate: %v", err)
|
||||
}
|
||||
if _, err = CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "*.example.com",
|
||||
CertID: &certificate.ID,
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to create managed domain: %v", err)
|
||||
}
|
||||
|
||||
result, err := MatchManagedDomainCertificate("edge.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("MatchManagedDomainCertificate failed: %v", err)
|
||||
}
|
||||
if !result.Matched || result.Candidate == nil {
|
||||
t.Fatal("expected wildcard domain to be matched")
|
||||
}
|
||||
if result.Candidate.MatchType != ManagedDomainMatchTypeWildcard {
|
||||
t.Fatalf("expected wildcard match, got %s", result.Candidate.MatchType)
|
||||
}
|
||||
|
||||
deepResult, err := MatchManagedDomainCertificate("deep.edge.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("MatchManagedDomainCertificate failed: %v", err)
|
||||
}
|
||||
if deepResult.Matched {
|
||||
t.Fatal("expected single-level wildcard not to match deep subdomain")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateManagedDomainRejectsInvalidWildcard(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "*.*.example.com",
|
||||
Enabled: true,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid wildcard domain to fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,661 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip/iputil"
|
||||
)
|
||||
|
||||
type NodeInput struct {
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
IPManualOverride *bool `json:"ip_manual_override"`
|
||||
AutoUpdateEnabled bool `json:"auto_update_enabled"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
GeoManualOverride bool `json:"geo_manual_override"`
|
||||
// TunnelRelay fields
|
||||
NodeType string `json:"node_type"`
|
||||
RelayBindPort int `json:"relay_bind_port"`
|
||||
RelayVhostHTTPPort int `json:"relay_vhost_http_port"`
|
||||
RelayAgentAccessAddr string `json:"relay_agent_access_addr"`
|
||||
RelayClientAccessAddr string `json:"relay_client_access_addr"`
|
||||
RelayClientProxyURL string `json:"relay_client_proxy_url"`
|
||||
RelayWebServerEnabled bool `json:"relay_web_server_enabled"`
|
||||
}
|
||||
|
||||
type NodeAgentUpdateInput struct {
|
||||
Channel string `json:"channel"`
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
|
||||
type NodeAgentReleaseInfo struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Body string `json:"body"`
|
||||
HTMLURL string `json:"html_url"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
HasUpdate bool `json:"has_update"`
|
||||
Channel string `json:"channel"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
UpdateRequested bool `json:"update_requested"`
|
||||
RequestedChannel string `json:"requested_channel"`
|
||||
RequestedTag string `json:"requested_tag"`
|
||||
}
|
||||
|
||||
type NodeBootstrapView struct {
|
||||
DiscoveryToken string `json:"discovery_token"`
|
||||
}
|
||||
|
||||
type AgentRegistrationResponse struct {
|
||||
NodeID string `json:"node_id"`
|
||||
AccessToken string `json:"access_token"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
func CreateNode(input NodeInput) (*NodeView, error) {
|
||||
name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input)
|
||||
if name == "" {
|
||||
return nil, errors.New("节点名不能为空")
|
||||
}
|
||||
ipManualOverride := resolveNodeIPManualOverride(input, nil, ip)
|
||||
node := &model.Node{
|
||||
Name: name,
|
||||
IP: ip,
|
||||
IPManualOverride: ipManualOverride,
|
||||
GeoName: geoName,
|
||||
GeoLatitude: geoLatitude,
|
||||
GeoLongitude: geoLongitude,
|
||||
GeoManualOverride: geoManualOverride,
|
||||
Version: "",
|
||||
ExtVersion: "",
|
||||
Status: NodeStatusPending,
|
||||
AutoUpdateEnabled: input.AutoUpdateEnabled,
|
||||
NodeType: normalizeNodeType(input.NodeType),
|
||||
}
|
||||
node.NodeID, err = newServerNodeID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.AccessToken, err = newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, 7000)
|
||||
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, 8080)
|
||||
node.RelayAuthToken, err = newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr)
|
||||
node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr)
|
||||
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
|
||||
node.RelayWebServerEnabled = input.RelayWebServerEnabled
|
||||
}
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("节点标识生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
slog.Info("node created", "name", node.Name, "node_id", node.NodeID)
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
func UpdateNode(id uint, input NodeInput) (*NodeView, error) {
|
||||
name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input)
|
||||
if name == "" {
|
||||
return nil, errors.New("节点名不能为空")
|
||||
}
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ipManualOverride := resolveNodeIPManualOverride(input, node, ip)
|
||||
node.Name = name
|
||||
node.IP = ip
|
||||
node.IPManualOverride = ipManualOverride
|
||||
node.GeoName = geoName
|
||||
node.GeoLatitude = geoLatitude
|
||||
node.GeoLongitude = geoLongitude
|
||||
node.GeoManualOverride = geoManualOverride
|
||||
node.AutoUpdateEnabled = input.AutoUpdateEnabled
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr)
|
||||
node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr)
|
||||
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
|
||||
node.RelayWebServerEnabled = input.RelayWebServerEnabled
|
||||
if input.RelayBindPort > 0 {
|
||||
node.RelayBindPort = input.RelayBindPort
|
||||
}
|
||||
if input.RelayVhostHTTPPort > 0 {
|
||||
node.RelayVhostHTTPPort = input.RelayVhostHTTPPort
|
||||
}
|
||||
}
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, strings.TrimSpace(node.IP))
|
||||
}
|
||||
if err = node.Update(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
slog.Info("node updated", "name", node.Name, "node_id", node.NodeID)
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
func DeleteNode(id uint) error {
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
slog.Info("node deleted", "name", node.Name, "node_id", node.NodeID)
|
||||
if err := node.Delete(); err != nil {
|
||||
return err
|
||||
}
|
||||
invalidateAccessTokenCache(node.AccessToken)
|
||||
DisconnectAgentWSClient(node.NodeID)
|
||||
DisconnectFlaredWSClient(node.NodeID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func GetNodeAgentRelease(ctx context.Context, id uint, channel string) (*NodeAgentReleaseInfo, error) {
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
release, err := fetchLatestGitHubRelease(ctx, common.AgentUpdateRepo, normalizeReleaseChannel(channel))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildNodeAgentReleaseView(node, release, normalizeReleaseChannel(channel)), nil
|
||||
}
|
||||
|
||||
func RequestNodeAgentUpdate(id uint, input NodeAgentUpdateInput) (*NodeView, error) {
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
channel := normalizeReleaseChannel(input.Channel)
|
||||
tagName := strings.TrimSpace(input.TagName)
|
||||
if tagName != "" {
|
||||
release, releaseErr := fetchGitHubReleaseByTag(context.Background(), common.AgentUpdateRepo, tagName)
|
||||
if releaseErr != nil {
|
||||
return nil, releaseErr
|
||||
}
|
||||
if channel == ReleaseChannelPreview && !release.Prerelease {
|
||||
return nil, errors.New("指定版本不是 preview 发布")
|
||||
}
|
||||
if channel == ReleaseChannelStable && release.Prerelease {
|
||||
return nil, errors.New("正式版更新不能选择 preview 发布")
|
||||
}
|
||||
}
|
||||
node.UpdateRequested = true
|
||||
node.UpdateChannel = channel.String()
|
||||
node.UpdateTag = tagName
|
||||
if err = model.DB.Model(node).Select("update_requested", "update_channel", "update_tag").Updates(node).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
if SendAgentWSSettings(node.NodeID, buildAgentSettings(node, true, channel.String(), tagName, node.RestartOpenrestyRequested)) {
|
||||
slog.Debug("agent manual update pushed via ws", "node_id", node.NodeID, "channel", channel.String(), "tag", tagName)
|
||||
} else {
|
||||
slog.Debug("agent manual update waiting for next heartbeat", "node_id", node.NodeID, "channel", channel.String(), "tag", tagName)
|
||||
}
|
||||
slog.Info("agent manual update requested", "node_id", node.NodeID, "name", node.Name, "channel", channel.String(), "tag", tagName)
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
func RequestNodeOpenrestyRestart(id uint) (*NodeView, error) {
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.RestartOpenrestyRequested = true
|
||||
if err = model.DB.Model(node).Select("restart_openresty_requested").Updates(node).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
slog.Info("openresty restart requested", "node_id", node.NodeID, "name", node.Name)
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
func RequestNodeForceSync(id uint) (*NodeView, error) {
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
activeConfig, err := GetActiveConfigMetaForAgent()
|
||||
if err != nil {
|
||||
return nil, errors.New("无法获取当前激活的配置版本:" + err.Error())
|
||||
}
|
||||
if !SendAgentWSForceSyncConfig(node.NodeID, activeConfig) {
|
||||
return nil, errors.New("节点不在线或通过 WebSocket 发送同步指令失败")
|
||||
}
|
||||
slog.Info("force sync requested via ws", "node_id", node.NodeID, "name", node.Name)
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
func AuthenticateAccessToken(token string) (*model.Node, error) {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return nil, errors.New("缺少 Agent Token")
|
||||
}
|
||||
return authenticateAccessTokenWithCache(token)
|
||||
}
|
||||
|
||||
func ValidateDiscoveryToken(token string) error {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return errors.New("缺少 Discovery Token")
|
||||
}
|
||||
discoveryToken, err := EnsureGlobalDiscoveryToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if token != discoveryToken {
|
||||
return errors.New("Discovery Token 无效")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func EnsureGlobalDiscoveryToken() (string, error) {
|
||||
common.OptionMapRWMutex.RLock()
|
||||
needsInit := common.OptionMap == nil
|
||||
common.OptionMapRWMutex.RUnlock()
|
||||
if needsInit {
|
||||
model.InitOptionMap()
|
||||
}
|
||||
common.OptionMapRWMutex.RLock()
|
||||
token := strings.TrimSpace(common.AgentDiscoveryToken)
|
||||
common.OptionMapRWMutex.RUnlock()
|
||||
if token != "" {
|
||||
return token, nil
|
||||
}
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err = model.UpdateOption("AgentDiscoveryToken", token); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func GetNodeBootstrapView() (*NodeBootstrapView, error) {
|
||||
token, err := EnsureGlobalDiscoveryToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &NodeBootstrapView{DiscoveryToken: token}, nil
|
||||
}
|
||||
|
||||
func RotateGlobalDiscoveryToken() (*NodeBootstrapView, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.UpdateOption("AgentDiscoveryToken", token); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &NodeBootstrapView{DiscoveryToken: token}, nil
|
||||
}
|
||||
|
||||
func buildNodeView(node *model.Node) *NodeView {
|
||||
status := computeNodeStatus(node)
|
||||
view := &NodeView{
|
||||
ID: node.ID,
|
||||
NodeID: node.NodeID,
|
||||
Name: node.Name,
|
||||
IP: node.IP,
|
||||
IPManualOverride: node.IPManualOverride,
|
||||
GeoName: strings.TrimSpace(node.GeoName),
|
||||
GeoLatitude: node.GeoLatitude,
|
||||
GeoLongitude: node.GeoLongitude,
|
||||
GeoManualOverride: node.GeoManualOverride,
|
||||
AccessToken: node.AccessToken,
|
||||
UpdateChannel: strings.TrimSpace(node.UpdateChannel),
|
||||
UpdateTag: strings.TrimSpace(node.UpdateTag),
|
||||
RestartOpenrestyRequested: node.RestartOpenrestyRequested,
|
||||
Version: node.Version,
|
||||
ExtVersion: node.ExtVersion,
|
||||
OpenrestyStatus: normalizeOpenrestyStatus(node.OpenrestyStatus),
|
||||
OpenrestyMessage: strings.TrimSpace(node.OpenrestyMessage),
|
||||
Status: status,
|
||||
CurrentVersion: node.CurrentVersion,
|
||||
LastSeenAt: nodeViewLastSeenAt(node),
|
||||
LastError: node.LastError,
|
||||
CreatedAt: node.CreatedAt,
|
||||
UpdatedAt: node.UpdatedAt,
|
||||
AutoUpdateEnabled: node.AutoUpdateEnabled,
|
||||
UpdateRequested: node.UpdateRequested,
|
||||
}
|
||||
if view.UpdateChannel == "" {
|
||||
view.UpdateChannel = ReleaseChannelStable.String()
|
||||
}
|
||||
view.NodeType = node.NodeType
|
||||
if view.NodeType == "" {
|
||||
view.NodeType = "edge_node"
|
||||
}
|
||||
view.RelayBindPort = node.RelayBindPort
|
||||
view.RelayVhostHTTPPort = node.RelayVhostHTTPPort
|
||||
view.RelayAgentAccessAddr = node.RelayAgentAccessAddr
|
||||
view.RelayClientAccessAddr = node.RelayClientAccessAddr
|
||||
view.RelayClientProxyURL = node.RelayClientProxyURL
|
||||
view.RelayStatus = node.RelayStatus
|
||||
view.RelayWebServerEnabled = node.RelayWebServerEnabled
|
||||
view.Version = node.Version
|
||||
view.ExtVersion = node.ExtVersion
|
||||
return view
|
||||
}
|
||||
|
||||
func nodeViewLastSeenAt(node *model.Node) any {
|
||||
if node == nil {
|
||||
return time.Time{}
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" && IsRelayWSConnected(node.NodeID) {
|
||||
return RelayWSConnectedLastSeenValue
|
||||
}
|
||||
if node.NodeType == "tunnel_client" && IsFlaredWSConnected(node.NodeID) {
|
||||
return FlaredWSConnectedLastSeenValue
|
||||
}
|
||||
if IsAgentWSConnected(node.NodeID) {
|
||||
return AgentWSConnectedLastSeenValue
|
||||
}
|
||||
return node.LastSeenAt
|
||||
}
|
||||
|
||||
func normalizeNodeInput(input NodeInput) (string, string, string, *float64, *float64, bool, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
ip := strings.TrimSpace(input.IP)
|
||||
geoName := strings.TrimSpace(input.GeoName)
|
||||
manualOverride := input.GeoManualOverride || geoName != "" || input.GeoLatitude != nil || input.GeoLongitude != nil
|
||||
if len(ip) > 64 {
|
||||
return "", "", "", nil, nil, false, errors.New("节点 IP 不能超过 64 个字符")
|
||||
}
|
||||
if ip != "" && net.ParseIP(ip) == nil {
|
||||
return "", "", "", nil, nil, false, errors.New("节点 IP 格式无效")
|
||||
}
|
||||
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
|
||||
return "", "", "", nil, nil, false, errors.New("锁定节点 IP 时必须填写节点 IP")
|
||||
}
|
||||
if len(geoName) > 128 {
|
||||
return "", "", "", nil, nil, false, errors.New("节点位置名不能超过 128 个字符")
|
||||
}
|
||||
|
||||
geoLatitude := cloneCoordinate(input.GeoLatitude)
|
||||
geoLongitude := cloneCoordinate(input.GeoLongitude)
|
||||
if (geoLatitude == nil) != (geoLongitude == nil) {
|
||||
return "", "", "", nil, nil, false, errors.New("地图坐标必须同时填写纬度和经度")
|
||||
}
|
||||
if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) {
|
||||
return "", "", "", nil, nil, false, errors.New("纬度必须在 -90 到 90 之间")
|
||||
}
|
||||
if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) {
|
||||
return "", "", "", nil, nil, false, errors.New("经度必须在 -180 到 180 之间")
|
||||
}
|
||||
|
||||
if !manualOverride {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
if geoLatitude == nil && geoLongitude == nil && geoName == "" {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
|
||||
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
|
||||
}
|
||||
|
||||
func resolveNodeIPManualOverride(input NodeInput, existing *model.Node, normalizedIP string) bool {
|
||||
if input.IPManualOverride != nil {
|
||||
return *input.IPManualOverride
|
||||
}
|
||||
if existing == nil {
|
||||
return strings.TrimSpace(normalizedIP) != ""
|
||||
}
|
||||
if existing.IPManualOverride {
|
||||
return true
|
||||
}
|
||||
return strings.TrimSpace(normalizedIP) != "" && strings.TrimSpace(normalizedIP) != strings.TrimSpace(existing.IP)
|
||||
}
|
||||
|
||||
func cloneCoordinate(value *float64) *float64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func ResolveReportedNodeIP(reportedIP string, remoteAddr string) string {
|
||||
reported := iputil.NormalizeIP(reportedIP)
|
||||
remote := iputil.NormalizeRemoteAddr(remoteAddr)
|
||||
if reported == "" {
|
||||
return remote
|
||||
}
|
||||
if !shouldPreferRemoteNodeIP(reported) {
|
||||
return reported
|
||||
}
|
||||
if isPublicNodeIP(remote) {
|
||||
return remote
|
||||
}
|
||||
return reported
|
||||
}
|
||||
|
||||
func shouldPreferRemoteNodeIP(ip string) bool {
|
||||
return !isPublicNodeIP(ip)
|
||||
}
|
||||
|
||||
func isPublicNodeIP(raw string) bool {
|
||||
return iputil.IsPublicString(raw)
|
||||
}
|
||||
|
||||
func buildNodeAgentReleaseView(node *model.Node, release *githubReleaseResponse, channel ReleaseChannel) *NodeAgentReleaseInfo {
|
||||
currentVersion := strings.TrimSpace(node.Version)
|
||||
view := &NodeAgentReleaseInfo{
|
||||
CurrentVersion: currentVersion,
|
||||
Channel: channel.String(),
|
||||
UpdateRequested: node.UpdateRequested,
|
||||
RequestedChannel: normalizeReleaseChannel(node.UpdateChannel).String(),
|
||||
RequestedTag: strings.TrimSpace(node.UpdateTag),
|
||||
}
|
||||
if release == nil {
|
||||
return view
|
||||
}
|
||||
view.TagName = release.TagName
|
||||
view.Body = release.Body
|
||||
view.HTMLURL = release.HTMLURL
|
||||
view.PublishedAt = release.PublishedAt
|
||||
view.Prerelease = release.Prerelease
|
||||
view.HasUpdate = isVersionNewer(currentVersion, release.TagName)
|
||||
return view
|
||||
}
|
||||
|
||||
func RegisterNodeWithAccessToken(node *model.Node, payload AgentNodePayload) (*AgentRegistrationResponse, error) {
|
||||
payload = normalizeAgentNodePayload(payload)
|
||||
if node == nil {
|
||||
return nil, errors.New("节点不存在")
|
||||
}
|
||||
if err := validateAgentNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applyNodeRuntime(node, payload, true)
|
||||
if err := node.Update(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
slog.Info("agent register succeeded on reserved node", "node_id", node.NodeID, "name", node.Name)
|
||||
return &AgentRegistrationResponse{
|
||||
NodeID: node.NodeID,
|
||||
AccessToken: node.AccessToken,
|
||||
Name: node.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func RegisterNodeWithDiscovery(payload AgentNodePayload) (*AgentRegistrationResponse, error) {
|
||||
payload = normalizeAgentNodePayload(payload)
|
||||
if err := validateAgentNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeID, err := newServerNodeID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
agentToken, err := newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeName := payload.Name
|
||||
if nodeName == "" {
|
||||
nodeName = nodeID
|
||||
}
|
||||
node := &model.Node{
|
||||
NodeID: nodeID,
|
||||
Name: nodeName,
|
||||
AccessToken: agentToken,
|
||||
}
|
||||
applyNodeRuntime(node, payload, false)
|
||||
if err = node.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("节点标识生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
slog.Info("agent discovery register succeeded", "node_id", node.NodeID, "name", node.Name)
|
||||
return &AgentRegistrationResponse{
|
||||
NodeID: node.NodeID,
|
||||
AccessToken: node.AccessToken,
|
||||
Name: node.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeAgentNodePayload(payload AgentNodePayload) AgentNodePayload {
|
||||
payload.Name = strings.TrimSpace(payload.Name)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
return payload
|
||||
}
|
||||
|
||||
func validateAgentNodePayload(payload AgentNodePayload) error {
|
||||
if payload.IP == "" {
|
||||
return errors.New("ip 不能为空")
|
||||
}
|
||||
if net.ParseIP(payload.IP) == nil {
|
||||
return errors.New("ip 格式无效")
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return errors.New("version 不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func applyNodeRuntime(node *model.Node, payload AgentNodePayload, preserveName bool) {
|
||||
if !preserveName || strings.TrimSpace(node.Name) == "" {
|
||||
if strings.TrimSpace(payload.Name) != "" {
|
||||
node.Name = strings.TrimSpace(payload.Name)
|
||||
}
|
||||
}
|
||||
if !node.IPManualOverride {
|
||||
node.IP = strings.TrimSpace(payload.IP)
|
||||
}
|
||||
node.Version = strings.TrimSpace(payload.Version)
|
||||
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
node.Status = NodeStatusOnline
|
||||
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
node.LastSeenAt = time.Now()
|
||||
node.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
}
|
||||
|
||||
func applyGeoInfoFromIP(node *model.Node, rawIP string) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
node.GeoName = ""
|
||||
node.GeoLatitude = nil
|
||||
node.GeoLongitude = nil
|
||||
ip := net.ParseIP(strings.TrimSpace(rawIP))
|
||||
if ip == nil {
|
||||
return
|
||||
}
|
||||
info, err := geoip.GetGeoInfo(ip)
|
||||
if err != nil || info == nil {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(info.Name) != "" {
|
||||
node.GeoName = strings.TrimSpace(info.Name)
|
||||
}
|
||||
if info.Latitude != nil && info.Longitude != nil {
|
||||
node.GeoLatitude = cloneCoordinate(info.Latitude)
|
||||
node.GeoLongitude = cloneCoordinate(info.Longitude)
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeOpenrestyStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case OpenrestyStatusHealthy:
|
||||
return OpenrestyStatusHealthy
|
||||
case OpenrestyStatusUnhealthy:
|
||||
return OpenrestyStatusUnhealthy
|
||||
default:
|
||||
return OpenrestyStatusUnknown
|
||||
}
|
||||
}
|
||||
|
||||
func newRandomToken() (string, error) {
|
||||
buf := make([]byte, 16)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func newServerNodeID() (string, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + token, nil
|
||||
}
|
||||
|
||||
func normalizeNodeType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "tunnel_relay":
|
||||
return "tunnel_relay"
|
||||
case "tunnel_client":
|
||||
return "tunnel_client"
|
||||
default:
|
||||
return "edge_node"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeRelayPort(port int, defaultPort int) int {
|
||||
if port <= 0 || port > 65535 {
|
||||
return defaultPort
|
||||
}
|
||||
return port
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
ristretto "github.com/dgraph-io/ristretto/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
agentTokenPositiveCacheTTL = 2 * time.Minute
|
||||
agentTokenNegativeCacheTTL = 10 * time.Minute
|
||||
agentTokenNegativeCacheCap = 10000
|
||||
)
|
||||
|
||||
type cachedAgentNode struct {
|
||||
node *model.Node
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type cachedMissingAccessToken struct {
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type agentTokenAuthCache struct {
|
||||
positive *ristretto.Cache[string, cachedAgentNode]
|
||||
negative *ristretto.Cache[string, cachedMissingAccessToken]
|
||||
now func() time.Time
|
||||
loadNodeByToken func(string) (*model.Node, error)
|
||||
}
|
||||
|
||||
var nodeAccessTokenCache = newAccessTokenAuthCache()
|
||||
|
||||
func newAccessTokenAuthCache() *agentTokenAuthCache {
|
||||
return &agentTokenAuthCache{
|
||||
positive: mustNewAccessTokenPositiveCache(),
|
||||
negative: mustNewAccessTokenNegativeCache(),
|
||||
now: time.Now,
|
||||
loadNodeByToken: func(token string) (*model.Node, error) {
|
||||
return model.GetNodeByAccessToken(token)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func mustNewAccessTokenPositiveCache() *ristretto.Cache[string, cachedAgentNode] {
|
||||
cache, err := ristretto.NewCache(&ristretto.Config[string, cachedAgentNode]{
|
||||
NumCounters: 1e5,
|
||||
MaxCost: 2e4,
|
||||
BufferItems: 64,
|
||||
})
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return cache
|
||||
}
|
||||
|
||||
func mustNewAccessTokenNegativeCache() *ristretto.Cache[string, cachedMissingAccessToken] {
|
||||
cache, err := ristretto.NewCache(&ristretto.Config[string, cachedMissingAccessToken]{
|
||||
NumCounters: 1e5,
|
||||
MaxCost: agentTokenNegativeCacheCap,
|
||||
BufferItems: 64,
|
||||
})
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return cache
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) authenticate(token string) (*model.Node, error) {
|
||||
now := c.now()
|
||||
if node, ok := c.getNode(token, now); ok {
|
||||
return node, nil
|
||||
}
|
||||
if c.isMissing(token, now) {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
node, err := c.loadNodeByToken(token)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.storeNode(token, node, now.Add(agentTokenPositiveCacheTTL))
|
||||
return cloneCachedNode(node), nil
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) getNode(token string, now time.Time) (*model.Node, bool) {
|
||||
entry, ok := c.positive.Get(token)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if now.After(entry.expiresAt) {
|
||||
c.positive.Del(token)
|
||||
return nil, false
|
||||
}
|
||||
return cloneCachedNode(entry.node), true
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) isMissing(token string, now time.Time) bool {
|
||||
entry, ok := c.negative.Get(token)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if now.After(entry.expiresAt) {
|
||||
c.negative.Del(token)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) storeNode(token string, node *model.Node, expiresAt time.Time) {
|
||||
if token == "" || node == nil {
|
||||
return
|
||||
}
|
||||
c.negative.Del(token)
|
||||
c.positive.Set(token, cachedAgentNode{
|
||||
node: cloneCachedNode(node),
|
||||
expiresAt: expiresAt,
|
||||
}, 1)
|
||||
c.positive.Wait()
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
|
||||
if token == "" {
|
||||
return
|
||||
}
|
||||
c.positive.Del(token)
|
||||
c.negative.Set(token, cachedMissingAccessToken{
|
||||
expiresAt: expiresAt,
|
||||
}, 1)
|
||||
c.negative.Wait()
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) invalidate(token string) {
|
||||
if token == "" {
|
||||
return
|
||||
}
|
||||
c.positive.Del(token)
|
||||
c.negative.Del(token)
|
||||
}
|
||||
|
||||
func (c *agentTokenAuthCache) reset() {
|
||||
c.positive.Clear()
|
||||
c.negative.Clear()
|
||||
}
|
||||
|
||||
func cloneCachedNode(node *model.Node) *model.Node {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *node
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func authenticateAccessTokenWithCache(token string) (*model.Node, error) {
|
||||
return nodeAccessTokenCache.authenticate(token)
|
||||
}
|
||||
|
||||
func refreshAccessTokenCache(node *model.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
nodeAccessTokenCache.storeNode(
|
||||
node.AccessToken,
|
||||
node,
|
||||
nodeAccessTokenCache.now().Add(agentTokenPositiveCacheTTL),
|
||||
)
|
||||
}
|
||||
|
||||
func invalidateAccessTokenCache(token string) {
|
||||
nodeAccessTokenCache.invalidate(token)
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestAccessTokenAuthCacheUsesPositiveCacheUntilLogicalExpiry(t *testing.T) {
|
||||
cache := newAccessTokenAuthCache()
|
||||
cache.reset()
|
||||
baseTime := time.Date(2026, 3, 14, 16, 0, 0, 0, time.UTC)
|
||||
currentTime := baseTime
|
||||
cache.now = func() time.Time {
|
||||
return currentTime
|
||||
}
|
||||
|
||||
loadCount := 0
|
||||
cache.loadNodeByToken = func(token string) (*model.Node, error) {
|
||||
loadCount++
|
||||
return &model.Node{
|
||||
NodeID: fmt.Sprintf("node-%d", loadCount),
|
||||
Name: "edge",
|
||||
AccessToken: token,
|
||||
}, nil
|
||||
}
|
||||
|
||||
first, err := cache.authenticate("token-a")
|
||||
if err != nil {
|
||||
t.Fatalf("expected first auth to succeed: %v", err)
|
||||
}
|
||||
if loadCount != 1 {
|
||||
t.Fatalf("expected one db load, got %d", loadCount)
|
||||
}
|
||||
|
||||
second, err := cache.authenticate("token-a")
|
||||
if err != nil {
|
||||
t.Fatalf("expected cached auth to succeed: %v", err)
|
||||
}
|
||||
if loadCount != 1 {
|
||||
t.Fatalf("expected cache hit without db load, got %d", loadCount)
|
||||
}
|
||||
if first.NodeID != second.NodeID {
|
||||
t.Fatalf("expected cached node to match original, got %s and %s", first.NodeID, second.NodeID)
|
||||
}
|
||||
|
||||
currentTime = baseTime.Add(agentTokenPositiveCacheTTL + time.Second)
|
||||
third, err := cache.authenticate("token-a")
|
||||
if err != nil {
|
||||
t.Fatalf("expected auth after expiry to succeed: %v", err)
|
||||
}
|
||||
if loadCount != 2 {
|
||||
t.Fatalf("expected reload after logical expiry, got %d loads", loadCount)
|
||||
}
|
||||
if third.NodeID == second.NodeID {
|
||||
t.Fatalf("expected refreshed cache entry after expiry, got unchanged node id %s", third.NodeID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccessTokenAuthCacheRefreshesAfterMissingEntryExpires(t *testing.T) {
|
||||
cache := newAccessTokenAuthCache()
|
||||
cache.reset()
|
||||
baseTime := time.Date(2026, 3, 14, 16, 30, 0, 0, time.UTC)
|
||||
currentTime := baseTime
|
||||
cache.now = func() time.Time {
|
||||
return currentTime
|
||||
}
|
||||
|
||||
loadCount := 0
|
||||
cache.loadNodeByToken = func(token string) (*model.Node, error) {
|
||||
loadCount++
|
||||
if loadCount == 1 {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
return &model.Node{
|
||||
NodeID: "node-recovered",
|
||||
Name: "edge",
|
||||
AccessToken: token,
|
||||
}, nil
|
||||
}
|
||||
|
||||
_, err := cache.authenticate("token-missing")
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("expected first lookup to miss, got %v", err)
|
||||
}
|
||||
if loadCount != 1 {
|
||||
t.Fatalf("expected one db load for first miss, got %d", loadCount)
|
||||
}
|
||||
|
||||
_, err = cache.authenticate("token-missing")
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("expected cached missing lookup to miss, got %v", err)
|
||||
}
|
||||
if loadCount != 1 {
|
||||
t.Fatalf("expected missing cache hit without db load, got %d", loadCount)
|
||||
}
|
||||
|
||||
currentTime = baseTime.Add(agentTokenNegativeCacheTTL + time.Second)
|
||||
node, err := cache.authenticate("token-missing")
|
||||
if err != nil {
|
||||
t.Fatalf("expected lookup after missing expiry to reload successfully: %v", err)
|
||||
}
|
||||
if loadCount != 2 {
|
||||
t.Fatalf("expected db reload after missing cache expiry, got %d", loadCount)
|
||||
}
|
||||
if node.NodeID != "node-recovered" {
|
||||
t.Fatalf("unexpected recovered node: %+v", node)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultObservabilityWindow = 24 * time.Hour
|
||||
defaultObservabilityLimit = 120
|
||||
maxObservabilityLimit = 500
|
||||
)
|
||||
|
||||
type NodeObservabilityQuery struct {
|
||||
Hours int `json:"hours"`
|
||||
Limit int `json:"limit"`
|
||||
}
|
||||
|
||||
type NodeObservabilityView struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Profile *model.NodeSystemProfile `json:"profile"`
|
||||
MetricSnapshots []*model.NodeMetricSnapshot `json:"metric_snapshots"`
|
||||
TrafficReports []*model.NodeRequestReport `json:"traffic_reports"`
|
||||
HealthEvents []*model.NodeHealthEvent `json:"health_events"`
|
||||
Analytics NodeObservabilityAnalytics `json:"analytics"`
|
||||
Trends NodeObservabilityTrends `json:"trends"`
|
||||
RelayDashboard *RelayDashboardSnapshot `json:"relay_dashboard,omitempty"`
|
||||
}
|
||||
|
||||
type NodeObservabilityAnalytics struct {
|
||||
Traffic TrafficWindowSummary `json:"traffic"`
|
||||
Distributions TrafficDistributions `json:"distributions"`
|
||||
Health ObservabilityHealthSummary `json:"health"`
|
||||
}
|
||||
|
||||
type NodeObservabilityTrends struct {
|
||||
Traffic24h []TrafficTrendPoint `json:"traffic_24h"`
|
||||
Capacity24h []CapacityTrendPoint `json:"capacity_24h"`
|
||||
Network24h []NetworkTrendPoint `json:"network_24h"`
|
||||
DiskIO24h []DiskIOTrendPoint `json:"disk_io_24h"`
|
||||
}
|
||||
|
||||
type NodeHealthEventCleanupResult struct {
|
||||
NodeID string `json:"node_id"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
}
|
||||
|
||||
type RelayDashboardSnapshot struct {
|
||||
TotalProxies int `json:"total_proxies"`
|
||||
OnlineProxies int `json:"online_proxies"`
|
||||
OfflineProxies int `json:"offline_proxies"`
|
||||
Proxies []RelayProxyStat `json:"proxies"`
|
||||
TotalConnections int `json:"total_connections"`
|
||||
ClientCounts int `json:"client_counts"`
|
||||
}
|
||||
|
||||
type RelayProxyStat struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Status string `json:"status"`
|
||||
ClientVersion string `json:"client_version"`
|
||||
LastStartTime string `json:"last_start_time"`
|
||||
LastCloseTime string `json:"last_close_time"`
|
||||
ClientAddr string `json:"client_addr"`
|
||||
}
|
||||
|
||||
func GetNodeObservability(id uint, query NodeObservabilityQuery) (*NodeObservabilityView, error) {
|
||||
now := time.Now()
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
limit := normalizeObservabilityLimit(query.Limit)
|
||||
since := now.Add(-normalizeObservabilityWindow(query.Hours))
|
||||
|
||||
profile, err := model.GetNodeSystemProfile(node.NodeID)
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
profile = nil
|
||||
}
|
||||
|
||||
snapshots, err := model.ListNodeMetricSnapshots(node.NodeID, since, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reports, err := model.ListNodeRequestReports(node.NodeID, since, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessLogRegions, err := model.ListNodeAccessLogRegionCounts(node.NodeID, since, 8)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trendSnapshots, err := model.ListNodeMetricSnapshots(node.NodeID, now.Add(-24*time.Hour), 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trendOpenresty, _ := model.ListNodeObservationOpenresty(node.NodeID, now.Add(-24*time.Hour), 0)
|
||||
trendReports, err := model.ListNodeRequestReports(node.NodeID, now.Add(-24*time.Hour), 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
events, err := model.ListNodeHealthEvents(node.NodeID, false, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
view := &NodeObservabilityView{
|
||||
NodeID: node.NodeID,
|
||||
Profile: profile,
|
||||
MetricSnapshots: snapshots,
|
||||
TrafficReports: reports,
|
||||
HealthEvents: events,
|
||||
Analytics: NodeObservabilityAnalytics{
|
||||
Traffic: buildTrafficWindowSummary(latestTrafficReport(reports)),
|
||||
Distributions: buildTrafficDistributions(reports, accessLogRegions, 8),
|
||||
Health: buildObservabilityHealthSummary(latestMetricSnapshot(snapshots), latestTrafficReport(reports), events),
|
||||
},
|
||||
Trends: NodeObservabilityTrends{
|
||||
Traffic24h: buildTrafficTrendPoints(now, trendReports),
|
||||
Capacity24h: buildCapacityTrendPoints(now, trendSnapshots),
|
||||
Network24h: buildNetworkTrendPoints(now, trendSnapshots, trendOpenresty),
|
||||
DiskIO24h: buildDiskIOTrendPoints(now, trendSnapshots),
|
||||
},
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
frpsObs, _ := model.ListNodeObservationFrps(node.NodeID, time.Time{}, 1)
|
||||
var latestFrps *model.NodeObservationFrps
|
||||
if len(frpsObs) > 0 {
|
||||
latestFrps = frpsObs[0]
|
||||
}
|
||||
view.RelayDashboard = buildRelayDashboardSnapshot(node, latestFrps)
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func buildRelayDashboardSnapshot(node *model.Node, obs *model.NodeObservationFrps) *RelayDashboardSnapshot {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
totalProxies := 0
|
||||
totalConnections := 0
|
||||
clientCounts := 0
|
||||
proxies := []RelayProxyStat{}
|
||||
|
||||
if obs != nil {
|
||||
totalProxies = obs.FrpsProxyCount
|
||||
totalConnections = obs.FrpsConnections
|
||||
clientCounts = obs.FrpsClientCount
|
||||
if obs.FrpsProxies != "" {
|
||||
var decoded []RelayProxyStat
|
||||
if err := json.Unmarshal([]byte(obs.FrpsProxies), &decoded); err == nil {
|
||||
proxies = decoded
|
||||
}
|
||||
}
|
||||
}
|
||||
if totalProxies < 0 {
|
||||
totalProxies = 0
|
||||
}
|
||||
onlineProxies := 0
|
||||
for _, p := range proxies {
|
||||
if p.Status == "online" {
|
||||
onlineProxies++
|
||||
}
|
||||
}
|
||||
// Fallback for backward compatibility
|
||||
if len(proxies) == 0 {
|
||||
onlineProxies = totalProxies
|
||||
if node.RelayStatus != "healthy" {
|
||||
onlineProxies = 0
|
||||
}
|
||||
}
|
||||
|
||||
return &RelayDashboardSnapshot{
|
||||
TotalProxies: totalProxies,
|
||||
OnlineProxies: onlineProxies,
|
||||
OfflineProxies: totalProxies - onlineProxies,
|
||||
Proxies: proxies,
|
||||
TotalConnections: maxInt(totalConnections, 0),
|
||||
ClientCounts: maxInt(clientCounts, 0),
|
||||
}
|
||||
}
|
||||
|
||||
func maxInt(a int, b int) int {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func CleanupNodeHealthEvents(id uint) (*NodeHealthEventCleanupResult, error) {
|
||||
node, err := model.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deletedCount, err := model.DeleteNodeHealthEvents(node.NodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &NodeHealthEventCleanupResult{
|
||||
NodeID: node.NodeID,
|
||||
DeletedCount: deletedCount,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func latestMetricSnapshot(snapshots []*model.NodeMetricSnapshot) *model.NodeMetricSnapshot {
|
||||
for _, snapshot := range snapshots {
|
||||
if snapshot != nil {
|
||||
return snapshot
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func latestTrafficReport(reports []*model.NodeRequestReport) *model.NodeRequestReport {
|
||||
for _, report := range reports {
|
||||
if report != nil {
|
||||
return report
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeObservabilityLimit(limit int) int {
|
||||
if limit <= 0 {
|
||||
return defaultObservabilityLimit
|
||||
}
|
||||
if limit > maxObservabilityLimit {
|
||||
return maxObservabilityLimit
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func normalizeObservabilityWindow(hours int) time.Duration {
|
||||
if hours <= 0 {
|
||||
return defaultObservabilityWindow
|
||||
}
|
||||
return time.Duration(hours) * time.Hour
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,410 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
NodeHealthEventStatusActive = "active"
|
||||
NodeHealthEventStatusResolved = "resolved"
|
||||
NodeHealthSeverityInfo = "info"
|
||||
NodeHealthSeverityWarning = "warning"
|
||||
NodeHealthSeverityCritical = "critical"
|
||||
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
|
||||
nodeAccessLogPathMaxLength = 100
|
||||
)
|
||||
|
||||
type AgentNodeSystemProfile struct {
|
||||
Hostname string `json:"hostname"`
|
||||
OSName string `json:"os_name"`
|
||||
OSVersion string `json:"os_version"`
|
||||
KernelVersion string `json:"kernel_version"`
|
||||
Architecture string `json:"architecture"`
|
||||
CPUModel string `json:"cpu_model"`
|
||||
CPUCores int `json:"cpu_cores"`
|
||||
TotalMemoryBytes int64 `json:"total_memory_bytes"`
|
||||
TotalDiskBytes int64 `json:"total_disk_bytes"`
|
||||
UptimeSeconds int64 `json:"uptime_seconds"`
|
||||
ReportedAtUnix int64 `json:"reported_at_unix"`
|
||||
}
|
||||
|
||||
type AgentNodeMetricSnapshot struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
CPUUsagePercent float64 `json:"cpu_usage_percent"`
|
||||
MemoryUsedBytes int64 `json:"memory_used_bytes"`
|
||||
MemoryTotalBytes int64 `json:"memory_total_bytes"`
|
||||
StorageUsedBytes int64 `json:"storage_used_bytes"`
|
||||
StorageTotalBytes int64 `json:"storage_total_bytes"`
|
||||
DiskReadBytes int64 `json:"disk_read_bytes"`
|
||||
DiskWriteBytes int64 `json:"disk_write_bytes"`
|
||||
NetworkRxBytes int64 `json:"network_rx_bytes"`
|
||||
NetworkTxBytes int64 `json:"network_tx_bytes"`
|
||||
}
|
||||
|
||||
type AgentNodeOpenrestyObservation struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
OpenrestyConnections int64 `json:"openresty_connections"`
|
||||
}
|
||||
|
||||
type AgentNodeTrafficReport struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
StatusCodes map[string]int64 `json:"status_codes"`
|
||||
TopDomains map[string]int64 `json:"top_domains"`
|
||||
SourceCountries map[string]int64 `json:"source_countries"`
|
||||
}
|
||||
|
||||
type AgentNodeAccessLog struct {
|
||||
LoggedAtUnix int64 `json:"logged_at_unix"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
}
|
||||
|
||||
type AgentBufferedObservabilityRecord struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
Snapshot *AgentNodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
OpenrestyObservation *AgentNodeOpenrestyObservation `json:"openresty_observation,omitempty"`
|
||||
TrafficReport *AgentNodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []AgentNodeAccessLog `json:"access_logs,omitempty"`
|
||||
}
|
||||
|
||||
type AgentNodeHealthEvent struct {
|
||||
EventType string `json:"event_type"`
|
||||
Severity string `json:"severity"`
|
||||
Message string `json:"message"`
|
||||
TriggeredAtUnix int64 `json:"triggered_at_unix"`
|
||||
Metadata map[string]string `json:"metadata"`
|
||||
}
|
||||
|
||||
func persistHeartbeatObservability(nodeID string, payload AgentNodePayload, reportedAt time.Time) {
|
||||
if strings.TrimSpace(nodeID) == "" {
|
||||
return
|
||||
}
|
||||
if payload.Profile == nil && payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0 && len(payload.BufferedObservability) == 0 && payload.HealthEvents == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistBufferedObservability(tx, nodeID, payload.BufferedObservability, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeMetricSnapshot(tx, nodeID, payload.Snapshot, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeOpenrestyObservation(tx, nodeID, payload.OpenrestyObservation, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeTrafficReport(tx, nodeID, payload.TrafficReport, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeAccessLogs(tx, nodeID, payload.AccessLogs, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if payload.HealthEvents != nil {
|
||||
if err := reconcileNodeHealthEvents(tx, nodeID, payload.HealthEvents, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
slog.Error("persist heartbeat observability failed", "node_id", nodeID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func persistBufferedObservability(tx *gorm.DB, nodeID string, records []AgentBufferedObservabilityRecord, reportedAt time.Time) error {
|
||||
for _, record := range records {
|
||||
if err := persistNodeMetricSnapshot(tx, nodeID, record.Snapshot, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeOpenrestyObservation(tx, nodeID, record.OpenrestyObservation, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeTrafficReport(tx, nodeID, record.TrafficReport, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeAccessLogs(tx, nodeID, record.AccessLogs, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *AgentNodeSystemProfile, reportedAt time.Time) error {
|
||||
if profile == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.NodeSystemProfile{
|
||||
NodeID: nodeID,
|
||||
Hostname: strings.TrimSpace(profile.Hostname),
|
||||
OSName: strings.TrimSpace(profile.OSName),
|
||||
OSVersion: strings.TrimSpace(profile.OSVersion),
|
||||
KernelVersion: strings.TrimSpace(profile.KernelVersion),
|
||||
Architecture: strings.TrimSpace(profile.Architecture),
|
||||
CPUModel: strings.TrimSpace(profile.CPUModel),
|
||||
CPUCores: profile.CPUCores,
|
||||
TotalMemoryBytes: profile.TotalMemoryBytes,
|
||||
TotalDiskBytes: profile.TotalDiskBytes,
|
||||
UptimeSeconds: profile.UptimeSeconds,
|
||||
ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt),
|
||||
}
|
||||
return tx.Model(&model.NodeSystemProfile{}).Where("node_id = ?", nodeID).Assign(record).FirstOrCreate(record).Error
|
||||
}
|
||||
|
||||
func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *AgentNodeMetricSnapshot, reportedAt time.Time) error {
|
||||
if snapshot == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.NodeMetricSnapshot{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: timeFromUnix(snapshot.CapturedAtUnix, reportedAt),
|
||||
CPUUsagePercent: snapshot.CPUUsagePercent,
|
||||
MemoryUsedBytes: snapshot.MemoryUsedBytes,
|
||||
MemoryTotalBytes: snapshot.MemoryTotalBytes,
|
||||
StorageUsedBytes: snapshot.StorageUsedBytes,
|
||||
StorageTotalBytes: snapshot.StorageTotalBytes,
|
||||
DiskReadBytes: snapshot.DiskReadBytes,
|
||||
DiskWriteBytes: snapshot.DiskWriteBytes,
|
||||
NetworkRxBytes: snapshot.NetworkRxBytes,
|
||||
NetworkTxBytes: snapshot.NetworkTxBytes,
|
||||
}
|
||||
exists, err := model.NodeMetricSnapshotExists(tx, nodeID, record.CapturedAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *AgentNodeOpenrestyObservation, reportedAt time.Time) error {
|
||||
if obs == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.NodeObservationOpenresty{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: timeFromUnix(obs.CapturedAtUnix, reportedAt),
|
||||
OpenrestyRxBytes: obs.OpenrestyRxBytes,
|
||||
OpenrestyTxBytes: obs.OpenrestyTxBytes,
|
||||
OpenrestyConnections: obs.OpenrestyConnections,
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *AgentNodeTrafficReport, reportedAt time.Time) error {
|
||||
if report == nil {
|
||||
return nil
|
||||
}
|
||||
if report.WindowEndedAtUnix > 0 && report.WindowStartedAtUnix > report.WindowEndedAtUnix {
|
||||
return errors.New("traffic report window_started_at_unix 不能大于 window_ended_at_unix")
|
||||
}
|
||||
record := &model.NodeRequestReport{
|
||||
NodeID: nodeID,
|
||||
WindowStartedAt: timeFromUnix(report.WindowStartedAtUnix, reportedAt),
|
||||
WindowEndedAt: timeFromUnix(report.WindowEndedAtUnix, reportedAt),
|
||||
RequestCount: report.RequestCount,
|
||||
ErrorCount: report.ErrorCount,
|
||||
UniqueVisitorCount: report.UniqueVisitorCount,
|
||||
StatusCodesJSON: marshalJSON(report.StatusCodes),
|
||||
TopDomainsJSON: marshalJSON(report.TopDomains),
|
||||
SourceCountriesJSON: marshalJSON(report.SourceCountries),
|
||||
}
|
||||
exists, err := model.NodeRequestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []AgentNodeAccessLog, reportedAt time.Time) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
resolver, err := newAccessLogRegionResolver()
|
||||
if err != nil {
|
||||
slog.Warn("initialize access log geo resolver failed", "node_id", nodeID, "error", err)
|
||||
}
|
||||
if resolver != nil {
|
||||
defer resolver.Close()
|
||||
}
|
||||
for _, item := range logs {
|
||||
record := &model.NodeAccessLog{
|
||||
NodeID: nodeID,
|
||||
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
|
||||
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
|
||||
Region: "",
|
||||
Host: strings.TrimSpace(item.Host),
|
||||
Path: truncateForDatabase(strings.TrimSpace(item.Path), nodeAccessLogPathMaxLength),
|
||||
StatusCode: item.StatusCode,
|
||||
}
|
||||
if resolver != nil {
|
||||
record.Region = resolver.Resolve(record.RemoteAddr)
|
||||
}
|
||||
exists, err := model.NodeAccessLogExists(tx, record)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
continue
|
||||
}
|
||||
if err := tx.Create(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
_, err = model.DeleteNodeAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
|
||||
return err
|
||||
}
|
||||
|
||||
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentNodeHealthEvent, reportedAt time.Time) error {
|
||||
return reconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
|
||||
}
|
||||
|
||||
func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentNodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
|
||||
activeTypes := make(map[string]AgentNodeHealthEvent, len(events))
|
||||
for _, event := range events {
|
||||
eventType := normalizeHealthEventType(event.EventType)
|
||||
if eventType == "" {
|
||||
continue
|
||||
}
|
||||
if len(managedEventTypes) > 0 {
|
||||
if _, ok := managedEventTypes[eventType]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
event.EventType = eventType
|
||||
event.Severity = normalizeHealthSeverity(event.Severity)
|
||||
if event.TriggeredAtUnix <= 0 {
|
||||
event.TriggeredAtUnix = reportedAt.Unix()
|
||||
}
|
||||
activeTypes[eventType] = event
|
||||
}
|
||||
|
||||
var activeEvents []*model.NodeHealthEvent
|
||||
query := tx.Where("node_id = ? AND status = ?", nodeID, NodeHealthEventStatusActive)
|
||||
if len(managedEventTypes) > 0 {
|
||||
scopedTypes := make([]string, 0, len(managedEventTypes))
|
||||
for eventType := range managedEventTypes {
|
||||
eventType = normalizeHealthEventType(eventType)
|
||||
if eventType != "" {
|
||||
scopedTypes = append(scopedTypes, eventType)
|
||||
}
|
||||
}
|
||||
if len(scopedTypes) == 0 {
|
||||
return nil
|
||||
}
|
||||
query = query.Where("event_type IN ?", scopedTypes)
|
||||
}
|
||||
if err := query.Find(&activeEvents).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
activeByType := make(map[string]*model.NodeHealthEvent, len(activeEvents))
|
||||
for _, event := range activeEvents {
|
||||
activeByType[event.EventType] = event
|
||||
}
|
||||
|
||||
for eventType, event := range activeTypes {
|
||||
triggeredAt := timeFromUnix(event.TriggeredAtUnix, reportedAt)
|
||||
if existing, ok := activeByType[eventType]; ok {
|
||||
existing.Severity = event.Severity
|
||||
existing.Message = normalizeHealthEventMessage(event.Message)
|
||||
existing.LastTriggeredAt = triggeredAt
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.MetadataJSON = marshalJSON(event.Metadata)
|
||||
existing.ResolvedAt = nil
|
||||
if err := tx.Save(existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
record := &model.NodeHealthEvent{
|
||||
NodeID: nodeID,
|
||||
EventType: eventType,
|
||||
Severity: event.Severity,
|
||||
Status: NodeHealthEventStatusActive,
|
||||
Message: normalizeHealthEventMessage(event.Message),
|
||||
FirstTriggeredAt: triggeredAt,
|
||||
LastTriggeredAt: triggeredAt,
|
||||
ReportedAt: reportedAt,
|
||||
MetadataJSON: marshalJSON(event.Metadata),
|
||||
}
|
||||
if err := tx.Create(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
for _, existing := range activeEvents {
|
||||
if _, ok := activeTypes[existing.EventType]; ok {
|
||||
continue
|
||||
}
|
||||
resolvedAt := reportedAt
|
||||
existing.Status = NodeHealthEventStatusResolved
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.ResolvedAt = &resolvedAt
|
||||
if err := tx.Save(existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeHealthEventType(eventType string) string {
|
||||
eventType = strings.TrimSpace(strings.ToLower(eventType))
|
||||
eventType = strings.ReplaceAll(eventType, " ", "_")
|
||||
return eventType
|
||||
}
|
||||
|
||||
func normalizeHealthSeverity(severity string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(severity)) {
|
||||
case NodeHealthSeverityCritical:
|
||||
return NodeHealthSeverityCritical
|
||||
case NodeHealthSeverityInfo:
|
||||
return NodeHealthSeverityInfo
|
||||
default:
|
||||
return NodeHealthSeverityWarning
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeHealthEventMessage(message string) string {
|
||||
return truncateForDatabase(message, 4096)
|
||||
}
|
||||
|
||||
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
|
||||
if unixSeconds <= 0 {
|
||||
return fallback
|
||||
}
|
||||
return time.Unix(unixSeconds, 0).UTC()
|
||||
}
|
||||
|
||||
func marshalJSON(value any) string {
|
||||
if value == nil {
|
||||
return ""
|
||||
}
|
||||
raw, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
type DistributionItem struct {
|
||||
Key string `json:"key"`
|
||||
Value int64 `json:"value"`
|
||||
}
|
||||
|
||||
type TrafficDistributions struct {
|
||||
StatusCodes []DistributionItem `json:"status_codes"`
|
||||
TopDomains []DistributionItem `json:"top_domains"`
|
||||
SourceCountries []DistributionItem `json:"source_countries"`
|
||||
}
|
||||
|
||||
type TrafficWindowSummary struct {
|
||||
WindowStartedAt time.Time `json:"window_started_at"`
|
||||
WindowEndedAt time.Time `json:"window_ended_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
EstimatedQPS float64 `json:"estimated_qps"`
|
||||
ErrorRatePercent float64 `json:"error_rate_percent"`
|
||||
}
|
||||
|
||||
type ObservabilityHealthSummary struct {
|
||||
ActiveAlerts int `json:"active_alerts"`
|
||||
CriticalAlerts int `json:"critical_alerts"`
|
||||
WarningAlerts int `json:"warning_alerts"`
|
||||
InfoAlerts int `json:"info_alerts"`
|
||||
ResolvedAlerts int `json:"resolved_alerts"`
|
||||
HasCapacityRisk bool `json:"has_capacity_risk"`
|
||||
HasTrafficRisk bool `json:"has_traffic_risk"`
|
||||
HasRuntimeRisk bool `json:"has_runtime_risk"`
|
||||
}
|
||||
|
||||
type distributionAccumulator map[string]int64
|
||||
|
||||
func buildTrafficWindowSummary(report *model.NodeRequestReport) TrafficWindowSummary {
|
||||
if report == nil {
|
||||
return TrafficWindowSummary{}
|
||||
}
|
||||
summary := TrafficWindowSummary{
|
||||
WindowStartedAt: report.WindowStartedAt,
|
||||
WindowEndedAt: report.WindowEndedAt,
|
||||
RequestCount: report.RequestCount,
|
||||
UniqueVisitorCount: report.UniqueVisitorCount,
|
||||
ErrorCount: report.ErrorCount,
|
||||
}
|
||||
if duration := report.WindowEndedAt.Sub(report.WindowStartedAt).Seconds(); duration > 0 {
|
||||
summary.EstimatedQPS = float64(report.RequestCount) / duration
|
||||
}
|
||||
if report.RequestCount > 0 {
|
||||
summary.ErrorRatePercent = (float64(report.ErrorCount) / float64(report.RequestCount)) * 100
|
||||
}
|
||||
return summary
|
||||
}
|
||||
|
||||
func buildTrafficDistributions(
|
||||
reports []*model.NodeRequestReport,
|
||||
accessLogRegions []*model.NodeAccessLogRegionCount,
|
||||
limit int,
|
||||
) TrafficDistributions {
|
||||
statusCodes := make(distributionAccumulator)
|
||||
topDomains := make(distributionAccumulator)
|
||||
reportSourceCountries := make(distributionAccumulator)
|
||||
for _, report := range reports {
|
||||
mergeJSONCounts(statusCodes, report.StatusCodesJSON)
|
||||
mergeJSONCounts(topDomains, report.TopDomainsJSON)
|
||||
mergeJSONCounts(reportSourceCountries, report.SourceCountriesJSON)
|
||||
}
|
||||
sourceCountries := reportSourceCountries
|
||||
if len(accessLogRegions) > 0 {
|
||||
sourceCountries = make(distributionAccumulator, len(accessLogRegions))
|
||||
for _, item := range accessLogRegions {
|
||||
if item == nil || strings.TrimSpace(item.Region) == "" || item.Count <= 0 {
|
||||
continue
|
||||
}
|
||||
sourceCountries[item.Region] = item.Count
|
||||
}
|
||||
}
|
||||
return TrafficDistributions{
|
||||
StatusCodes: toDistributionItems(statusCodes, limit),
|
||||
TopDomains: toDistributionItems(topDomains, limit),
|
||||
SourceCountries: toDistributionItems(sourceCountries, limit),
|
||||
}
|
||||
}
|
||||
|
||||
func buildObservabilityHealthSummary(snapshot *model.NodeMetricSnapshot, report *model.NodeRequestReport, events []*model.NodeHealthEvent) ObservabilityHealthSummary {
|
||||
summary := ObservabilityHealthSummary{}
|
||||
for _, event := range events {
|
||||
if event == nil {
|
||||
continue
|
||||
}
|
||||
if event.Status == NodeHealthEventStatusResolved {
|
||||
summary.ResolvedAlerts++
|
||||
continue
|
||||
}
|
||||
summary.ActiveAlerts++
|
||||
switch event.Severity {
|
||||
case NodeHealthSeverityCritical:
|
||||
summary.CriticalAlerts++
|
||||
case NodeHealthSeverityWarning:
|
||||
summary.WarningAlerts++
|
||||
default:
|
||||
summary.InfoAlerts++
|
||||
}
|
||||
}
|
||||
if snapshot != nil {
|
||||
memoryUsage := percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes)
|
||||
storageUsage := percentage(snapshot.StorageUsedBytes, snapshot.StorageTotalBytes)
|
||||
summary.HasCapacityRisk = snapshot.CPUUsagePercent >= 80 || memoryUsage >= 85 || storageUsage >= 85
|
||||
}
|
||||
if report != nil && report.RequestCount >= 100 {
|
||||
summary.HasTrafficRisk = (float64(report.ErrorCount) / float64(report.RequestCount)) >= 0.05
|
||||
}
|
||||
summary.HasRuntimeRisk = summary.ActiveAlerts > 0 || summary.HasCapacityRisk || summary.HasTrafficRisk
|
||||
return summary
|
||||
}
|
||||
|
||||
func mergeJSONCounts(target distributionAccumulator, raw string) {
|
||||
if len(target) == 0 && strings.TrimSpace(raw) == "" {
|
||||
return
|
||||
}
|
||||
values := parseJSONCounts(raw)
|
||||
for key, value := range values {
|
||||
if strings.TrimSpace(key) == "" || value <= 0 {
|
||||
continue
|
||||
}
|
||||
target[key] += value
|
||||
}
|
||||
}
|
||||
|
||||
func parseJSONCounts(raw string) map[string]int64 {
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return nil
|
||||
}
|
||||
values := make(map[string]int64)
|
||||
if err := json.Unmarshal([]byte(raw), &values); err != nil {
|
||||
return nil
|
||||
}
|
||||
return values
|
||||
}
|
||||
|
||||
func toDistributionItems(values distributionAccumulator, limit int) []DistributionItem {
|
||||
if len(values) == 0 {
|
||||
return []DistributionItem{}
|
||||
}
|
||||
items := make([]DistributionItem, 0, len(values))
|
||||
for key, value := range values {
|
||||
if strings.TrimSpace(key) == "" || value <= 0 {
|
||||
continue
|
||||
}
|
||||
items = append(items, DistributionItem{Key: key, Value: value})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
if items[i].Value == items[j].Value {
|
||||
return items[i].Key < items[j].Key
|
||||
}
|
||||
return items[i].Value > items[j].Value
|
||||
})
|
||||
if limit > 0 && len(items) > limit {
|
||||
items = items[:limit]
|
||||
}
|
||||
return items
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
const observabilityTrendBuckets = 24
|
||||
|
||||
type TrafficTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
}
|
||||
|
||||
type CapacityTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
AverageCPUUsagePercent float64 `json:"average_cpu_usage_percent"`
|
||||
AverageMemoryUsagePercent float64 `json:"average_memory_usage_percent"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
type NetworkTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
NetworkRxBytes int64 `json:"network_rx_bytes"`
|
||||
NetworkTxBytes int64 `json:"network_tx_bytes"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
type DiskIOTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
DiskReadBytes int64 `json:"disk_read_bytes"`
|
||||
DiskWriteBytes int64 `json:"disk_write_bytes"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
type capacityTrendAccumulator struct {
|
||||
cpuSum float64
|
||||
cpuCount int
|
||||
memSum float64
|
||||
memCount int
|
||||
nodes map[string]struct{}
|
||||
}
|
||||
|
||||
type snapshotTrendAccumulator struct {
|
||||
nodes map[string]struct{}
|
||||
}
|
||||
|
||||
func buildTrafficTrendPoints(now time.Time, reports []*model.NodeRequestReport) []TrafficTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]TrafficTrendPoint, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
}
|
||||
|
||||
for _, report := range reports {
|
||||
index, ok := trendBucketIndex(report.WindowEndedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
points[index].RequestCount += report.RequestCount
|
||||
points[index].ErrorCount += report.ErrorCount
|
||||
points[index].UniqueVisitorCount += report.UniqueVisitorCount
|
||||
}
|
||||
|
||||
return points
|
||||
}
|
||||
|
||||
func buildCapacityTrendPoints(now time.Time, snapshots []*model.NodeMetricSnapshot) []CapacityTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]CapacityTrendPoint, observabilityTrendBuckets)
|
||||
accumulators := make([]capacityTrendAccumulator, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
accumulators[index].nodes = make(map[string]struct{})
|
||||
}
|
||||
|
||||
for _, snapshot := range snapshots {
|
||||
index, ok := trendBucketIndex(snapshot.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if snapshot.CPUUsagePercent > 0 {
|
||||
accumulators[index].cpuSum += snapshot.CPUUsagePercent
|
||||
accumulators[index].cpuCount++
|
||||
}
|
||||
if memoryUsage := percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes); memoryUsage > 0 {
|
||||
accumulators[index].memSum += memoryUsage
|
||||
accumulators[index].memCount++
|
||||
}
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
for index := range points {
|
||||
if accumulators[index].cpuCount > 0 {
|
||||
points[index].AverageCPUUsagePercent = accumulators[index].cpuSum / float64(accumulators[index].cpuCount)
|
||||
}
|
||||
if accumulators[index].memCount > 0 {
|
||||
points[index].AverageMemoryUsagePercent = accumulators[index].memSum / float64(accumulators[index].memCount)
|
||||
}
|
||||
points[index].ReportedNodes = len(accumulators[index].nodes)
|
||||
}
|
||||
|
||||
return points
|
||||
}
|
||||
|
||||
func buildNetworkTrendPoints(now time.Time, snapshots []*model.NodeMetricSnapshot, openrestyObs []*model.NodeObservationOpenresty) []NetworkTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]NetworkTrendPoint, observabilityTrendBuckets)
|
||||
accumulators := make([]snapshotTrendAccumulator, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
accumulators[index].nodes = make(map[string]struct{})
|
||||
}
|
||||
|
||||
for _, snapshot := range snapshots {
|
||||
index, ok := trendBucketIndex(snapshot.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
points[index].NetworkRxBytes += snapshot.NetworkRxBytes
|
||||
points[index].NetworkTxBytes += snapshot.NetworkTxBytes
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
for _, obs := range openrestyObs {
|
||||
index, ok := trendBucketIndex(obs.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
points[index].OpenrestyRxBytes += obs.OpenrestyRxBytes
|
||||
points[index].OpenrestyTxBytes += obs.OpenrestyTxBytes
|
||||
if obs.NodeID != "" {
|
||||
accumulators[index].nodes[obs.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
for index := range points {
|
||||
points[index].ReportedNodes = len(accumulators[index].nodes)
|
||||
}
|
||||
|
||||
return points
|
||||
}
|
||||
|
||||
func buildDiskIOTrendPoints(now time.Time, snapshots []*model.NodeMetricSnapshot) []DiskIOTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]DiskIOTrendPoint, observabilityTrendBuckets)
|
||||
accumulators := make([]snapshotTrendAccumulator, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
accumulators[index].nodes = make(map[string]struct{})
|
||||
}
|
||||
|
||||
sort.Slice(snapshots, func(i int, j int) bool {
|
||||
if snapshots[i].CapturedAt.Equal(snapshots[j].CapturedAt) {
|
||||
return snapshots[i].NodeID < snapshots[j].NodeID
|
||||
}
|
||||
return snapshots[i].CapturedAt.Before(snapshots[j].CapturedAt)
|
||||
})
|
||||
|
||||
type diskCounterState struct {
|
||||
read int64
|
||||
write int64
|
||||
seen bool
|
||||
}
|
||||
|
||||
previousByNode := make(map[string]diskCounterState, len(snapshots))
|
||||
|
||||
for _, snapshot := range snapshots {
|
||||
nodeKey := snapshot.NodeID
|
||||
if nodeKey == "" {
|
||||
nodeKey = "__unknown__"
|
||||
}
|
||||
|
||||
previous := previousByNode[nodeKey]
|
||||
previousByNode[nodeKey] = diskCounterState{
|
||||
read: snapshot.DiskReadBytes,
|
||||
write: snapshot.DiskWriteBytes,
|
||||
seen: true,
|
||||
}
|
||||
if !previous.seen {
|
||||
continue
|
||||
}
|
||||
|
||||
index, ok := trendBucketIndex(snapshot.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
readDelta := snapshot.DiskReadBytes - previous.read
|
||||
writeDelta := snapshot.DiskWriteBytes - previous.write
|
||||
if readDelta < 0 {
|
||||
readDelta = 0
|
||||
}
|
||||
if writeDelta < 0 {
|
||||
writeDelta = 0
|
||||
}
|
||||
|
||||
points[index].DiskReadBytes += readDelta
|
||||
points[index].DiskWriteBytes += writeDelta
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
for index := range points {
|
||||
points[index].ReportedNodes = len(accumulators[index].nodes)
|
||||
}
|
||||
|
||||
return points
|
||||
}
|
||||
|
||||
func trendWindowStart(now time.Time) time.Time {
|
||||
return now.Truncate(time.Hour).Add(-(observabilityTrendBuckets - 1) * time.Hour)
|
||||
}
|
||||
|
||||
func trendBucketIndex(timestamp time.Time, start time.Time) (int, bool) {
|
||||
if timestamp.Before(start) {
|
||||
return 0, false
|
||||
}
|
||||
delta := timestamp.Sub(start)
|
||||
index := int(delta / time.Hour)
|
||||
if index < 0 || index >= observabilityTrendBuckets {
|
||||
return 0, false
|
||||
}
|
||||
return index, true
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestBuildDiskIOTrendPointsUsesCounterDelta(t *testing.T) {
|
||||
now := time.Date(2026, 3, 14, 18, 30, 0, 0, time.UTC)
|
||||
start := trendWindowStart(now)
|
||||
|
||||
points := buildDiskIOTrendPoints(now, []*model.NodeMetricSnapshot{
|
||||
{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: start.Add(22 * time.Hour),
|
||||
DiskReadBytes: 100,
|
||||
DiskWriteBytes: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: start.Add(23 * time.Hour),
|
||||
DiskReadBytes: 250,
|
||||
DiskWriteBytes: 260,
|
||||
},
|
||||
})
|
||||
|
||||
last := points[len(points)-1]
|
||||
if last.DiskReadBytes != 150 || last.DiskWriteBytes != 60 {
|
||||
t.Fatalf("expected disk io trend to use counter delta, got %+v", last)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type OriginInput struct {
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
type OriginRouteSummary struct {
|
||||
ID uint `json:"id"`
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
type OriginView struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
Remark string `json:"remark"`
|
||||
RouteCount int64 `json:"route_count"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type OriginDetailView struct {
|
||||
OriginView
|
||||
Routes []OriginRouteSummary `json:"routes"`
|
||||
}
|
||||
|
||||
func ListOrigins() ([]OriginView, error) {
|
||||
origins, err := model.ListOrigins()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildOriginViews(origins)
|
||||
}
|
||||
|
||||
func GetOriginDetail(id uint) (*OriginDetailView, error) {
|
||||
origin, err := model.GetOriginByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views, err := buildOriginViews([]*model.Origin{origin})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
routes, err := model.ListProxyRoutesByOriginID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]OriginRouteSummary, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
items = append(items, OriginRouteSummary{
|
||||
ID: route.ID,
|
||||
Domain: route.Domain,
|
||||
OriginURL: route.OriginURL,
|
||||
Enabled: route.Enabled,
|
||||
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
return items[i].Domain < items[j].Domain
|
||||
})
|
||||
detail := &OriginDetailView{
|
||||
OriginView: views[0],
|
||||
Routes: items,
|
||||
}
|
||||
return detail, nil
|
||||
}
|
||||
|
||||
func CreateOrigin(input OriginInput) (*model.Origin, error) {
|
||||
origin, err := buildOrigin(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = origin.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("源站地址已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return origin, nil
|
||||
}
|
||||
|
||||
func UpdateOrigin(id uint, input OriginInput) (*model.Origin, error) {
|
||||
origin, err := model.GetOriginByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
previousAddress := origin.Address
|
||||
nextOrigin, err := buildOrigin(origin, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Save(nextOrigin).Error; err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return errors.New("源站地址已存在")
|
||||
}
|
||||
return err
|
||||
}
|
||||
if previousAddress == nextOrigin.Address {
|
||||
return nil
|
||||
}
|
||||
return updateRoutesForOriginAddress(tx, nextOrigin.ID, nextOrigin.Address)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nextOrigin, nil
|
||||
}
|
||||
|
||||
func DeleteOrigin(id uint) error {
|
||||
routes, err := model.ListProxyRoutesByOriginID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(routes) > 0 {
|
||||
return errors.New("该源站仍被规则引用,无法删除")
|
||||
}
|
||||
origin, err := model.GetOriginByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return origin.Delete()
|
||||
}
|
||||
|
||||
func buildOrigin(existing *model.Origin, input OriginInput) (*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 getOrCreateOriginByAddress(address string) (*model.Origin, error) {
|
||||
normalizedAddress := normalizeOriginAddress(address)
|
||||
if err := validateOriginAddress(normalizedAddress); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
existing, err := model.GetOriginByAddress(normalizedAddress)
|
||||
if err == nil {
|
||||
return existing, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
origin := &model.Origin{
|
||||
Name: normalizedAddress,
|
||||
Address: normalizedAddress,
|
||||
Remark: "",
|
||||
}
|
||||
if err := origin.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return model.GetOriginByAddress(normalizedAddress)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return origin, nil
|
||||
}
|
||||
|
||||
func updateRoutesForOriginAddress(tx *gorm.DB, originID uint, address string) error {
|
||||
var routes []*model.ProxyRoute
|
||||
if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; 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 := tx.Model(&model.ProxyRoute{}).
|
||||
Where("id = ?", route.ID).
|
||||
Updates(map[string]any{
|
||||
"origin_url": rewrittenOriginURL,
|
||||
"upstreams": string(upstreamsJSON),
|
||||
}).Error; err != nil {
|
||||
return fmt.Errorf("update route %d origin address failed: %w", route.ID, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildOriginViews(origins []*model.Origin) ([]OriginView, error) {
|
||||
countRows, err := model.ListOriginRouteCounts()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
countMap := make(map[uint]int64, len(countRows))
|
||||
for _, row := range countRows {
|
||||
countMap[row.OriginID] = row.RouteCount
|
||||
}
|
||||
views := make([]OriginView, 0, len(origins))
|
||||
for _, origin := range origins {
|
||||
views = append(views, OriginView{
|
||||
ID: origin.ID,
|
||||
Name: origin.Name,
|
||||
Address: origin.Address,
|
||||
Remark: origin.Remark,
|
||||
RouteCount: countMap[origin.ID],
|
||||
CreatedAt: origin.CreatedAt,
|
||||
UpdatedAt: origin.UpdatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
func normalizeOriginAddress(raw string) string {
|
||||
return strings.ToLower(strings.TrimSpace(raw))
|
||||
}
|
||||
|
||||
func validateOriginAddress(address string) error {
|
||||
if address == "" {
|
||||
return errors.New("源站地址不能为空")
|
||||
}
|
||||
if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") {
|
||||
return errors.New("源站地址无需包含 IPv6 方括号")
|
||||
}
|
||||
if ip := net.ParseIP(address); ip != nil {
|
||||
return nil
|
||||
}
|
||||
if len(address) > 253 {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
for _, label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
if label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
for _, r := range label {
|
||||
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
|
||||
continue
|
||||
}
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeOriginName(name string, address string) string {
|
||||
normalized := strings.TrimSpace(name)
|
||||
if normalized != "" {
|
||||
return normalized
|
||||
}
|
||||
return address
|
||||
}
|
||||
|
||||
func normalizeOriginPort(raw string) (string, error) {
|
||||
port := strings.TrimSpace(raw)
|
||||
if port == "" {
|
||||
return "", errors.New("端口不能为空")
|
||||
}
|
||||
value, err := strconv.Atoi(port)
|
||||
if err != nil || value < 1 || value > 65535 {
|
||||
return "", errors.New("端口格式不合法")
|
||||
}
|
||||
return strconv.Itoa(value), nil
|
||||
}
|
||||
|
||||
func normalizeOriginScheme(raw string) (string, error) {
|
||||
scheme := strings.ToLower(strings.TrimSpace(raw))
|
||||
switch scheme {
|
||||
case "http", "https":
|
||||
return scheme, nil
|
||||
default:
|
||||
return "", errors.New("源站协议仅支持 http 或 https")
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeOriginURI(raw string) (string, error) {
|
||||
uri := strings.TrimSpace(raw)
|
||||
if uri == "" {
|
||||
return "", nil
|
||||
}
|
||||
if strings.Contains(uri, "://") {
|
||||
return "", errors.New("源站路径不能包含协议")
|
||||
}
|
||||
if !strings.HasPrefix(uri, "/") && !strings.HasPrefix(uri, "?") {
|
||||
return "", errors.New("源站路径需以 / 或 ? 开头")
|
||||
}
|
||||
return uri, nil
|
||||
}
|
||||
|
||||
func formatOriginHost(address string, port string) string {
|
||||
if ip := net.ParseIP(address); ip != nil && strings.Contains(address, ":") {
|
||||
return net.JoinHostPort(address, port)
|
||||
}
|
||||
return net.JoinHostPort(address, port)
|
||||
}
|
||||
|
||||
func buildOriginURLFromParts(
|
||||
scheme string,
|
||||
address string,
|
||||
port string,
|
||||
uri string,
|
||||
) (string, error) {
|
||||
normalizedScheme, err := normalizeOriginScheme(scheme)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
normalizedAddress := normalizeOriginAddress(address)
|
||||
if err := validateOriginAddress(normalizedAddress); err != nil {
|
||||
return "", err
|
||||
}
|
||||
normalizedPort, err := normalizeOriginPort(port)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
normalizedURI, err := normalizeOriginURI(uri)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
parsed := &url.URL{
|
||||
Scheme: normalizedScheme,
|
||||
Host: formatOriginHost(normalizedAddress, normalizedPort),
|
||||
}
|
||||
if normalizedURI != "" {
|
||||
if strings.HasPrefix(normalizedURI, "?") {
|
||||
parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?")
|
||||
} else {
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", 2)
|
||||
parsed.Path = pathQuery[0]
|
||||
if len(pathQuery) > 1 {
|
||||
parsed.RawQuery = pathQuery[1]
|
||||
}
|
||||
}
|
||||
}
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
func extractOriginAddress(rawURL string) (string, error) {
|
||||
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("源站地址格式不合法: %w", err)
|
||||
}
|
||||
address := normalizeOriginAddress(parsed.Hostname())
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return address, nil
|
||||
}
|
||||
|
||||
func rewriteOriginURLAddress(rawURL string, newAddress string) (string, error) {
|
||||
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("源站地址格式不合法: %w", err)
|
||||
}
|
||||
address := normalizeOriginAddress(newAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", err
|
||||
}
|
||||
port := parsed.Port()
|
||||
if port == "" {
|
||||
return "", errors.New("源站地址缺少端口")
|
||||
}
|
||||
parsed.Host = formatOriginHost(address, port)
|
||||
return parsed.String(), nil
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestCreateProxyRouteStructuredOriginAutoCreatesOrigin(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginScheme: "https",
|
||||
OriginAddress: "origin.internal",
|
||||
OriginPort: "8443",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if route.OriginID == nil || *route.OriginID == 0 {
|
||||
t.Fatal("expected route to be linked with an auto-created origin")
|
||||
}
|
||||
if route.OriginURL != "https://origin.internal:8443" {
|
||||
t.Fatalf("unexpected route origin url: %s", route.OriginURL)
|
||||
}
|
||||
|
||||
origin, err := model.GetOriginByID(*route.OriginID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetOriginByID failed: %v", err)
|
||||
}
|
||||
if origin.Address != "origin.internal" {
|
||||
t.Fatalf("unexpected origin address: %s", origin.Address)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateOriginRewritesLinkedRouteOriginURL(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
origin, err := CreateOrigin(OriginInput{
|
||||
Name: "primary-origin",
|
||||
Address: "origin-a.internal",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateOrigin failed: %v", err)
|
||||
}
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginID: &origin.ID,
|
||||
OriginScheme: "https",
|
||||
OriginPort: "8443",
|
||||
OriginURI: "/api",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
updatedOrigin, err := UpdateOrigin(origin.ID, OriginInput{
|
||||
Name: origin.Name,
|
||||
Address: "origin-c.internal",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateOrigin failed: %v", err)
|
||||
}
|
||||
if updatedOrigin.Address != "origin-c.internal" {
|
||||
t.Fatalf("unexpected updated origin address: %s", updatedOrigin.Address)
|
||||
}
|
||||
|
||||
reloadedRoute, err := model.GetProxyRouteByID(route.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetProxyRouteByID failed: %v", err)
|
||||
}
|
||||
if reloadedRoute.OriginURL != "https://origin-c.internal:8443/api" {
|
||||
t.Fatalf("expected route origin url to be rewritten, got %s", reloadedRoute.OriginURL)
|
||||
}
|
||||
if reloadedRoute.Upstreams == "" || reloadedRoute.Upstreams == "[]" {
|
||||
t.Fatalf("expected route upstreams to be preserved, got %s", reloadedRoute.Upstreams)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteOriginRejectsReferencedOrigin(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
origin, err := CreateOrigin(OriginInput{
|
||||
Address: "origin-a.internal",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateOrigin failed: %v", err)
|
||||
}
|
||||
if _, err = CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginID: &origin.ID,
|
||||
OriginScheme: "https",
|
||||
OriginPort: "443",
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
if err = DeleteOrigin(origin.ID); err == nil {
|
||||
t.Fatal("expected referenced origin deletion to fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,834 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/url"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
pagesMaxDeploymentFiles = 1000
|
||||
pagesMaxDeploymentBytes = 100 * 1024 * 1024
|
||||
defaultPagesEntryFile = "index.html"
|
||||
defaultPagesFallbackPath = "/index.html"
|
||||
)
|
||||
|
||||
var pagesSlugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,126}[a-z0-9]$|^[a-z0-9]$`)
|
||||
|
||||
type PagesProjectInput struct {
|
||||
Name string `json:"name"`
|
||||
Slug string `json:"slug"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
|
||||
SPAFallbackPath string `json:"spa_fallback_path"`
|
||||
APIProxyEnabled bool `json:"api_proxy_enabled"`
|
||||
APIProxyPath string `json:"api_proxy_path"`
|
||||
APIProxyPass string `json:"api_proxy_pass"`
|
||||
APIProxyRewrite string `json:"api_proxy_rewrite"`
|
||||
RootDir string `json:"root_dir"`
|
||||
EntryFile string `json:"entry_file"`
|
||||
}
|
||||
|
||||
type PagesProjectView struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Slug string `json:"slug"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
|
||||
SPAFallbackPath string `json:"spa_fallback_path"`
|
||||
APIProxyEnabled bool `json:"api_proxy_enabled"`
|
||||
APIProxyPath string `json:"api_proxy_path"`
|
||||
APIProxyPass string `json:"api_proxy_pass"`
|
||||
APIProxyRewrite string `json:"api_proxy_rewrite"`
|
||||
RootDir string `json:"root_dir"`
|
||||
EntryFile string `json:"entry_file"`
|
||||
ActiveDeploymentID *uint `json:"active_deployment_id"`
|
||||
ActiveDeployment *PagesDeploymentView `json:"active_deployment,omitempty"`
|
||||
DeploymentCount int64 `json:"deployment_count"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type PagesDeploymentView struct {
|
||||
ID uint `json:"id"`
|
||||
ProjectID uint `json:"project_id"`
|
||||
DeploymentNumber int `json:"deployment_number"`
|
||||
Checksum string `json:"checksum"`
|
||||
Status string `json:"status"`
|
||||
FileCount int `json:"file_count"`
|
||||
TotalSize int64 `json:"total_size"`
|
||||
CreatedBy string `json:"created_by"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
ActivatedAt *time.Time `json:"activated_at"`
|
||||
}
|
||||
|
||||
type PagesDeploymentFileView struct {
|
||||
ID uint `json:"id"`
|
||||
DeploymentID uint `json:"deployment_id"`
|
||||
Path string `json:"path"`
|
||||
Size int64 `json:"size"`
|
||||
Checksum string `json:"checksum"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type pagesDeploymentManifest struct {
|
||||
Files []model.PagesDeploymentFile
|
||||
FileCount int
|
||||
TotalSize int64
|
||||
EntryFile string
|
||||
}
|
||||
|
||||
func ListPagesProjects() ([]*PagesProjectView, error) {
|
||||
projects, err := model.ListPagesProjects()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]*PagesProjectView, 0, len(projects))
|
||||
for _, project := range projects {
|
||||
view, err := buildPagesProjectView(project)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func GetPagesProject(id uint) (*PagesProjectView, error) {
|
||||
project, err := model.GetPagesProjectByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildPagesProjectView(project)
|
||||
}
|
||||
|
||||
func CreatePagesProject(input PagesProjectInput) (*PagesProjectView, error) {
|
||||
project, err := buildPagesProject(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.DB.Create(project).Error; err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("Pages 项目标识已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildPagesProjectView(project)
|
||||
}
|
||||
|
||||
func UpdatePagesProject(id uint, input PagesProjectInput) (*PagesProjectView, error) {
|
||||
project, err := model.GetPagesProjectByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
project, err = buildPagesProject(project, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.DB.Model(project).Updates(map[string]any{
|
||||
"name": project.Name,
|
||||
"slug": project.Slug,
|
||||
"description": project.Description,
|
||||
"enabled": project.Enabled,
|
||||
"spa_fallback_enabled": project.SPAFallbackEnabled,
|
||||
"spa_fallback_path": project.SPAFallbackPath,
|
||||
"api_proxy_enabled": project.APIProxyEnabled,
|
||||
"api_proxy_path": project.APIProxyPath,
|
||||
"api_proxy_pass": project.APIProxyPass,
|
||||
"api_proxy_rewrite": project.APIProxyRewrite,
|
||||
"root_dir": project.RootDir,
|
||||
"entry_file": project.EntryFile,
|
||||
}).Error; err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("Pages 项目标识已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildPagesProjectView(project)
|
||||
}
|
||||
|
||||
func DeletePagesProject(id uint) error {
|
||||
project, err := model.GetPagesProjectByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var routeCount int64
|
||||
if err = model.DB.Model(&model.ProxyRoute{}).Where("pages_project_id = ?", project.ID).Count(&routeCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if routeCount > 0 {
|
||||
return errors.New("Pages 项目已被规则引用,不能删除")
|
||||
}
|
||||
deployments, err := model.ListPagesDeployments(project.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("deployment_id IN (?)", tx.Model(&model.PagesDeployment{}).Select("id").Where("project_id = ?", project.ID)).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("project_id = ?", project.ID).Delete(&model.PagesDeployment{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Delete(project).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, deployment := range deployments {
|
||||
_ = os.Remove(deployment.ArtifactPath)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func ListPagesProjectDeployments(projectID uint) ([]*PagesDeploymentView, error) {
|
||||
if _, err := model.GetPagesProjectByID(projectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deployments, err := model.ListPagesDeployments(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]*PagesDeploymentView, 0, len(deployments))
|
||||
for _, deployment := range deployments {
|
||||
views = append(views, buildPagesDeploymentView(deployment))
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func ListPagesDeploymentFiles(deploymentID uint) ([]*PagesDeploymentFileView, error) {
|
||||
if _, err := model.GetPagesDeploymentByID(deploymentID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
files, err := model.ListPagesDeploymentFiles(deploymentID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]*PagesDeploymentFileView, 0, len(files))
|
||||
for _, file := range files {
|
||||
views = append(views, &PagesDeploymentFileView{
|
||||
ID: file.ID,
|
||||
DeploymentID: file.DeploymentID,
|
||||
Path: file.Path,
|
||||
Size: file.Size,
|
||||
Checksum: file.Checksum,
|
||||
CreatedAt: file.CreatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func UploadPagesDeployment(projectID uint, fileHeader *multipart.FileHeader, rootDir string, entryFile string, createdBy string) (*PagesDeploymentView, error) {
|
||||
project, err := model.GetPagesProjectByID(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if fileHeader == nil {
|
||||
return nil, errors.New("缺少 Pages 部署包")
|
||||
}
|
||||
if !strings.EqualFold(filepath.Ext(fileHeader.Filename), ".zip") {
|
||||
return nil, errors.New("Pages 部署包必须是 .zip 文件")
|
||||
}
|
||||
rootDir, err = validateAndNormalizePagesRootDir(project.RootDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
entryFile = normalizePagesEntryFile(project.EntryFile)
|
||||
tempPath, checksum, err := persistPagesUploadTemp(fileHeader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer os.Remove(tempPath)
|
||||
manifest, err := inspectPagesZip(tempPath, rootDir, entryFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
artifactPath, err := pagesArtifactPath(project.Slug, checksum)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = os.MkdirAll(filepath.Dir(artifactPath), 0o755); err != nil {
|
||||
return nil, fmt.Errorf("创建 Pages 存储目录失败: %w", err)
|
||||
}
|
||||
if err = copyFile(tempPath, artifactPath); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deployment := &model.PagesDeployment{}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var maxNumber int
|
||||
if err := tx.Model(&model.PagesDeployment{}).
|
||||
Where("project_id = ?", project.ID).
|
||||
Select("COALESCE(MAX(deployment_number), 0)").
|
||||
Scan(&maxNumber).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
deployment = &model.PagesDeployment{
|
||||
ProjectID: project.ID,
|
||||
DeploymentNumber: maxNumber + 1,
|
||||
Checksum: checksum,
|
||||
Status: model.PagesDeploymentStatusUploaded,
|
||||
ArtifactPath: artifactPath,
|
||||
FileCount: manifest.FileCount,
|
||||
TotalSize: manifest.TotalSize,
|
||||
CreatedBy: strings.TrimSpace(createdBy),
|
||||
}
|
||||
if err := tx.Create(deployment).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for index := range manifest.Files {
|
||||
manifest.Files[index].DeploymentID = deployment.ID
|
||||
}
|
||||
if len(manifest.Files) > 0 {
|
||||
if err := tx.Create(&manifest.Files).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
_ = os.Remove(artifactPath)
|
||||
return nil, err
|
||||
}
|
||||
return buildPagesDeploymentView(deployment), nil
|
||||
}
|
||||
|
||||
func validateAndNormalizePagesRootDir(raw string) (string, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
return "", nil
|
||||
}
|
||||
if len(value) > 512 {
|
||||
return "", errors.New("Pages 根目录长度不能超过 512")
|
||||
}
|
||||
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
|
||||
return "", errors.New("Pages 根目录包含不支持的字符")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r <= 0x20 || r == 0x7f {
|
||||
return "", errors.New("Pages 根目录不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(filepath.ToSlash(value))
|
||||
if cleaned == "." || cleaned == "/" {
|
||||
return "", nil
|
||||
}
|
||||
for _, segment := range strings.Split(cleaned, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("Pages 根目录不能包含 . 或 .. 路径段")
|
||||
}
|
||||
}
|
||||
return strings.TrimPrefix(cleaned, "/"), nil
|
||||
}
|
||||
|
||||
func ActivatePagesDeployment(projectID uint, deploymentID uint) (*PagesProjectView, error) {
|
||||
project, err := model.GetPagesProjectByID(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deployment, err := model.GetPagesDeploymentByID(deploymentID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if deployment.ProjectID != project.ID {
|
||||
return nil, errors.New("Pages 部署不属于该项目")
|
||||
}
|
||||
now := time.Now()
|
||||
if err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&model.PagesDeployment{}).
|
||||
Where("project_id = ?", project.ID).
|
||||
Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(deployment).Updates(map[string]any{
|
||||
"status": model.PagesDeploymentStatusActive,
|
||||
"activated_at": &now,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(project).Updates(map[string]any{
|
||||
"active_deployment_id": deployment.ID,
|
||||
}).Error
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetPagesProject(project.ID)
|
||||
}
|
||||
|
||||
func DeletePagesDeployment(projectID uint, deploymentID uint) error {
|
||||
project, err := model.GetPagesProjectByID(projectID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
deployment, err := model.GetPagesDeploymentByID(deploymentID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deployment.ProjectID != project.ID {
|
||||
return errors.New("Pages 部署不属于该项目")
|
||||
}
|
||||
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID == deployment.ID {
|
||||
return errors.New("不能删除当前激活的 Pages 部署")
|
||||
}
|
||||
return model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("deployment_id = ?", deployment.ID).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Delete(deployment).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
_ = os.Remove(deployment.ArtifactPath)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func GetPagesDeploymentPackagePath(deploymentID uint) (string, string, error) {
|
||||
deployment, err := model.GetPagesDeploymentByID(deploymentID)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
if err = ensurePagesDeploymentInActiveSnapshot(deployment.ID); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
if strings.TrimSpace(deployment.ArtifactPath) == "" {
|
||||
return "", "", errors.New("Pages 部署包路径为空")
|
||||
}
|
||||
if _, err = os.Stat(deployment.ArtifactPath); err != nil {
|
||||
return "", "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
}
|
||||
return deployment.ArtifactPath, fmt.Sprintf("pages-deployment-%d.zip", deployment.ID), nil
|
||||
}
|
||||
|
||||
func ensurePagesDeploymentInActiveSnapshot(deploymentID uint) error {
|
||||
version, err := model.GetActiveConfigVersion()
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return errors.New("Pages 部署尚未进入激活配置")
|
||||
}
|
||||
return err
|
||||
}
|
||||
snapshot, err := parseSnapshotDocument(version.SnapshotJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, route := range snapshot.Routes {
|
||||
if route.UpstreamType != "pages" || route.PagesDeployment == nil {
|
||||
continue
|
||||
}
|
||||
if route.PagesDeployment.DeploymentID == deploymentID {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return errors.New("Pages 部署尚未进入激活配置")
|
||||
}
|
||||
|
||||
func buildPagesProject(project *model.PagesProject, input PagesProjectInput) (*model.PagesProject, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, errors.New("Pages 项目名称不能为空")
|
||||
}
|
||||
slug := normalizePagesSlug(input.Slug)
|
||||
if slug == "" {
|
||||
slug = normalizePagesSlug(name)
|
||||
}
|
||||
if !pagesSlugPattern.MatchString(slug) {
|
||||
return nil, errors.New("Pages 项目标识只能包含小写字母、数字和连字符")
|
||||
}
|
||||
if project == nil {
|
||||
project = &model.PagesProject{}
|
||||
}
|
||||
project.Name = name
|
||||
project.Slug = slug
|
||||
project.Description = strings.TrimSpace(input.Description)
|
||||
project.Enabled = input.Enabled
|
||||
project.SPAFallbackEnabled = input.SPAFallbackEnabled
|
||||
fallbackPath, err := normalizePagesFallbackPath(input.SPAFallbackPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
project.SPAFallbackPath = fallbackPath
|
||||
|
||||
project.APIProxyEnabled = input.APIProxyEnabled
|
||||
apiProxyPath := strings.TrimSpace(input.APIProxyPath)
|
||||
apiProxyPass := strings.TrimSpace(input.APIProxyPass)
|
||||
apiProxyRewrite := strings.TrimSpace(input.APIProxyRewrite)
|
||||
|
||||
if project.APIProxyEnabled {
|
||||
if apiProxyPath == "" {
|
||||
return nil, errors.New("启用 API 反代时,匹配路径不能为空")
|
||||
}
|
||||
if !strings.HasPrefix(apiProxyPath, "/") {
|
||||
return nil, errors.New("API 反代匹配路径必须以 '/' 开头")
|
||||
}
|
||||
if apiProxyPass == "" {
|
||||
return nil, errors.New("启用 API 反代时,后端服务地址不能为空")
|
||||
}
|
||||
parsedURL, err := url.Parse(apiProxyPass)
|
||||
if err != nil || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") || parsedURL.Host == "" {
|
||||
return nil, errors.New("API 反代后端服务地址必须是有效的 HTTP/HTTPS URL")
|
||||
}
|
||||
}
|
||||
project.APIProxyPath = apiProxyPath
|
||||
project.APIProxyPass = apiProxyPass
|
||||
project.APIProxyRewrite = apiProxyRewrite
|
||||
|
||||
rootDir, err := validateAndNormalizePagesRootDir(input.RootDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
project.RootDir = rootDir
|
||||
project.EntryFile = normalizePagesEntryFile(input.EntryFile)
|
||||
|
||||
return project, nil
|
||||
}
|
||||
|
||||
func buildPagesProjectView(project *model.PagesProject) (*PagesProjectView, error) {
|
||||
if project == nil {
|
||||
return nil, errors.New("Pages 项目为空")
|
||||
}
|
||||
view := &PagesProjectView{
|
||||
ID: project.ID,
|
||||
Name: project.Name,
|
||||
Slug: project.Slug,
|
||||
Description: project.Description,
|
||||
Enabled: project.Enabled,
|
||||
SPAFallbackEnabled: project.SPAFallbackEnabled,
|
||||
SPAFallbackPath: normalizeStoredPagesFallbackPath(project.SPAFallbackPath),
|
||||
APIProxyEnabled: project.APIProxyEnabled,
|
||||
APIProxyPath: project.APIProxyPath,
|
||||
APIProxyPass: project.APIProxyPass,
|
||||
APIProxyRewrite: project.APIProxyRewrite,
|
||||
RootDir: project.RootDir,
|
||||
EntryFile: project.EntryFile,
|
||||
ActiveDeploymentID: project.ActiveDeploymentID,
|
||||
CreatedAt: project.CreatedAt,
|
||||
UpdatedAt: project.UpdatedAt,
|
||||
}
|
||||
if err := model.DB.Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&view.DeploymentCount).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID != 0 {
|
||||
deployment, err := model.GetPagesDeploymentByID(*project.ActiveDeploymentID)
|
||||
if err == nil {
|
||||
view.ActiveDeployment = buildPagesDeploymentView(deployment)
|
||||
}
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func buildPagesDeploymentView(deployment *model.PagesDeployment) *PagesDeploymentView {
|
||||
if deployment == nil {
|
||||
return nil
|
||||
}
|
||||
return &PagesDeploymentView{
|
||||
ID: deployment.ID,
|
||||
ProjectID: deployment.ProjectID,
|
||||
DeploymentNumber: deployment.DeploymentNumber,
|
||||
Checksum: deployment.Checksum,
|
||||
Status: deployment.Status,
|
||||
FileCount: deployment.FileCount,
|
||||
TotalSize: deployment.TotalSize,
|
||||
CreatedBy: deployment.CreatedBy,
|
||||
CreatedAt: deployment.CreatedAt,
|
||||
ActivatedAt: deployment.ActivatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizePagesSlug(raw string) string {
|
||||
value := strings.ToLower(strings.TrimSpace(raw))
|
||||
var builder strings.Builder
|
||||
lastDash := false
|
||||
for _, r := range value {
|
||||
valid := (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9')
|
||||
if valid {
|
||||
builder.WriteRune(r)
|
||||
lastDash = false
|
||||
continue
|
||||
}
|
||||
if !lastDash {
|
||||
builder.WriteByte('-')
|
||||
lastDash = true
|
||||
}
|
||||
}
|
||||
return strings.Trim(builder.String(), "-")
|
||||
}
|
||||
|
||||
func normalizePagesFallbackPath(raw string) (string, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
value = defaultPagesFallbackPath
|
||||
}
|
||||
if len(value) > 512 {
|
||||
return "", errors.New("SPA fallback 回退路径长度不能超过 512")
|
||||
}
|
||||
if !strings.HasPrefix(value, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须以 / 开头")
|
||||
}
|
||||
if value == "/" || strings.HasSuffix(value, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
|
||||
}
|
||||
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
|
||||
return "", errors.New("SPA fallback 回退路径包含不支持的字符")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r <= 0x20 || r == 0x7f {
|
||||
return "", errors.New("SPA fallback 回退路径不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
for _, segment := range strings.Split(value, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("SPA fallback 回退路径不能包含 . 或 .. 路径段")
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(value)
|
||||
if cleaned == "." || !strings.HasPrefix(cleaned, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径不合法")
|
||||
}
|
||||
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
|
||||
}
|
||||
return cleaned, nil
|
||||
}
|
||||
|
||||
func normalizeStoredPagesFallbackPath(value string) string {
|
||||
normalized, err := normalizePagesFallbackPath(value)
|
||||
if err != nil {
|
||||
return defaultPagesFallbackPath
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func normalizePagesEntryFile(raw string) string {
|
||||
value := path.Clean(strings.TrimSpace(filepath.ToSlash(raw)))
|
||||
if value == "." || value == "/" {
|
||||
return defaultPagesEntryFile
|
||||
}
|
||||
return strings.TrimPrefix(value, "/")
|
||||
}
|
||||
|
||||
func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, error) {
|
||||
file, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
defer file.Close()
|
||||
temp, err := os.CreateTemp("", "openflare-pages-*.zip")
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
defer temp.Close()
|
||||
hash := sha256.New()
|
||||
limited := io.LimitReader(file, pagesMaxDeploymentBytes+1)
|
||||
written, err := io.Copy(io.MultiWriter(temp, hash), limited)
|
||||
if err != nil {
|
||||
_ = os.Remove(temp.Name())
|
||||
return "", "", err
|
||||
}
|
||||
if written > pagesMaxDeploymentBytes {
|
||||
_ = os.Remove(temp.Name())
|
||||
return "", "", fmt.Errorf("Pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
|
||||
}
|
||||
return temp.Name(), hex.EncodeToString(hash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
func findCommonRootPrefix(files []*zip.File) (string, error) {
|
||||
var firstFilePath string
|
||||
hasMultipleFiles := false
|
||||
for _, item := range files {
|
||||
normalizedPath, skip, err := normalizePagesZipPath(item.Name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
if firstFilePath == "" {
|
||||
firstFilePath = normalizedPath
|
||||
} else {
|
||||
hasMultipleFiles = true
|
||||
}
|
||||
}
|
||||
if firstFilePath == "" {
|
||||
return "", nil
|
||||
}
|
||||
parts := strings.Split(firstFilePath, "/")
|
||||
if len(parts) <= 1 {
|
||||
return "", nil
|
||||
}
|
||||
commonPrefix := parts[0] + "/"
|
||||
if hasMultipleFiles {
|
||||
for _, item := range files {
|
||||
normalizedPath, skip, err := normalizePagesZipPath(item.Name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
if !strings.HasPrefix(normalizedPath, commonPrefix) {
|
||||
return "", nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return commonPrefix, nil
|
||||
}
|
||||
|
||||
func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*pagesDeploymentManifest, error) {
|
||||
reader, err := zip.OpenReader(zipPath)
|
||||
if err != nil {
|
||||
return nil, errors.New("Pages 部署包不是有效 zip 文件")
|
||||
}
|
||||
defer reader.Close()
|
||||
|
||||
commonPrefix, err := findCommonRootPrefix(reader.File)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
manifest := &pagesDeploymentManifest{
|
||||
Files: []model.PagesDeploymentFile{},
|
||||
EntryFile: entryFile,
|
||||
}
|
||||
targetEntryPath := entryFile
|
||||
if rootDir != "" {
|
||||
targetEntryPath = path.Join(rootDir, entryFile)
|
||||
}
|
||||
entrySeen := false
|
||||
for _, item := range reader.File {
|
||||
normalizedPath, skip, err := normalizePagesZipPath(item.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
|
||||
if commonPrefix != "" {
|
||||
normalizedPath = strings.TrimPrefix(normalizedPath, commonPrefix)
|
||||
}
|
||||
|
||||
if item.FileInfo().Mode()&os.ModeSymlink != 0 {
|
||||
return nil, fmt.Errorf("Pages 部署包不支持符号链接: %s", normalizedPath)
|
||||
}
|
||||
if item.UncompressedSize64 > pagesMaxDeploymentBytes {
|
||||
return nil, fmt.Errorf("Pages 文件过大: %s", normalizedPath)
|
||||
}
|
||||
manifest.FileCount++
|
||||
if manifest.FileCount > pagesMaxDeploymentFiles {
|
||||
return nil, fmt.Errorf("Pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles)
|
||||
}
|
||||
manifest.TotalSize += int64(item.UncompressedSize64)
|
||||
if manifest.TotalSize > pagesMaxDeploymentBytes {
|
||||
return nil, fmt.Errorf("Pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
|
||||
}
|
||||
checksum, err := checksumZipFile(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if normalizedPath == targetEntryPath {
|
||||
entrySeen = true
|
||||
}
|
||||
manifest.Files = append(manifest.Files, model.PagesDeploymentFile{
|
||||
Path: normalizedPath,
|
||||
Size: int64(item.UncompressedSize64),
|
||||
Checksum: checksum,
|
||||
})
|
||||
}
|
||||
if manifest.FileCount == 0 {
|
||||
return nil, errors.New("Pages 部署包不能为空")
|
||||
}
|
||||
if !entrySeen {
|
||||
return nil, fmt.Errorf("Pages 部署包缺少入口文件 %s", targetEntryPath)
|
||||
}
|
||||
return manifest, nil
|
||||
}
|
||||
|
||||
func normalizePagesZipPath(raw string) (string, bool, error) {
|
||||
name := strings.TrimSpace(filepath.ToSlash(raw))
|
||||
if name == "" {
|
||||
return "", true, nil
|
||||
}
|
||||
if strings.HasSuffix(name, "/") {
|
||||
return "", true, nil
|
||||
}
|
||||
if strings.HasPrefix(name, "/") || path.IsAbs(name) {
|
||||
return "", false, fmt.Errorf("Pages 部署包不能包含绝对路径: %s", raw)
|
||||
}
|
||||
cleaned := path.Clean(name)
|
||||
if cleaned == "." {
|
||||
return "", true, nil
|
||||
}
|
||||
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") {
|
||||
return "", false, fmt.Errorf("Pages 部署包路径不能逃逸目录: %s", raw)
|
||||
}
|
||||
return cleaned, false, nil
|
||||
}
|
||||
|
||||
func checksumZipFile(item *zip.File) (string, error) {
|
||||
file, err := item.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer file.Close()
|
||||
hash := sha256.New()
|
||||
if _, err = io.Copy(hash, file); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(hash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
func pagesArtifactPath(projectSlug string, checksum string) (string, error) {
|
||||
root, err := pagesStorageRoot()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(root, "artifacts", projectSlug, checksum+".zip"), nil
|
||||
}
|
||||
|
||||
func pagesStorageRoot() (string, error) {
|
||||
if common.SQLDSN != "" {
|
||||
return filepath.Abs(filepath.Join("data", "pages"))
|
||||
}
|
||||
dbPath := strings.TrimSpace(common.SQLitePath)
|
||||
if dbPath == "" {
|
||||
return filepath.Abs(filepath.Join("data", "pages"))
|
||||
}
|
||||
dir := filepath.Dir(dbPath)
|
||||
if dir == "." || dir == "" {
|
||||
dir = "data"
|
||||
}
|
||||
return filepath.Abs(filepath.Join(dir, "pages"))
|
||||
}
|
||||
|
||||
func copyFile(src string, dst string) error {
|
||||
input, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer input.Close()
|
||||
output, err := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o644)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer output.Close()
|
||||
if _, err = io.Copy(output, input); err != nil {
|
||||
return err
|
||||
}
|
||||
return output.Sync()
|
||||
}
|
||||
@@ -0,0 +1,420 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestPagesUploadActivateAndPublishStaticRoute(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "Marketing Site",
|
||||
Slug: "marketing-site",
|
||||
Enabled: true,
|
||||
SPAFallbackEnabled: true,
|
||||
SPAFallbackPath: "/app.html",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
uploadHeader := multipartFileHeader(t, "site.zip", testPagesZip(t, map[string]string{
|
||||
"index.html": "<h1>Hello Pages</h1>",
|
||||
"assets/app.js": "console.log('pages')",
|
||||
"assets/style.css": "body{color:#111}",
|
||||
}))
|
||||
deployment, err := UploadPagesDeployment(project.ID, uploadHeader, "", "index.html", "root")
|
||||
if err != nil {
|
||||
t.Fatalf("UploadPagesDeployment failed: %v", err)
|
||||
}
|
||||
if deployment.FileCount != 3 || deployment.TotalSize == 0 {
|
||||
t.Fatalf("unexpected deployment manifest: %+v", deployment)
|
||||
}
|
||||
project, err = ActivatePagesDeployment(project.ID, deployment.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("ActivatePagesDeployment failed: %v", err)
|
||||
}
|
||||
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID != deployment.ID {
|
||||
t.Fatalf("expected active deployment %d, got %+v", deployment.ID, project.ActiveDeploymentID)
|
||||
}
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "pages.example.com",
|
||||
Enabled: true,
|
||||
UpstreamType: "pages",
|
||||
PagesProjectID: &project.ID,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if route.UpstreamType != "pages" || route.PagesProjectID == nil || *route.PagesProjectID != project.ID {
|
||||
t.Fatalf("expected route to bind Pages project, got %+v", route)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root", false)
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"upstream_type":"pages"`) {
|
||||
t.Fatalf("expected snapshot to include pages route, got %s", result.Version.SnapshotJSON)
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"deployment_id":`) {
|
||||
t.Fatalf("expected snapshot to include pages deployment, got %s", result.Version.SnapshotJSON)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "root \"__OPENFLARE_PAGES_DIR__/deployments/") {
|
||||
t.Fatalf("expected rendered config to use pages dir placeholder, got:\n%s", result.Version.RenderedConfig)
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"spa_fallback_path":"/app.html"`) {
|
||||
t.Fatalf("expected snapshot to include custom SPA fallback path, got %s", result.Version.SnapshotJSON)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "try_files $uri $uri/ /app.html;") {
|
||||
t.Fatalf("expected SPA fallback try_files, got:\n%s", result.Version.RenderedConfig)
|
||||
}
|
||||
if strings.Contains(result.Version.RenderedConfig, "proxy_pass") {
|
||||
t.Fatalf("Pages route must not render proxy_pass, got:\n%s", result.Version.RenderedConfig)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPagesProjectRejectsUnsafeFallbackPath(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "Unsafe Fallback",
|
||||
Slug: "unsafe-fallback",
|
||||
Enabled: true,
|
||||
SPAFallbackEnabled: true,
|
||||
SPAFallbackPath: "/index.html; proxy_pass http://evil",
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "回退路径") {
|
||||
t.Fatalf("expected unsafe SPA fallback path rejection, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadPagesDeploymentRejectsZipSlip(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "Unsafe Site",
|
||||
Slug: "unsafe-site",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
_, err = UploadPagesDeployment(project.ID, multipartFileHeader(t, "bad.zip", testPagesZip(t, map[string]string{
|
||||
"../escape.html": "bad",
|
||||
"index.html": "ok",
|
||||
})), "", "index.html", "root")
|
||||
if err == nil || !strings.Contains(err.Error(), "逃逸目录") {
|
||||
t.Fatalf("expected zip-slip rejection, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPagesRouteRequiresActiveDeployment(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "Draft Site",
|
||||
Slug: "draft-site",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
if _, err = CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "draft.example.com",
|
||||
Enabled: true,
|
||||
UpstreamType: "pages",
|
||||
PagesProjectID: &project.ID,
|
||||
}); err == nil || !strings.Contains(err.Error(), "没有激活部署") {
|
||||
t.Fatalf("expected active deployment validation, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPagesDeploymentPackageRequiresActiveConfigSnapshot(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{Name: "Published Site", Slug: "published-site", Enabled: true})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
deployment, err := UploadPagesDeployment(project.ID, multipartFileHeader(t, "site.zip", testPagesZip(t, map[string]string{
|
||||
"index.html": "ok",
|
||||
})), "", "index.html", "root")
|
||||
if err != nil {
|
||||
t.Fatalf("UploadPagesDeployment failed: %v", err)
|
||||
}
|
||||
if _, err = ActivatePagesDeployment(project.ID, deployment.ID); err != nil {
|
||||
t.Fatalf("ActivatePagesDeployment failed: %v", err)
|
||||
}
|
||||
if _, _, err = GetPagesDeploymentPackagePath(deployment.ID); err == nil || !strings.Contains(err.Error(), "激活配置") {
|
||||
t.Fatalf("expected package download to require active config, got %v", err)
|
||||
}
|
||||
if _, err = CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "published.example.com",
|
||||
Enabled: true,
|
||||
UpstreamType: "pages",
|
||||
PagesProjectID: &project.ID,
|
||||
}); err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if _, err = PublishConfigVersion("root", false); err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
filePath, fileName, err := GetPagesDeploymentPackagePath(deployment.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetPagesDeploymentPackagePath failed after publish: %v", err)
|
||||
}
|
||||
if filePath == "" || fileName == "" {
|
||||
t.Fatalf("expected package path and file name, got path=%q name=%q", filePath, fileName)
|
||||
}
|
||||
}
|
||||
|
||||
func testPagesZip(t *testing.T, files map[string]string) []byte {
|
||||
t.Helper()
|
||||
var buffer bytes.Buffer
|
||||
writer := zip.NewWriter(&buffer)
|
||||
for name, content := range files {
|
||||
file, err := writer.Create(name)
|
||||
if err != nil {
|
||||
t.Fatalf("create zip entry failed: %v", err)
|
||||
}
|
||||
if _, err := file.Write([]byte(content)); err != nil {
|
||||
t.Fatalf("write zip entry failed: %v", err)
|
||||
}
|
||||
}
|
||||
if err := writer.Close(); err != nil {
|
||||
t.Fatalf("close zip failed: %v", err)
|
||||
}
|
||||
return buffer.Bytes()
|
||||
}
|
||||
|
||||
func multipartFileHeader(t *testing.T, fileName string, content []byte) *multipart.FileHeader {
|
||||
t.Helper()
|
||||
var body bytes.Buffer
|
||||
writer := multipart.NewWriter(&body)
|
||||
part, err := writer.CreateFormFile("package", fileName)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateFormFile failed: %v", err)
|
||||
}
|
||||
if _, err = part.Write(content); err != nil {
|
||||
t.Fatalf("write multipart file failed: %v", err)
|
||||
}
|
||||
if err = writer.Close(); err != nil {
|
||||
t.Fatalf("close multipart writer failed: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest("POST", "/", &body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
if err = req.ParseMultipartForm(int64(len(content)) + 1024); err != nil {
|
||||
t.Fatalf("ParseMultipartForm failed: %v", err)
|
||||
}
|
||||
file, header, err := req.FormFile("package")
|
||||
if err != nil {
|
||||
t.Fatalf("FormFile failed: %v", err)
|
||||
}
|
||||
file.Close()
|
||||
return header
|
||||
}
|
||||
|
||||
func TestDeletePagesDeploymentRejectsActiveDeployment(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{Name: "Active", Slug: "active", Enabled: true})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
deployment, err := UploadPagesDeployment(project.ID, multipartFileHeader(t, "site.zip", testPagesZip(t, map[string]string{"index.html": "ok"})), "", "index.html", "root")
|
||||
if err != nil {
|
||||
t.Fatalf("UploadPagesDeployment failed: %v", err)
|
||||
}
|
||||
if _, err = ActivatePagesDeployment(project.ID, deployment.ID); err != nil {
|
||||
t.Fatalf("ActivatePagesDeployment failed: %v", err)
|
||||
}
|
||||
if err = DeletePagesDeployment(project.ID, deployment.ID); err == nil {
|
||||
t.Fatal("expected active deployment deletion to fail")
|
||||
}
|
||||
var stored model.PagesDeployment
|
||||
if err = model.DB.First(&stored, deployment.ID).Error; err != nil {
|
||||
t.Fatalf("expected active deployment to remain: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadPagesDeploymentWithTopLevelFolder(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "Folder Site",
|
||||
Slug: "folder-site",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
// Upload a zip with all files inside a top-level directory "Speed-Test-source/"
|
||||
uploadHeader := multipartFileHeader(t, "site.zip", testPagesZip(t, map[string]string{
|
||||
"Speed-Test-source/index.html": "<h1>Hello Pages</h1>",
|
||||
"Speed-Test-source/assets/app.js": "console.log('pages')",
|
||||
}))
|
||||
deployment, err := UploadPagesDeployment(project.ID, uploadHeader, "", "index.html", "root")
|
||||
if err != nil {
|
||||
t.Fatalf("UploadPagesDeployment with folder failed: %v", err)
|
||||
}
|
||||
if deployment.FileCount != 2 {
|
||||
t.Fatalf("expected 2 files, got %d", deployment.FileCount)
|
||||
}
|
||||
if project.EntryFile != "index.html" {
|
||||
t.Fatalf("expected EntryFile to be index.html, got %q", project.EntryFile)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPagesProjectAPIProxyValidation(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// 1. Invalid configuration: enabled but empty fields
|
||||
_, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "API Proxy 1",
|
||||
Enabled: true,
|
||||
APIProxyEnabled: true,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "匹配路径不能为空") {
|
||||
t.Fatalf("expected error for empty match path, got: %v", err)
|
||||
}
|
||||
|
||||
// 2. Invalid path: must start with '/'
|
||||
_, err = CreatePagesProject(PagesProjectInput{
|
||||
Name: "API Proxy 2",
|
||||
Enabled: true,
|
||||
APIProxyEnabled: true,
|
||||
APIProxyPath: "api",
|
||||
APIProxyPass: "http://127.0.0.1:8080",
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "必须以 '/' 开头") {
|
||||
t.Fatalf("expected error for path not starting with /, got: %v", err)
|
||||
}
|
||||
|
||||
// 3. Invalid target URL
|
||||
_, err = CreatePagesProject(PagesProjectInput{
|
||||
Name: "API Proxy 3",
|
||||
Enabled: true,
|
||||
APIProxyEnabled: true,
|
||||
APIProxyPath: "/api",
|
||||
APIProxyPass: "127.0.0.1:8080",
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "有效的 HTTP/HTTPS URL") {
|
||||
t.Fatalf("expected error for invalid pass URL, got: %v", err)
|
||||
}
|
||||
|
||||
// 4. Valid configuration
|
||||
project, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "API Proxy Valid",
|
||||
Enabled: true,
|
||||
APIProxyEnabled: true,
|
||||
APIProxyPath: "/api",
|
||||
APIProxyPass: "http://127.0.0.1:8080",
|
||||
APIProxyRewrite: "/",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error creating valid project: %v", err)
|
||||
}
|
||||
if !project.APIProxyEnabled || project.APIProxyPath != "/api" || project.APIProxyPass != "http://127.0.0.1:8080" || project.APIProxyRewrite != "/" {
|
||||
t.Fatalf("unexpected project state: %+v", project)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadPagesDeploymentWithRootDir(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
project, err := CreatePagesProject(PagesProjectInput{
|
||||
Name: "App Site",
|
||||
Slug: "app-site",
|
||||
Enabled: true,
|
||||
RootDir: "build",
|
||||
EntryFile: "index.html",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePagesProject failed: %v", err)
|
||||
}
|
||||
|
||||
// 1. Upload a zip with files inside a subfolder.
|
||||
uploadHeader := multipartFileHeader(t, "site.zip", testPagesZip(t, map[string]string{
|
||||
"build/index.html": "<h1>App Root</h1>",
|
||||
"build/static/bundle.js": "console.log('app')",
|
||||
"README.md": "README info",
|
||||
}))
|
||||
deployment, err := UploadPagesDeployment(project.ID, uploadHeader, "build", "index.html", "root")
|
||||
if err != nil {
|
||||
t.Fatalf("UploadPagesDeployment with rootDir failed: %v", err)
|
||||
}
|
||||
if deployment.FileCount != 3 {
|
||||
t.Fatalf("expected 3 files, got %d", deployment.FileCount)
|
||||
}
|
||||
if project.RootDir != "build" {
|
||||
t.Fatalf("expected RootDir to be 'build', got %q", project.RootDir)
|
||||
}
|
||||
if project.EntryFile != "index.html" {
|
||||
t.Fatalf("expected EntryFile to be 'index.html', got %q", project.EntryFile)
|
||||
}
|
||||
|
||||
// 2. Update project configuration to a wrong entry file relative to root directory, upload should fail
|
||||
project, err = UpdatePagesProject(project.ID, PagesProjectInput{
|
||||
Name: "App Site",
|
||||
Slug: "app-site",
|
||||
Enabled: true,
|
||||
RootDir: "build",
|
||||
EntryFile: "missing.html",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdatePagesProject failed: %v", err)
|
||||
}
|
||||
_, err = UploadPagesDeployment(project.ID, uploadHeader, "build", "missing.html", "root")
|
||||
if err == nil || !strings.Contains(err.Error(), "缺少入口文件") {
|
||||
t.Fatalf("expected failure for missing entry file, got %v", err)
|
||||
}
|
||||
|
||||
// Revert to correct config for snapshot check
|
||||
project, err = UpdatePagesProject(project.ID, PagesProjectInput{
|
||||
Name: "App Site",
|
||||
Slug: "app-site",
|
||||
Enabled: true,
|
||||
RootDir: "build",
|
||||
EntryFile: "index.html",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdatePagesProject failed: %v", err)
|
||||
}
|
||||
|
||||
// 3. Test config snapshot LocalRoot path rendering
|
||||
project, err = ActivatePagesDeployment(project.ID, deployment.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("ActivatePagesDeployment failed: %v", err)
|
||||
}
|
||||
_, err = CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
Enabled: true,
|
||||
UpstreamType: "pages",
|
||||
PagesProjectID: &project.ID,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
result, err := PublishConfigVersion("root", false)
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
// Verify LocalRoot contains the rootDir
|
||||
expectedLocalRoot := fmt.Sprintf("deployments/%d/current/build", deployment.ID)
|
||||
if !strings.Contains(result.Version.SnapshotJSON, expectedLocalRoot) {
|
||||
t.Fatalf("expected snapshot JSON to include %q, got %s", expectedLocalRoot, result.Version.SnapshotJSON)
|
||||
}
|
||||
|
||||
if !strings.Contains(result.Version.RenderedConfig, "current/build") {
|
||||
t.Fatalf("expected rendered config to point to current/build, got:\n%s", result.Version.RenderedConfig)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,569 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// RelayHeartbeatPayload is the payload sent by OpenFlareRelay in each heartbeat.
|
||||
type RelayHeartbeatPayload struct {
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"frp_version"`
|
||||
RelayStatus string `json:"relay_status"`
|
||||
FrpsConnCount int `json:"frps_connections"`
|
||||
FrpsProxyCount int `json:"frps_proxy_count"`
|
||||
FrpsClientCount int `json:"frps_client_count"`
|
||||
FrpsProxies []RelayProxyStat `json:"frps_proxies,omitempty"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
Profile *AgentNodeSystemProfile `json:"profile,omitempty"`
|
||||
Snapshot *AgentNodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
HealthEvents []AgentNodeHealthEvent `json:"health_events,omitempty"`
|
||||
}
|
||||
|
||||
const relayFrpsUnhealthyEventType = "frps_unhealthy"
|
||||
|
||||
// RelayConfig is the frps configuration sent to the Relay.
|
||||
type RelayConfig struct {
|
||||
BindPort int `json:"bind_port"`
|
||||
VhostHTTPPort int `json:"vhost_http_port"`
|
||||
AuthToken string `json:"auth_token"`
|
||||
LogLevel string `json:"log_level"`
|
||||
WebServerEnabled bool `json:"web_server_enabled"`
|
||||
}
|
||||
|
||||
// RelaySettings contains runtime settings for the Relay.
|
||||
type RelaySettings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
UpdateRepo string `json:"update_repo"`
|
||||
UpdateNow bool `json:"update_now"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
}
|
||||
|
||||
// RelayHeartbeatResponse is the response returned to the Relay from a heartbeat.
|
||||
type RelayHeartbeatResponse struct {
|
||||
RelayConfig *RelayConfig `json:"relay_config"`
|
||||
RelaySettings *RelaySettings `json:"relay_settings"`
|
||||
}
|
||||
|
||||
// HeartbeatRelay processes a relay heartbeat, updates node status, and returns config.
|
||||
func HeartbeatRelay(node *model.Node, payload RelayHeartbeatPayload) (*RelayHeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("relay node is nil")
|
||||
}
|
||||
slog.Debug("relay heartbeat received", "node_id", node.NodeID)
|
||||
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
payload.RelayStatus = normalizeRelayStatus(payload.RelayStatus)
|
||||
payload.Name = strings.TrimSpace(payload.Name)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
|
||||
previous := *node
|
||||
updateNow := node.UpdateRequested
|
||||
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(node.UpdateTag)
|
||||
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = ReleaseChannelStable.String()
|
||||
node.UpdateTag = ""
|
||||
|
||||
changes := make(map[string]any)
|
||||
appendRelayChange := func(key string, before any, after any) {
|
||||
if before != after {
|
||||
changes[key] = after
|
||||
}
|
||||
}
|
||||
now := time.Now()
|
||||
appendRelayChange("version", node.Version, payload.Version)
|
||||
appendRelayChange("ext_version", node.ExtVersion, payload.ExtVersion)
|
||||
appendRelayChange("relay_status", node.RelayStatus, payload.RelayStatus)
|
||||
|
||||
if previous.UpdateRequested {
|
||||
appendRelayChange("update_requested", previous.UpdateRequested, false)
|
||||
}
|
||||
if previous.UpdateChannel != ReleaseChannelStable.String() {
|
||||
appendRelayChange("update_channel", previous.UpdateChannel, ReleaseChannelStable.String())
|
||||
}
|
||||
if previous.UpdateTag != "" {
|
||||
appendRelayChange("update_tag", previous.UpdateTag, "")
|
||||
}
|
||||
|
||||
if payload.Name != "" && strings.TrimSpace(node.Name) == "" {
|
||||
appendRelayChange("name", node.Name, payload.Name)
|
||||
node.Name = payload.Name
|
||||
}
|
||||
if payload.IP != "" && !node.IPManualOverride {
|
||||
appendRelayChange("ip", node.IP, payload.IP)
|
||||
node.IP = payload.IP
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
changes["geo_name"] = node.GeoName
|
||||
changes["geo_latitude"] = node.GeoLatitude
|
||||
changes["geo_longitude"] = node.GeoLongitude
|
||||
}
|
||||
}
|
||||
if !node.LastSeenAt.Equal(now) {
|
||||
changes["last_seen_at"] = now
|
||||
}
|
||||
changes["status"] = NodeStatusOnline
|
||||
|
||||
node.Version = payload.Version
|
||||
node.ExtVersion = payload.ExtVersion
|
||||
node.RelayStatus = payload.RelayStatus
|
||||
|
||||
node.LastSeenAt = now
|
||||
node.Status = NodeStatusOnline
|
||||
|
||||
if len(changes) > 0 {
|
||||
if err := model.DB.Model(node).Updates(changes).Error; err != nil {
|
||||
return nil, fmt.Errorf("update relay heartbeat: %w", err)
|
||||
}
|
||||
}
|
||||
if err := reconcileRelayHealthEvents(node.NodeID, payload.RelayStatus, now); err != nil {
|
||||
return nil, fmt.Errorf("reconcile relay health events: %w", err)
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
persistRelayHeartbeatObservability(node.NodeID, payload, node.LastSeenAt)
|
||||
|
||||
return &RelayHeartbeatResponse{
|
||||
RelayConfig: buildRelayConfig(node),
|
||||
RelaySettings: buildRelaySettings(node, updateNow, updateChannel.String(), updateTag),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func persistRelayHeartbeatObservability(nodeID string, payload RelayHeartbeatPayload, reportedAt time.Time) {
|
||||
persistHeartbeatObservability(nodeID, AgentNodePayload{
|
||||
Profile: payload.Profile,
|
||||
Snapshot: payload.Snapshot,
|
||||
HealthEvents: payload.HealthEvents,
|
||||
}, reportedAt)
|
||||
|
||||
frpsObs := &model.NodeObservationFrps{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: reportedAt,
|
||||
FrpsConnections: 0,
|
||||
FrpsProxyCount: 0,
|
||||
FrpsClientCount: 0,
|
||||
FrpsProxies: "",
|
||||
}
|
||||
_ = frpsObs.Insert()
|
||||
}
|
||||
|
||||
func buildRelayConfig(node *model.Node) *RelayConfig {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
return &RelayConfig{
|
||||
BindPort: node.RelayBindPort,
|
||||
VhostHTTPPort: node.RelayVhostHTTPPort,
|
||||
AuthToken: node.RelayAuthToken,
|
||||
LogLevel: "info",
|
||||
WebServerEnabled: node.RelayWebServerEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func buildRelaySettings(node *model.Node, updateNow bool, updateChannel string, updateTag string) *RelaySettings {
|
||||
autoUpdate := false
|
||||
if node != nil {
|
||||
autoUpdate = node.AutoUpdateEnabled
|
||||
}
|
||||
if strings.TrimSpace(updateChannel) == "" {
|
||||
updateChannel = ReleaseChannelStable.String()
|
||||
}
|
||||
return &RelaySettings{
|
||||
HeartbeatInterval: common.AgentHeartbeatInterval,
|
||||
WebsocketUpgradeEnabled: common.AgentWebsocketUpgradeEnabled,
|
||||
AutoUpdate: autoUpdate,
|
||||
UpdateRepo: common.AgentUpdateRepo,
|
||||
UpdateNow: updateNow,
|
||||
UpdateChannel: updateChannel,
|
||||
UpdateTag: strings.TrimSpace(updateTag),
|
||||
}
|
||||
}
|
||||
|
||||
func reconcileRelayHealthEvents(nodeID string, relayStatus string, reportedAt time.Time) error {
|
||||
if relayStatus == "unknown" {
|
||||
return nil
|
||||
}
|
||||
managedTypes := map[string]struct{}{
|
||||
relayFrpsUnhealthyEventType: {},
|
||||
}
|
||||
events := []AgentNodeHealthEvent{}
|
||||
if relayStatus == "unhealthy" {
|
||||
events = append(events, AgentNodeHealthEvent{
|
||||
EventType: relayFrpsUnhealthyEventType,
|
||||
Severity: NodeHealthSeverityCritical,
|
||||
Message: "frps runtime is not healthy",
|
||||
TriggeredAtUnix: reportedAt.Unix(),
|
||||
Metadata: map[string]string{
|
||||
"relay_status": relayStatus,
|
||||
},
|
||||
})
|
||||
}
|
||||
return model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
return reconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, managedTypes)
|
||||
})
|
||||
}
|
||||
|
||||
func normalizeRelayStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case "healthy":
|
||||
return "healthy"
|
||||
case "unhealthy":
|
||||
return "unhealthy"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
// FlaredHeartbeatPayload is the payload sent by OpenFlared in each heartbeat.
|
||||
type FlaredHeartbeatPayload struct {
|
||||
ClientVersion string `json:"client_version"`
|
||||
FrpVersion string `json:"frp_version"`
|
||||
IP string `json:"ip"`
|
||||
TunnelStatus string `json:"tunnel_status"`
|
||||
ConnectedRelays []FlaredConnectedRelay `json:"connected_relays"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
CurrentChecksum string `json:"current_checksum"`
|
||||
}
|
||||
|
||||
func normalizeFlaredHeartbeatPayload(payload FlaredHeartbeatPayload) FlaredHeartbeatPayload {
|
||||
payload.ClientVersion = strings.TrimSpace(payload.ClientVersion)
|
||||
payload.FrpVersion = strings.TrimSpace(payload.FrpVersion)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
payload.TunnelStatus = strings.ToLower(strings.TrimSpace(payload.TunnelStatus))
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.CurrentChecksum = strings.TrimSpace(payload.CurrentChecksum)
|
||||
cleaned := make([]FlaredConnectedRelay, 0, len(payload.ConnectedRelays))
|
||||
for _, relay := range payload.ConnectedRelays {
|
||||
relay.RelayNodeID = strings.TrimSpace(relay.RelayNodeID)
|
||||
relay.Status = strings.ToLower(strings.TrimSpace(relay.Status))
|
||||
if relay.RelayNodeID == "" {
|
||||
continue
|
||||
}
|
||||
if relay.Status == "" {
|
||||
relay.Status = "unknown"
|
||||
}
|
||||
cleaned = append(cleaned, relay)
|
||||
}
|
||||
payload.ConnectedRelays = cleaned
|
||||
return payload
|
||||
}
|
||||
|
||||
// HeartbeatFlared processes an OpenFlared heartbeat, refreshes node status,
|
||||
// persists the connected relay snapshot, and returns the active tunnel
|
||||
// config summary plus runtime settings.
|
||||
func HeartbeatFlared(node *model.Node, payload FlaredHeartbeatPayload) (*FlaredHeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("tunnel client node is nil")
|
||||
}
|
||||
if node.NodeType != "tunnel_client" {
|
||||
return nil, fmt.Errorf("node %s is not a tunnel_client", node.NodeID)
|
||||
}
|
||||
slog.Debug("flared heartbeat received", "node_id", node.NodeID, "client_version", payload.ClientVersion)
|
||||
payload = normalizeFlaredHeartbeatPayload(payload)
|
||||
|
||||
now := time.Now()
|
||||
previous := *node
|
||||
updateNow := node.UpdateRequested
|
||||
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(node.UpdateTag)
|
||||
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = ReleaseChannelStable.String()
|
||||
node.UpdateTag = ""
|
||||
|
||||
changes := make(map[string]any)
|
||||
if previous.Version != payload.ClientVersion {
|
||||
changes["version"] = payload.ClientVersion
|
||||
}
|
||||
if previous.ExtVersion != payload.FrpVersion {
|
||||
changes["ext_version"] = payload.FrpVersion
|
||||
}
|
||||
if previous.CurrentVersion != payload.CurrentVersion {
|
||||
changes["current_version"] = payload.CurrentVersion
|
||||
}
|
||||
if !previous.LastSeenAt.Equal(now) {
|
||||
changes["last_seen_at"] = now
|
||||
}
|
||||
changes["status"] = NodeStatusOnline
|
||||
|
||||
node.Version = payload.ClientVersion
|
||||
node.ExtVersion = payload.FrpVersion
|
||||
node.CurrentVersion = payload.CurrentVersion
|
||||
node.LastSeenAt = now
|
||||
node.Status = NodeStatusOnline
|
||||
|
||||
if previous.UpdateRequested {
|
||||
changes["update_requested"] = false
|
||||
}
|
||||
if previous.UpdateChannel != ReleaseChannelStable.String() {
|
||||
changes["update_channel"] = ReleaseChannelStable.String()
|
||||
}
|
||||
if previous.UpdateTag != "" {
|
||||
changes["update_tag"] = ""
|
||||
}
|
||||
|
||||
if !node.IPManualOverride && payload.IP != "" && previous.IP != payload.IP {
|
||||
changes["ip"] = payload.IP
|
||||
node.IP = payload.IP
|
||||
}
|
||||
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
if previous.GeoName != node.GeoName {
|
||||
changes["geo_name"] = node.GeoName
|
||||
}
|
||||
if !coordinatesEqual(previous.GeoLatitude, node.GeoLatitude) {
|
||||
changes["geo_latitude"] = node.GeoLatitude
|
||||
}
|
||||
if !coordinatesEqual(previous.GeoLongitude, node.GeoLongitude) {
|
||||
changes["geo_longitude"] = node.GeoLongitude
|
||||
}
|
||||
}
|
||||
|
||||
if len(changes) > 0 {
|
||||
if err := model.DB.Model(node).Updates(changes).Error; err != nil {
|
||||
return nil, fmt.Errorf("update flared heartbeat: %w", err)
|
||||
}
|
||||
}
|
||||
refreshAccessTokenCache(node)
|
||||
persistFlaredObservability(node.NodeID, payload, now)
|
||||
|
||||
activeConfig, err := GetActiveConfigMetaForAgent()
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
return &FlaredHeartbeatResponse{
|
||||
ActiveConfig: activeConfig,
|
||||
TunnelSettings: buildRelaySettings(node, updateNow, updateChannel.String(), updateTag),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// persistFlaredObservability records the latest connection snapshot and
|
||||
// health event for the OpenFlared client.
|
||||
func persistFlaredObservability(nodeID string, payload FlaredHeartbeatPayload, reportedAt time.Time) {
|
||||
connected := make([]string, 0, len(payload.ConnectedRelays))
|
||||
for _, relay := range payload.ConnectedRelays {
|
||||
connected = append(connected, fmt.Sprintf("%s:%s", relay.RelayNodeID, relay.Status))
|
||||
}
|
||||
managedTypes := map[string]struct{}{
|
||||
"flared_runtime_unhealthy": {},
|
||||
}
|
||||
var events []AgentNodeHealthEvent
|
||||
if payload.TunnelStatus == "unhealthy" {
|
||||
events = append(events, AgentNodeHealthEvent{
|
||||
EventType: "flared_runtime_unhealthy",
|
||||
Severity: NodeHealthSeverityCritical,
|
||||
Message: "openflared runtime is not healthy",
|
||||
TriggeredAtUnix: reportedAt.Unix(),
|
||||
Metadata: map[string]string{
|
||||
"tunnel_status": payload.TunnelStatus,
|
||||
"client_version": payload.ClientVersion,
|
||||
"current_version": payload.CurrentVersion,
|
||||
"current_checksum": payload.CurrentChecksum,
|
||||
"connected_relays": strings.Join(connected, ","),
|
||||
},
|
||||
})
|
||||
}
|
||||
_ = reconcileScopedNodeHealthEvents(model.DB, nodeID, events, reportedAt, managedTypes)
|
||||
}
|
||||
|
||||
// FlaredConnectedRelay describes the status of a relay connection from a client.
|
||||
type FlaredConnectedRelay struct {
|
||||
RelayNodeID string `json:"relay_node_id"`
|
||||
Status string `json:"status"`
|
||||
ProxyCount int `json:"proxy_count"`
|
||||
}
|
||||
|
||||
// FlaredHeartbeatResponse is the response returned to the OpenFlared client.
|
||||
type FlaredHeartbeatResponse struct {
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
TunnelSettings *RelaySettings `json:"tunnel_settings"`
|
||||
}
|
||||
|
||||
// FlaredTunnelConfigResponse is the full tunnel routing config sent to the client.
|
||||
type FlaredTunnelConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
Relays []FlaredRelayInfo `json:"relays"`
|
||||
Proxies []FlaredProxyEntry `json:"proxies"`
|
||||
}
|
||||
|
||||
// FlaredRelayInfo describes a relay that the client should connect to.
|
||||
type FlaredRelayInfo struct {
|
||||
RelayNodeID string `json:"relay_node_id"`
|
||||
Address string `json:"address"`
|
||||
AuthToken string `json:"auth_token"`
|
||||
ProxyURL string `json:"proxy_url"`
|
||||
}
|
||||
|
||||
// FlaredProxyEntry describes a single frpc proxy definition.
|
||||
type FlaredProxyEntry struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
LocalAddr string `json:"local_addr"`
|
||||
LocalPort int `json:"local_port"`
|
||||
CustomDomains []string `json:"custom_domains"`
|
||||
}
|
||||
|
||||
// GetFlaredTunnelConfig builds the full tunnel routing config for an OpenFlared client.
|
||||
func GetFlaredTunnelConfig(node *model.Node) (*FlaredTunnelConfigResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("node is nil")
|
||||
}
|
||||
|
||||
activeVersion, err := model.GetActiveConfigVersion()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("no active config version: %w", err)
|
||||
}
|
||||
|
||||
// Get all enabled proxy routes with tunnel upstream targeting this tunnel
|
||||
routes, err := model.GetEnabledProxyRoutes()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get proxy routes: %w", err)
|
||||
}
|
||||
|
||||
// Get all online tunnel relay nodes
|
||||
relayNodes, err := model.ListNodesByType("tunnel_relay")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get relay nodes: %w", err)
|
||||
}
|
||||
|
||||
// Build relay info
|
||||
relays := make([]FlaredRelayInfo, 0, len(relayNodes))
|
||||
for _, node := range relayNodes {
|
||||
if node.RelayStatus == "healthy" || node.Status == NodeStatusOnline {
|
||||
addr := relayClientAddress(node)
|
||||
relays = append(relays, FlaredRelayInfo{
|
||||
RelayNodeID: node.NodeID,
|
||||
Address: addr,
|
||||
AuthToken: node.RelayAuthToken,
|
||||
ProxyURL: strings.TrimSpace(node.RelayClientProxyURL),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Build proxy entries from routes
|
||||
proxies := make([]FlaredProxyEntry, 0)
|
||||
for _, route := range routes {
|
||||
if route.UpstreamType != "tunnel" || route.TunnelNodeID == nil || *route.TunnelNodeID != node.ID {
|
||||
continue
|
||||
}
|
||||
if !route.Enabled {
|
||||
continue
|
||||
}
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
localAddr, localPort := parseTunnelTargetAddr(route.TunnelTargetAddr)
|
||||
proxies = append(proxies, FlaredProxyEntry{
|
||||
Name: fmt.Sprintf("%s-%s", node.NodeID, sanitizeProxyName(domains[0])),
|
||||
Type: "http",
|
||||
LocalAddr: localAddr,
|
||||
LocalPort: localPort,
|
||||
CustomDomains: domains,
|
||||
})
|
||||
}
|
||||
|
||||
return &FlaredTunnelConfigResponse{
|
||||
Version: activeVersion.Version,
|
||||
Checksum: activeVersion.Checksum,
|
||||
Relays: relays,
|
||||
Proxies: proxies,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func relayClientAddress(node *model.Node) string {
|
||||
if node == nil {
|
||||
return ""
|
||||
}
|
||||
port := node.RelayBindPort
|
||||
if port <= 0 {
|
||||
port = 7000
|
||||
}
|
||||
addr := strings.TrimSpace(node.RelayClientAccessAddr)
|
||||
if addr == "" {
|
||||
addr = strings.TrimSpace(node.IP)
|
||||
}
|
||||
if addr == "" {
|
||||
return fmt.Sprintf("127.0.0.1:%d", port)
|
||||
}
|
||||
if _, _, err := net.SplitHostPort(addr); err == nil {
|
||||
return addr
|
||||
}
|
||||
if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 {
|
||||
return net.JoinHostPort(addr, strconv.Itoa(port))
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", addr, port)
|
||||
}
|
||||
|
||||
func relayAgentAddress(node *model.Node) string {
|
||||
if node == nil {
|
||||
return ""
|
||||
}
|
||||
port := node.RelayVhostHTTPPort
|
||||
if port <= 0 {
|
||||
port = 8080
|
||||
}
|
||||
addr := strings.TrimSpace(node.RelayAgentAccessAddr)
|
||||
if addr == "" {
|
||||
addr = strings.TrimSpace(node.RelayClientAccessAddr)
|
||||
}
|
||||
if addr == "" {
|
||||
addr = strings.TrimSpace(node.IP)
|
||||
}
|
||||
if addr == "" {
|
||||
return fmt.Sprintf("127.0.0.1:%d", port)
|
||||
}
|
||||
if _, _, err := net.SplitHostPort(addr); err == nil {
|
||||
return addr
|
||||
}
|
||||
if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 {
|
||||
return net.JoinHostPort(addr, strconv.Itoa(port))
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", addr, port)
|
||||
}
|
||||
|
||||
func parseTunnelTargetAddr(addr string) (string, int) {
|
||||
addr = strings.TrimSpace(addr)
|
||||
if addr == "" {
|
||||
return "127.0.0.1", 80
|
||||
}
|
||||
host, portStr, err := splitHostPort(addr)
|
||||
if err != nil {
|
||||
return addr, 80
|
||||
}
|
||||
port := 80
|
||||
if _, err := fmt.Sscanf(portStr, "%d", &port); err != nil {
|
||||
port = 80
|
||||
}
|
||||
return host, port
|
||||
}
|
||||
|
||||
func splitHostPort(addr string) (string, string, error) {
|
||||
lastColon := strings.LastIndex(addr, ":")
|
||||
if lastColon < 0 {
|
||||
return addr, "", fmt.Errorf("no port")
|
||||
}
|
||||
return addr[:lastColon], addr[lastColon+1:], nil
|
||||
}
|
||||
|
||||
func sanitizeProxyName(domain string) string {
|
||||
return strings.ReplaceAll(strings.ReplaceAll(domain, ".", "-"), "*", "wildcard")
|
||||
}
|
||||
@@ -0,0 +1,529 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestHeartbeatRelayPersistsRuntimeAndObservability(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "node-relay-observe",
|
||||
Name: "relay-1",
|
||||
IP: "",
|
||||
AccessToken: "relay-token",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "tunnel_relay",
|
||||
RelayStatus: "unknown",
|
||||
Version: "",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed relay node: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UTC()
|
||||
_, err := HeartbeatRelay(node, RelayHeartbeatPayload{
|
||||
Version: "v0.1.0",
|
||||
ExtVersion: "0.61.0",
|
||||
RelayStatus: "healthy",
|
||||
FrpsConnCount: 7,
|
||||
FrpsProxyCount: 3,
|
||||
Name: "relay-runtime",
|
||||
IP: "203.0.113.9",
|
||||
Profile: &AgentNodeSystemProfile{
|
||||
Hostname: "relay-runtime",
|
||||
OSName: "Ubuntu",
|
||||
OSVersion: "24.04",
|
||||
Architecture: "amd64",
|
||||
CPUCores: 4,
|
||||
ReportedAtUnix: now.Unix(),
|
||||
},
|
||||
Snapshot: &AgentNodeMetricSnapshot{
|
||||
CapturedAtUnix: now.Unix(),
|
||||
CPUUsagePercent: 12.5,
|
||||
NetworkRxBytes: 1024,
|
||||
NetworkTxBytes: 2048,
|
||||
},
|
||||
HealthEvents: []AgentNodeHealthEvent{},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("HeartbeatRelay failed: %v", err)
|
||||
}
|
||||
|
||||
updated, err := model.GetNodeByNodeID(node.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload node: %v", err)
|
||||
}
|
||||
if updated.Status != NodeStatusOnline || updated.RelayStatus != "healthy" {
|
||||
t.Fatalf("unexpected relay status: %+v", updated)
|
||||
}
|
||||
if updated.IP != "203.0.113.9" {
|
||||
t.Fatalf("expected relay IP to be updated, got %q", updated.IP)
|
||||
}
|
||||
if updated.Version != "v0.1.0" || updated.ExtVersion != "0.61.0" {
|
||||
t.Fatalf("expected relay versions to be updated, got relay=%q frp=%q", updated.Version, updated.ExtVersion)
|
||||
}
|
||||
|
||||
profile, err := model.GetNodeSystemProfile(node.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("expected relay system profile: %v", err)
|
||||
}
|
||||
if profile.Hostname != "relay-runtime" || profile.OSName != "Ubuntu" {
|
||||
t.Fatalf("unexpected relay profile: %+v", profile)
|
||||
}
|
||||
|
||||
snapshots, err := model.ListNodeMetricSnapshots(node.NodeID, now.Add(-time.Minute), 10)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to list relay snapshots: %v", err)
|
||||
}
|
||||
if len(snapshots) != 1 || snapshots[0].CPUUsagePercent != 12.5 {
|
||||
t.Fatalf("unexpected relay snapshots: %+v", snapshots)
|
||||
}
|
||||
|
||||
observability, err := GetNodeObservability(updated.ID, NodeObservabilityQuery{Hours: 1, Limit: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeObservability failed: %v", err)
|
||||
}
|
||||
if observability.RelayDashboard == nil {
|
||||
t.Fatal("expected relay dashboard snapshot")
|
||||
}
|
||||
if observability.RelayDashboard.TotalConnections != 0 || observability.RelayDashboard.TotalProxies != 0 {
|
||||
// Frps telemetry collection is disabled; dashboard values are always zero.
|
||||
t.Fatalf("unexpected relay dashboard: %+v", observability.RelayDashboard)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredRejectsWrongNodeType(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "node-not-tunnel-client",
|
||||
Name: "edge",
|
||||
IP: "10.0.0.1",
|
||||
AccessToken: "edge-token",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "edge_node",
|
||||
Version: "v0.0.0",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed edge node: %v", err)
|
||||
}
|
||||
|
||||
_, err := HeartbeatFlared(node, FlaredHeartbeatPayload{
|
||||
ClientVersion: "v0.1.0",
|
||||
FrpVersion: "0.61.0",
|
||||
TunnelStatus: "running",
|
||||
CurrentVersion: "v1",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected error for non-tunnel_client node type")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredRejectsNilNode(t *testing.T) {
|
||||
if _, err := HeartbeatFlared(nil, FlaredHeartbeatPayload{}); err == nil {
|
||||
t.Fatal("expected error when node is nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredPersistsRuntime(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "node-flared-1",
|
||||
Name: "flared-1",
|
||||
IP: "",
|
||||
AccessToken: "tunnel-token-abc",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "tunnel_client",
|
||||
Version: "",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed flared node: %v", err)
|
||||
}
|
||||
|
||||
resp, err := HeartbeatFlared(node, FlaredHeartbeatPayload{
|
||||
ClientVersion: " v0.2.0 ",
|
||||
FrpVersion: " 0.61.1 ",
|
||||
IP: " 192.168.1.10 ",
|
||||
TunnelStatus: " RUNNING ",
|
||||
ConnectedRelays: []FlaredConnectedRelay{
|
||||
{RelayNodeID: " node-relay-1 ", Status: " HEALTHY ", ProxyCount: 3},
|
||||
{RelayNodeID: "", Status: "running"},
|
||||
{RelayNodeID: "node-relay-2", Status: ""},
|
||||
},
|
||||
CurrentVersion: "v1",
|
||||
CurrentChecksum: "checksum-1",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("HeartbeatFlared failed: %v", err)
|
||||
}
|
||||
if resp == nil {
|
||||
t.Fatal("expected non-nil response")
|
||||
}
|
||||
if resp.TunnelSettings == nil {
|
||||
t.Fatal("expected tunnel_settings in response")
|
||||
}
|
||||
if resp.TunnelSettings.HeartbeatInterval == 0 {
|
||||
t.Fatal("expected heartbeat interval to be set in tunnel_settings")
|
||||
}
|
||||
|
||||
updated, err := model.GetNodeByNodeID(node.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload flared node: %v", err)
|
||||
}
|
||||
if updated.Status != NodeStatusOnline {
|
||||
t.Fatalf("expected flared node to be online, got %q", updated.Status)
|
||||
}
|
||||
if updated.Version != "v0.2.0" {
|
||||
t.Fatalf("expected client_version to be trimmed and stored, got %q", updated.Version)
|
||||
}
|
||||
if updated.ExtVersion != "0.61.1" {
|
||||
t.Fatalf("expected frp_version to be trimmed and stored, got %q", updated.ExtVersion)
|
||||
}
|
||||
if updated.IP != "192.168.1.10" {
|
||||
t.Fatalf("expected IP to be trimmed and stored, got %q", updated.IP)
|
||||
}
|
||||
if updated.CurrentVersion != "v1" {
|
||||
t.Fatalf("expected current_version to be stored, got %q", updated.CurrentVersion)
|
||||
}
|
||||
if updated.LastSeenAt.IsZero() {
|
||||
t.Fatal("expected last_seen_at to be updated")
|
||||
}
|
||||
|
||||
// Test IPManualOverride
|
||||
updated.IPManualOverride = true
|
||||
if err := updated.Update(); err != nil {
|
||||
t.Fatalf("failed to lock IP: %v", err)
|
||||
}
|
||||
|
||||
_, err = HeartbeatFlared(updated, FlaredHeartbeatPayload{
|
||||
ClientVersion: "v0.2.0",
|
||||
FrpVersion: "0.61.1",
|
||||
IP: "10.0.0.99",
|
||||
TunnelStatus: "running",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("second HeartbeatFlared failed: %v", err)
|
||||
}
|
||||
|
||||
lockedNode, err := model.GetNodeByNodeID(node.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload locked node: %v", err)
|
||||
}
|
||||
if lockedNode.IP != "192.168.1.10" {
|
||||
t.Fatalf("expected IP to stay locked at 192.168.1.10, but got %q", lockedNode.IP)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredTrimsAndFiltersRelays(t *testing.T) {
|
||||
normalized := normalizeFlaredHeartbeatPayload(FlaredHeartbeatPayload{
|
||||
TunnelStatus: " UNHEALTHY ",
|
||||
ConnectedRelays: []FlaredConnectedRelay{
|
||||
{RelayNodeID: " node-a ", Status: " OK "},
|
||||
{RelayNodeID: "", Status: "running"},
|
||||
},
|
||||
})
|
||||
if normalized.TunnelStatus != "unhealthy" {
|
||||
t.Fatalf("expected tunnel_status to be lower-cased, got %q", normalized.TunnelStatus)
|
||||
}
|
||||
if len(normalized.ConnectedRelays) != 1 {
|
||||
t.Fatalf("expected empty relay_node_id to be dropped, got %+v", normalized.ConnectedRelays)
|
||||
}
|
||||
relay := normalized.ConnectedRelays[0]
|
||||
if relay.RelayNodeID != "node-a" {
|
||||
t.Fatalf("expected relay_node_id to be trimmed, got %q", relay.RelayNodeID)
|
||||
}
|
||||
if relay.Status != "ok" {
|
||||
t.Fatalf("expected status to be lower-cased, got %q", relay.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "node-flared-unhealthy",
|
||||
Name: "flared-unhealthy",
|
||||
IP: "",
|
||||
AccessToken: "tunnel-token-unhealthy",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "tunnel_client",
|
||||
Version: "",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed flared node: %v", err)
|
||||
}
|
||||
|
||||
if _, err := HeartbeatFlared(node, FlaredHeartbeatPayload{
|
||||
ClientVersion: "v0.2.0",
|
||||
FrpVersion: "0.61.0",
|
||||
TunnelStatus: "unhealthy",
|
||||
CurrentVersion: "v1",
|
||||
CurrentChecksum: "checksum-1",
|
||||
}); err != nil {
|
||||
t.Fatalf("HeartbeatFlared failed: %v", err)
|
||||
}
|
||||
|
||||
events, err := model.ListNodeHealthEvents(node.NodeID, false, 20)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodeHealthEvents failed: %v", err)
|
||||
}
|
||||
if len(events) == 0 {
|
||||
t.Fatal("expected unhealthy heartbeat to emit a node health event")
|
||||
}
|
||||
foundUnhealthy := false
|
||||
for _, event := range events {
|
||||
if event.EventType == "flared_runtime_unhealthy" {
|
||||
foundUnhealthy = true
|
||||
}
|
||||
}
|
||||
if !foundUnhealthy {
|
||||
t.Fatalf("expected flared_runtime_unhealthy event in %+v", events)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredEmitsEmptyConnectedRelays(t *testing.T) {
|
||||
normalized := normalizeFlaredHeartbeatPayload(FlaredHeartbeatPayload{})
|
||||
if normalized.ConnectedRelays == nil {
|
||||
t.Fatal("expected ConnectedRelays to be non-nil empty slice for nil input")
|
||||
}
|
||||
if len(normalized.ConnectedRelays) != 0 {
|
||||
t.Fatalf("expected empty ConnectedRelays, got %+v", normalized.ConnectedRelays)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetFlaredTunnelConfigRequiresActiveVersion(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "node-flared-noactive",
|
||||
Name: "flared-noactive",
|
||||
IP: "",
|
||||
AccessToken: "tunnel-token-na",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "tunnel_client",
|
||||
Version: "",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed flared node: %v", err)
|
||||
}
|
||||
|
||||
_, err := GetFlaredTunnelConfig(node)
|
||||
if err == nil {
|
||||
t.Fatal("expected error when no active config version exists")
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
// We accept either wrapping the underlying error or surfacing a friendly message.
|
||||
// Just ensure we surface a clear failure instead of a nil result.
|
||||
t.Logf("GetFlaredTunnelConfig returned wrapped error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelRoutePublishAndFlaredConfigUseRelayPorts(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
relayNode := &model.Node{
|
||||
NodeID: "node-relay-ports",
|
||||
Name: "relay-ports",
|
||||
IP: "85.235.64.179",
|
||||
AccessToken: "relay-token-ports",
|
||||
Status: NodeStatusOnline,
|
||||
NodeType: "tunnel_relay",
|
||||
RelayStatus: "healthy",
|
||||
RelayBindPort: 17000,
|
||||
RelayVhostHTTPPort: 18080,
|
||||
RelayAuthToken: "relay-auth-token",
|
||||
RelayClientAccessAddr: "de-e",
|
||||
}
|
||||
if err := relayNode.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed relay node: %v", err)
|
||||
}
|
||||
tunnelNode := &model.Node{
|
||||
NodeID: "node-flared-ports",
|
||||
Name: "flared-ports",
|
||||
IP: "",
|
||||
AccessToken: "tunnel-token-ports",
|
||||
Status: NodeStatusOnline,
|
||||
NodeType: "tunnel_client",
|
||||
Version: "v0.2.0",
|
||||
}
|
||||
if err := tunnelNode.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed tunnel client node: %v", err)
|
||||
}
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "flared.example.com",
|
||||
UpstreamType: "tunnel",
|
||||
TunnelID: &tunnelNode.ID,
|
||||
TunnelTargetAddr: "10.0.0.8:8080",
|
||||
TunnelTargetProtocol: "http",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if route.TunnelNodeID == nil || *route.TunnelNodeID != tunnelNode.ID {
|
||||
t.Fatalf("expected legacy tunnel_id to bind tunnel_node_id, got %+v", route.TunnelNodeID)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root", false)
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "server de-e:18080 max_fails=3 fail_timeout=10s;") {
|
||||
t.Fatalf("expected rendered OpenResty upstream to use relay vhost port, got:\n%s", result.Version.RenderedConfig)
|
||||
}
|
||||
|
||||
config, err := GetFlaredTunnelConfig(tunnelNode)
|
||||
if err != nil {
|
||||
t.Fatalf("GetFlaredTunnelConfig failed: %v", err)
|
||||
}
|
||||
if len(config.Relays) != 1 {
|
||||
t.Fatalf("expected one relay, got %+v", config.Relays)
|
||||
}
|
||||
if config.Relays[0].Address != "de-e:17000" {
|
||||
t.Fatalf("expected relay client address to include bind port, got %q", config.Relays[0].Address)
|
||||
}
|
||||
if len(config.Proxies) != 1 {
|
||||
t.Fatalf("expected one proxy, got %+v", config.Proxies)
|
||||
}
|
||||
proxy := config.Proxies[0]
|
||||
if proxy.LocalAddr != "10.0.0.8" || proxy.LocalPort != 8080 {
|
||||
t.Fatalf("unexpected proxy target: %+v", proxy)
|
||||
}
|
||||
if len(proxy.CustomDomains) != 1 || proxy.CustomDomains[0] != "flared.example.com" {
|
||||
t.Fatalf("unexpected proxy domains: %+v", proxy.CustomDomains)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatRelaySelfUpdatePropagationAndReset(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "relay-update-node",
|
||||
Name: "relay-u",
|
||||
IP: "1.1.1.1",
|
||||
AccessToken: "relay-update-token",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "tunnel_relay",
|
||||
AutoUpdateEnabled: true,
|
||||
UpdateRequested: true,
|
||||
UpdateChannel: "preview",
|
||||
UpdateTag: "v1.2.3",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed relay node: %v", err)
|
||||
}
|
||||
|
||||
resp, err := HeartbeatRelay(node, RelayHeartbeatPayload{
|
||||
Version: "v1.0.0",
|
||||
ExtVersion: "0.61.0",
|
||||
RelayStatus: "healthy",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("HeartbeatRelay failed: %v", err)
|
||||
}
|
||||
|
||||
if resp == nil || resp.RelaySettings == nil {
|
||||
t.Fatal("expected non-nil response with RelaySettings")
|
||||
}
|
||||
|
||||
settings := resp.RelaySettings
|
||||
if !settings.AutoUpdate {
|
||||
t.Error("expected AutoUpdate to be true")
|
||||
}
|
||||
if !settings.UpdateNow {
|
||||
t.Error("expected UpdateNow to be true")
|
||||
}
|
||||
if settings.UpdateChannel != "preview" {
|
||||
t.Errorf("expected UpdateChannel to be preview, got %q", settings.UpdateChannel)
|
||||
}
|
||||
if settings.UpdateTag != "v1.2.3" {
|
||||
t.Errorf("expected UpdateTag to be v1.2.3, got %q", settings.UpdateTag)
|
||||
}
|
||||
|
||||
// Verify that the requested update was cleared in the DB
|
||||
updated, err := model.GetNodeByNodeID(node.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload node: %v", err)
|
||||
}
|
||||
if updated.UpdateRequested {
|
||||
t.Error("expected UpdateRequested to be reset to false in the database")
|
||||
}
|
||||
if updated.UpdateChannel != "stable" {
|
||||
t.Errorf("expected UpdateChannel to be reset to stable, got %q", updated.UpdateChannel)
|
||||
}
|
||||
if updated.UpdateTag != "" {
|
||||
t.Errorf("expected UpdateTag to be reset to empty, got %q", updated.UpdateTag)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredSelfUpdatePropagationAndReset(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
node := &model.Node{
|
||||
NodeID: "flared-update-node",
|
||||
Name: "flared-u",
|
||||
IP: "1.1.1.2",
|
||||
AccessToken: "flared-update-token",
|
||||
Status: NodeStatusPending,
|
||||
NodeType: "tunnel_client",
|
||||
AutoUpdateEnabled: true,
|
||||
UpdateRequested: true,
|
||||
UpdateChannel: "stable",
|
||||
UpdateTag: "v2.3.4",
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
t.Fatalf("failed to seed flared node: %v", err)
|
||||
}
|
||||
|
||||
resp, err := HeartbeatFlared(node, FlaredHeartbeatPayload{
|
||||
ClientVersion: "v1.0.0",
|
||||
FrpVersion: "0.61.0",
|
||||
TunnelStatus: "running",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("HeartbeatFlared failed: %v", err)
|
||||
}
|
||||
|
||||
if resp == nil || resp.TunnelSettings == nil {
|
||||
t.Fatal("expected non-nil response with TunnelSettings")
|
||||
}
|
||||
|
||||
settings := resp.TunnelSettings
|
||||
if !settings.AutoUpdate {
|
||||
t.Error("expected AutoUpdate to be true")
|
||||
}
|
||||
if !settings.UpdateNow {
|
||||
t.Error("expected UpdateNow to be true")
|
||||
}
|
||||
if settings.UpdateChannel != "stable" {
|
||||
t.Errorf("expected UpdateChannel to be stable, got %q", settings.UpdateChannel)
|
||||
}
|
||||
if settings.UpdateTag != "v2.3.4" {
|
||||
t.Errorf("expected UpdateTag to be v2.3.4, got %q", settings.UpdateTag)
|
||||
}
|
||||
|
||||
// Verify that the requested update was cleared in the DB
|
||||
updated, err := model.GetNodeByNodeID(node.NodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload node: %v", err)
|
||||
}
|
||||
if updated.UpdateRequested {
|
||||
t.Error("expected UpdateRequested to be reset to false in the database")
|
||||
}
|
||||
if updated.UpdateChannel != "stable" {
|
||||
t.Errorf("expected UpdateChannel to be reset to stable, got %q", updated.UpdateChannel)
|
||||
}
|
||||
if updated.UpdateTag != "" {
|
||||
t.Errorf("expected UpdateTag to be reset to empty, got %q", updated.UpdateTag)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package service
|
||||
|
||||
const (
|
||||
RelayWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__"
|
||||
)
|
||||
|
||||
var DefaultRelayWSHub = NewWSHub("relay")
|
||||
|
||||
func RegisterRelayWSClient(nodeID string) *WSClient {
|
||||
return DefaultRelayWSHub.Register(nodeID)
|
||||
}
|
||||
|
||||
func UnregisterRelayWSClient(client *WSClient) {
|
||||
DefaultRelayWSHub.Unregister(client)
|
||||
}
|
||||
|
||||
func IsRelayWSConnected(nodeID string) bool {
|
||||
return DefaultRelayWSHub.IsConnected(nodeID)
|
||||
}
|
||||
|
||||
func SendRelayWSPing(nodeID string) bool {
|
||||
return DefaultRelayWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: "ping",
|
||||
})
|
||||
}
|
||||
|
||||
func SendRelayWSPong(nodeID string) bool {
|
||||
return DefaultRelayWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: "pong",
|
||||
})
|
||||
}
|
||||
|
||||
func SendRelayWSConfig(nodeID string, config *RelayConfig) bool {
|
||||
if config == nil {
|
||||
return false
|
||||
}
|
||||
return DefaultRelayWSHub.SendMessage(nodeID, WSMessage{
|
||||
Type: "relay_config",
|
||||
Payload: config,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,272 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestAcmeAndDnsIntegration(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
// 1. Create a DNS Account
|
||||
dnsAccount := &model.DnsAccount{
|
||||
Name: "Test Cloudflare",
|
||||
Type: "cloudflare",
|
||||
Authorization: `{"api_token": "dummy_token"}`,
|
||||
}
|
||||
if err := dnsAccount.Insert(); err != nil {
|
||||
t.Fatalf("Failed to insert DNS Account: %v", err)
|
||||
}
|
||||
|
||||
// 2. Apply for TLS Certificate (using the new ApplyTLSCertificate function)
|
||||
certInput := TLSApplyInput{
|
||||
Name: "Test ACME Cert",
|
||||
PrimaryDomain: "example.com",
|
||||
OtherDomains: "*.example.com",
|
||||
DnsAccountID: dnsAccount.ID,
|
||||
KeyAlgorithm: "RSA2048",
|
||||
AutoRenew: true,
|
||||
}
|
||||
|
||||
cert, err := ApplyTLSCertificate(certInput)
|
||||
if err != nil {
|
||||
t.Fatalf("ApplyTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
if cert.ApplyStatus != "applying" {
|
||||
t.Fatalf("Expected cert ApplyStatus to be applying, got %s", cert.ApplyStatus)
|
||||
}
|
||||
|
||||
if cert.Provider != "acme" {
|
||||
t.Fatalf("Expected cert Provider to be acme, got %s", cert.Provider)
|
||||
}
|
||||
|
||||
// 3. Try to delete the DNS account (should fail since it's used by the cert)
|
||||
// Actually, the delete logic is in the controller for the foreign key check.
|
||||
// But let's check if the controller logic can be tested here, or we just trust the DB setup.
|
||||
var count int64
|
||||
model.DB.Model(&model.TLSCertificate{}).Where("dns_account_id = ?", dnsAccount.ID).Count(&count)
|
||||
if count != 1 {
|
||||
t.Fatalf("Expected 1 certificate associated with DNS account, got %d", count)
|
||||
}
|
||||
|
||||
// 4. Test RenewTLSCertificate
|
||||
renewedCert, err := RenewTLSCertificate(cert.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("RenewTLSCertificate failed: %v", err)
|
||||
}
|
||||
if renewedCert.ApplyStatus != "applying" {
|
||||
t.Fatalf("Expected renewed cert ApplyStatus to be applying, got %s", renewedCert.ApplyStatus)
|
||||
}
|
||||
|
||||
// Wait for the async goroutine to fail (it now registers an LE account, which takes longer)
|
||||
time.Sleep(5 * time.Second)
|
||||
|
||||
// Reload cert and verify error status
|
||||
finalCert, err := model.GetTLSCertificateByID(renewedCert.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to reload cert: %v", err)
|
||||
}
|
||||
if finalCert.ApplyStatus != "error" {
|
||||
t.Fatalf("Expected final cert ApplyStatus to be error, got %s", finalCert.ApplyStatus)
|
||||
}
|
||||
if finalCert.ApplyMessage == "" {
|
||||
t.Fatalf("Expected final cert ApplyMessage to be populated, got empty")
|
||||
}
|
||||
|
||||
// Clean up
|
||||
if err := DeleteTLSCertificate(cert.ID); err != nil {
|
||||
t.Fatalf("DeleteTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
if err := dnsAccount.Delete(); err != nil {
|
||||
t.Fatalf("Failed to delete DNS Account after cert cleanup: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertTLSCertificateToAcmePreservesUploadUntilSuccess(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
originalCertPEM, originalKeyPEM := generateCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: originalCertPEM,
|
||||
KeyPEM: originalKeyPEM,
|
||||
Remark: "manual upload",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
originalCertPEM = cert.CertPEM
|
||||
originalKeyPEM = cert.KeyPEM
|
||||
|
||||
newCertPEM, newKeyPEM := generateCertificatePair(t, []string{"managed.example.com"})
|
||||
started := make(chan struct{}, 1)
|
||||
release := make(chan struct{})
|
||||
restore := SetTLSCertificateObtainFuncForTest(func(c *model.TLSCertificate) error {
|
||||
started <- struct{}{}
|
||||
<-release
|
||||
c.CertPEM = newCertPEM
|
||||
c.KeyPEM = newKeyPEM
|
||||
c.NotBefore = time.Now().Add(-time.Hour)
|
||||
c.NotAfter = time.Now().Add(90 * 24 * time.Hour)
|
||||
c.ApplyStatus = "ready"
|
||||
c.ApplyMessage = ""
|
||||
return model.DB.Save(c).Error
|
||||
})
|
||||
t.Cleanup(restore)
|
||||
|
||||
converted, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{
|
||||
Name: "managed-cert",
|
||||
Remark: "converted",
|
||||
AcmeAccountID: 1,
|
||||
DnsAccountID: 2,
|
||||
KeyAlgorithm: "EC256",
|
||||
AutoRenew: true,
|
||||
PrimaryDomain: "managed.example.com",
|
||||
OtherDomains: "www.managed.example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ConvertTLSCertificateToAcme failed: %v", err)
|
||||
}
|
||||
if converted.ID != cert.ID {
|
||||
t.Fatalf("expected converted certificate to keep id %d, got %d", cert.ID, converted.ID)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected conversion obtain task to start")
|
||||
}
|
||||
|
||||
applying, err := model.GetTLSCertificateByID(cert.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("reload applying certificate failed: %v", err)
|
||||
}
|
||||
if applying.Provider != "upload" {
|
||||
t.Fatalf("expected provider to remain upload while applying, got %s", applying.Provider)
|
||||
}
|
||||
if applying.ApplyStatus != "applying" {
|
||||
t.Fatalf("expected applying status, got %s", applying.ApplyStatus)
|
||||
}
|
||||
if applying.CertPEM != originalCertPEM || applying.KeyPEM != originalKeyPEM {
|
||||
t.Fatal("expected original PEM payloads to be preserved while applying")
|
||||
}
|
||||
|
||||
close(release)
|
||||
|
||||
finalCert := waitForCertificateState(t, cert.ID, func(c *model.TLSCertificate) bool {
|
||||
return c.Provider == "acme" && c.ApplyStatus == "ready"
|
||||
})
|
||||
if finalCert.CertPEM != newCertPEM || finalCert.KeyPEM != newKeyPEM {
|
||||
t.Fatal("expected successful conversion to replace PEM payloads")
|
||||
}
|
||||
if !finalCert.AutoRenew {
|
||||
t.Fatal("expected converted certificate to keep auto renew enabled")
|
||||
}
|
||||
if finalCert.PrimaryDomain != "managed.example.com" || finalCert.OtherDomains != "www.managed.example.com" {
|
||||
t.Fatalf("expected converted certificate to persist ACME domains, got %+v", finalCert)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertTLSCertificateToAcmePreservesUploadOnFailure(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
originalCertPEM, originalKeyPEM := generateCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: originalCertPEM,
|
||||
KeyPEM: originalKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
originalCertPEM = cert.CertPEM
|
||||
originalKeyPEM = cert.KeyPEM
|
||||
|
||||
restore := SetTLSCertificateObtainFuncForTest(func(c *model.TLSCertificate) error {
|
||||
err := errors.New("dns challenge failed")
|
||||
updateCertError(c, err.Error())
|
||||
return err
|
||||
})
|
||||
t.Cleanup(restore)
|
||||
|
||||
if _, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{
|
||||
Name: "manual-cert",
|
||||
DnsAccountID: 1,
|
||||
PrimaryDomain: "manual.example.com",
|
||||
}); err != nil {
|
||||
t.Fatalf("ConvertTLSCertificateToAcme failed: %v", err)
|
||||
}
|
||||
|
||||
finalCert := waitForCertificateState(t, cert.ID, func(c *model.TLSCertificate) bool {
|
||||
return c.ApplyStatus == "error"
|
||||
})
|
||||
if finalCert.Provider != "upload" {
|
||||
t.Fatalf("expected failed conversion to keep upload provider, got %s", finalCert.Provider)
|
||||
}
|
||||
if finalCert.CertPEM != originalCertPEM || finalCert.KeyPEM != originalKeyPEM {
|
||||
t.Fatal("expected failed conversion to preserve original PEM payloads")
|
||||
}
|
||||
if !strings.Contains(finalCert.ApplyMessage, "dns challenge failed") {
|
||||
t.Fatalf("expected conversion error message, got %q", finalCert.ApplyMessage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertTLSCertificateToAcmeRejectsInvalidStates(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
certPEM, keyPEM := generateCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
|
||||
cert.Provider = "acme"
|
||||
if err := cert.Update(); err != nil {
|
||||
t.Fatalf("failed to mark certificate acme: %v", err)
|
||||
}
|
||||
if _, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{Name: "manual-cert"}); err == nil || !strings.Contains(err.Error(), "only uploaded") {
|
||||
t.Fatalf("expected non-upload conversion to fail, got %v", err)
|
||||
}
|
||||
|
||||
cert.Provider = "upload"
|
||||
cert.ApplyStatus = "applying"
|
||||
if err := cert.Update(); err != nil {
|
||||
t.Fatalf("failed to mark certificate applying: %v", err)
|
||||
}
|
||||
if _, err := ConvertTLSCertificateToAcme(cert.ID, TLSApplyInput{Name: "manual-cert"}); err == nil || !strings.Contains(err.Error(), "already applying") {
|
||||
t.Fatalf("expected applying conversion to fail, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func waitForCertificateState(t *testing.T, id uint, matches func(*model.TLSCertificate) bool) *model.TLSCertificate {
|
||||
t.Helper()
|
||||
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
t.Fatalf("reload certificate %d failed: %v", id, err)
|
||||
}
|
||||
if matches(cert) {
|
||||
return cert
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
t.Fatalf("reload certificate %d failed: %v", id, err)
|
||||
}
|
||||
t.Fatalf("certificate %d did not reach expected state: %+v", id, cert)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,365 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"strings"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
type TLSCertificateInput struct {
|
||||
Name string `json:"name"`
|
||||
CertPEM string `json:"cert_pem"`
|
||||
KeyPEM string `json:"key_pem"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
type TLSCertificateContent struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
CertPEM string `json:"cert_pem"`
|
||||
KeyPEM string `json:"key_pem"`
|
||||
Remark string `json:"remark"`
|
||||
Provider string `json:"provider"`
|
||||
AcmeAccountID uint `json:"acme_account_id"`
|
||||
DnsAccountID uint `json:"dns_account_id"`
|
||||
KeyAlgorithm string `json:"key_algorithm"`
|
||||
AutoRenew bool `json:"auto_renew"`
|
||||
PrimaryDomain string `json:"primary_domain"`
|
||||
OtherDomains string `json:"other_domains"`
|
||||
DisableCNAME bool `json:"disable_cname"`
|
||||
SkipDNS bool `json:"skip_dns"`
|
||||
DNS1 string `json:"dns1"`
|
||||
DNS2 string `json:"dns2"`
|
||||
ApplyStatus string `json:"apply_status"`
|
||||
ApplyMessage string `json:"apply_message"`
|
||||
}
|
||||
|
||||
type TLSApplyInput struct {
|
||||
Name string `json:"name"`
|
||||
Remark string `json:"remark"`
|
||||
AcmeAccountID uint `json:"acme_account_id"`
|
||||
DnsAccountID uint `json:"dns_account_id"`
|
||||
KeyAlgorithm string `json:"key_algorithm"`
|
||||
AutoRenew bool `json:"auto_renew"`
|
||||
PrimaryDomain string `json:"primary_domain"`
|
||||
OtherDomains string `json:"other_domains"`
|
||||
DisableCNAME bool `json:"disable_cname"`
|
||||
SkipDNS bool `json:"skip_dns"`
|
||||
DNS1 string `json:"dns1"`
|
||||
DNS2 string `json:"dns2"`
|
||||
}
|
||||
|
||||
var obtainTLSCertificate = ObtainSSL
|
||||
|
||||
func SetTLSCertificateObtainFuncForTest(fn func(*model.TLSCertificate) error) func() {
|
||||
previous := obtainTLSCertificate
|
||||
obtainTLSCertificate = fn
|
||||
return func() {
|
||||
obtainTLSCertificate = previous
|
||||
}
|
||||
}
|
||||
|
||||
func ListTLSCertificates() ([]*model.TLSCertificate, error) {
|
||||
return model.ListTLSCertificates()
|
||||
}
|
||||
|
||||
func GetTLSCertificate(id uint) (*model.TLSCertificate, error) {
|
||||
return model.GetTLSCertificateByID(id)
|
||||
}
|
||||
|
||||
func GetTLSCertificateContent(id uint) (*TLSCertificateContent, error) {
|
||||
certificate, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &TLSCertificateContent{
|
||||
ID: certificate.ID,
|
||||
Name: certificate.Name,
|
||||
CertPEM: certificate.CertPEM,
|
||||
KeyPEM: certificate.KeyPEM,
|
||||
Remark: certificate.Remark,
|
||||
Provider: certificate.Provider,
|
||||
AcmeAccountID: certificate.AcmeAccountID,
|
||||
DnsAccountID: certificate.DnsAccountID,
|
||||
KeyAlgorithm: certificate.KeyAlgorithm,
|
||||
AutoRenew: certificate.AutoRenew,
|
||||
PrimaryDomain: certificate.PrimaryDomain,
|
||||
OtherDomains: certificate.OtherDomains,
|
||||
DisableCNAME: certificate.DisableCNAME,
|
||||
SkipDNS: certificate.SkipDNS,
|
||||
DNS1: certificate.DNS1,
|
||||
DNS2: certificate.DNS2,
|
||||
ApplyStatus: certificate.ApplyStatus,
|
||||
ApplyMessage: certificate.ApplyMessage,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func CreateTLSCertificate(input TLSCertificateInput) (*model.TLSCertificate, error) {
|
||||
certificate, err := buildTLSCertificate(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = certificate.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return certificate, nil
|
||||
}
|
||||
|
||||
func CreateTLSCertificateFromFiles(name string, certFile *multipart.FileHeader, keyFile *multipart.FileHeader, remark string) (*model.TLSCertificate, error) {
|
||||
if certFile == nil || keyFile == nil {
|
||||
return nil, errors.New("certificate file and key file cannot be empty")
|
||||
}
|
||||
certContent, err := readMultipartFile(certFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyContent, err := readMultipartFile(keyFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: name,
|
||||
CertPEM: certContent,
|
||||
KeyPEM: keyContent,
|
||||
Remark: remark,
|
||||
})
|
||||
}
|
||||
|
||||
func UpdateTLSCertificate(id uint, input TLSCertificateInput) (*model.TLSCertificate, error) {
|
||||
existing, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
certificate, err := buildTLSCertificate(existing, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = certificate.Update(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return certificate, nil
|
||||
}
|
||||
|
||||
func DeleteTLSCertificate(id uint) error {
|
||||
routes, err := model.ListProxyRoutes()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, route := range routes {
|
||||
if route == nil {
|
||||
continue
|
||||
}
|
||||
if route.CertID != nil && *route.CertID == id {
|
||||
return errors.New("certificate is still referenced by proxy routes")
|
||||
}
|
||||
if strings.TrimSpace(route.CertIDs) == "" {
|
||||
continue
|
||||
}
|
||||
var certIDs []uint
|
||||
if err := json.Unmarshal([]byte(route.CertIDs), &certIDs); err != nil {
|
||||
return fmt.Errorf("proxy route %d cert_ids payload is invalid: %w", route.ID, err)
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if certID == id {
|
||||
return errors.New("certificate is still referenced by proxy routes")
|
||||
}
|
||||
}
|
||||
domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, 0)
|
||||
if err != nil {
|
||||
return fmt.Errorf("proxy route %d domain_cert_ids payload is invalid: %w", route.ID, err)
|
||||
}
|
||||
for _, certID := range domainCertIDs {
|
||||
if certID == id {
|
||||
return errors.New("certificate is still referenced by proxy routes")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
certificate, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
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{
|
||||
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")
|
||||
}
|
||||
|
||||
if err := cert.Insert(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Async obtain SSL
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = obtainTLSCertificate(c)
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func UpdateAcmeCertificate(id uint, input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "acme" {
|
||||
return nil, errors.New("only acme certificates can be updated via this endpoint")
|
||||
}
|
||||
|
||||
fillAcmeCertificateFields(cert, input)
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New("certificate name cannot be empty")
|
||||
}
|
||||
|
||||
if err := cert.Update(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Async obtain SSL with updated config
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = obtainTLSCertificate(c)
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func ConvertTLSCertificateToAcme(id uint, input TLSApplyInput) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "upload" {
|
||||
return nil, errors.New("only uploaded certificates can be converted to acme")
|
||||
}
|
||||
if cert.ApplyStatus == "applying" {
|
||||
return nil, errors.New("certificate is already applying")
|
||||
}
|
||||
|
||||
fillAcmeCertificateFields(cert, input)
|
||||
if cert.Name == "" {
|
||||
return nil, errors.New("certificate name cannot be empty")
|
||||
}
|
||||
cert.ApplyMessage = ""
|
||||
|
||||
if err := cert.Update(); err != nil {
|
||||
if model.IsUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
go func(c *model.TLSCertificate) {
|
||||
if err := obtainTLSCertificate(c); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
latest, err := model.GetTLSCertificateByID(c.ID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
latest.Provider = "acme"
|
||||
latest.ApplyStatus = "ready"
|
||||
latest.ApplyMessage = ""
|
||||
_ = latest.Update()
|
||||
}(cert)
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func RenewTLSCertificate(id uint) (*model.TLSCertificate, error) {
|
||||
cert, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cert.Provider != "acme" {
|
||||
return nil, errors.New("only acme certificates can be renewed")
|
||||
}
|
||||
|
||||
// Async obtain SSL
|
||||
go func(c *model.TLSCertificate) {
|
||||
_ = obtainTLSCertificate(c)
|
||||
}(cert)
|
||||
|
||||
cert.ApplyStatus = "applying"
|
||||
cert.Update()
|
||||
|
||||
return cert, nil
|
||||
}
|
||||
|
||||
func buildTLSCertificate(existing *model.TLSCertificate, input TLSCertificateInput) (*model.TLSCertificate, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
certPEM := strings.TrimSpace(input.CertPEM)
|
||||
keyPEM := strings.TrimSpace(input.KeyPEM)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
if name == "" {
|
||||
return nil, errors.New("certificate name cannot be empty")
|
||||
}
|
||||
if certPEM == "" || keyPEM == "" {
|
||||
return nil, errors.New("certificate content and key content cannot be empty")
|
||||
}
|
||||
parsed, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("certificate or key format is invalid: %w", err)
|
||||
}
|
||||
if len(parsed.Certificate) == 0 {
|
||||
return nil, errors.New("certificate content is invalid")
|
||||
}
|
||||
leaf, err := parseLeafCertificate(certPEM)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.TLSCertificate{}
|
||||
}
|
||||
existing.Name = name
|
||||
existing.CertPEM = certPEM
|
||||
existing.KeyPEM = keyPEM
|
||||
existing.NotBefore = leaf.NotBefore
|
||||
existing.NotAfter = leaf.NotAfter
|
||||
existing.Remark = remark
|
||||
return existing, nil
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
)
|
||||
|
||||
func parseLeafCertificate(certPEM string) (*x509.Certificate, error) {
|
||||
certPEMBlock, _ := pem.Decode([]byte(certPEM))
|
||||
if certPEMBlock == nil {
|
||||
return nil, errors.New("证书 PEM 内容不合法")
|
||||
}
|
||||
leaf, err := x509.ParseCertificate(certPEMBlock.Bytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return leaf, nil
|
||||
}
|
||||
|
||||
func readMultipartFile(fileHeader *multipart.FileHeader) (string, error) {
|
||||
file, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer file.Close()
|
||||
data, err := io.ReadAll(file)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
@@ -0,0 +1,868 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils"
|
||||
)
|
||||
|
||||
const (
|
||||
serverReleaseRepo = "Rain-kl/OpenFlare"
|
||||
githubReleasesAPIBase = "https://api.github.com/repos/%s/releases"
|
||||
)
|
||||
|
||||
type ReleaseChannel string
|
||||
|
||||
const (
|
||||
ReleaseChannelStable ReleaseChannel = "stable"
|
||||
ReleaseChannelPreview ReleaseChannel = "preview"
|
||||
)
|
||||
|
||||
var updateHTTPClient = &http.Client{
|
||||
Timeout: 30 * time.Second,
|
||||
}
|
||||
|
||||
var serverUpgradeState struct {
|
||||
sync.Mutex
|
||||
inProgress bool
|
||||
status string
|
||||
logs []ServerUpgradeLogRecord
|
||||
}
|
||||
|
||||
var serverUpgradeSubscribers struct {
|
||||
sync.Mutex
|
||||
nextID int
|
||||
listeners map[int]chan ServerUpgradeStreamSnapshot
|
||||
}
|
||||
|
||||
var manualServerBinaryState struct {
|
||||
sync.Mutex
|
||||
candidate *manualServerBinaryCandidate
|
||||
}
|
||||
|
||||
var serverBinaryUpgradeExecutor = replaceAndRestartServer
|
||||
|
||||
var serverUpgradeDispatchDelay = 500 * time.Millisecond
|
||||
|
||||
type LatestServerRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Body string `json:"body"`
|
||||
HTMLURL string `json:"html_url"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Channel string `json:"channel"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
HasUpdate bool `json:"has_update"`
|
||||
UpgradeSupported bool `json:"upgrade_supported"`
|
||||
InProgress bool `json:"in_progress"`
|
||||
UpgradeStatus string `json:"upgrade_status"`
|
||||
UpgradeLogs []ServerUpgradeLogRecord `json:"upgrade_logs"`
|
||||
}
|
||||
|
||||
type ServerUpgradeLogRecord struct {
|
||||
Level string `json:"level"`
|
||||
Message string `json:"message"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type ServerUpgradeStreamSnapshot struct {
|
||||
InProgress bool `json:"in_progress"`
|
||||
UpgradeStatus string `json:"upgrade_status"`
|
||||
UpgradeLogs []ServerUpgradeLogRecord `json:"upgrade_logs"`
|
||||
}
|
||||
|
||||
type githubReleaseResponse struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Body string `json:"body"`
|
||||
HTMLURL string `json:"html_url"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
Assets []githubAsset `json:"assets"`
|
||||
}
|
||||
|
||||
type githubAsset struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
}
|
||||
|
||||
type preparedServerUpgrade struct {
|
||||
release *LatestServerRelease
|
||||
downloadURL string
|
||||
execPath string
|
||||
}
|
||||
|
||||
type UploadedServerBinary struct {
|
||||
UploadToken string `json:"upload_token"`
|
||||
FileName string `json:"file_name"`
|
||||
DetectedVersion string `json:"detected_version"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
HasUpdate bool `json:"has_update"`
|
||||
UpgradeSupported bool `json:"upgrade_supported"`
|
||||
ReadyToUpgrade bool `json:"ready_to_upgrade"`
|
||||
ComparisonMessage string `json:"comparison_message"`
|
||||
UploadedAt time.Time `json:"uploaded_at"`
|
||||
}
|
||||
|
||||
type manualServerBinaryCandidate struct {
|
||||
UploadToken string
|
||||
FileName string
|
||||
DetectedVersion string
|
||||
CurrentVersion string
|
||||
TempPath string
|
||||
ExecPath string
|
||||
UploadedAt time.Time
|
||||
}
|
||||
|
||||
func GetLatestServerRelease(ctx context.Context, channel string) (*LatestServerRelease, error) {
|
||||
normalizedChannel := normalizeReleaseChannel(channel)
|
||||
release, err := fetchLatestRelease(ctx, normalizedChannel)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildLatestServerReleaseView(release, normalizedChannel), nil
|
||||
}
|
||||
|
||||
func ScheduleServerUpgrade(channel string) (*LatestServerRelease, error) {
|
||||
normalizedChannel := normalizeReleaseChannel(channel)
|
||||
serverUpgradeState.Lock()
|
||||
if serverUpgradeState.inProgress {
|
||||
serverUpgradeState.Unlock()
|
||||
return nil, fmt.Errorf("服务升级正在执行中,请稍后再试")
|
||||
}
|
||||
|
||||
resetServerUpgradeLogsLocked()
|
||||
serverUpgradeState.status = "running"
|
||||
appendServerUpgradeLogLocked("info", fmt.Sprintf("Automatic upgrade scheduled for channel: %s.", normalizedChannel.String()))
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
|
||||
prepared, err := prepareServerUpgrade(context.Background(), normalizedChannel)
|
||||
if err != nil {
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.status = "failed"
|
||||
appendServerUpgradeLogLocked("error", err.Error())
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = true
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
|
||||
prepared.release.InProgress = true
|
||||
|
||||
go func(task *preparedServerUpgrade) {
|
||||
time.Sleep(serverUpgradeDispatchDelay)
|
||||
if err := executeServerUpgrade(task); err != nil {
|
||||
recordServerUpgradeFailure(err)
|
||||
slog.Error("server self-update failed", "error", err)
|
||||
}
|
||||
}(prepared)
|
||||
|
||||
return prepared.release, nil
|
||||
}
|
||||
|
||||
func UploadManualServerBinary(ctx context.Context, fileName string, reader io.Reader) (*UploadedServerBinary, error) {
|
||||
inProgress, _, _ := snapshotServerUpgradeState()
|
||||
if inProgress {
|
||||
return nil, fmt.Errorf("服务升级正在执行中,请稍后再试")
|
||||
}
|
||||
if strings.TrimSpace(fileName) == "" {
|
||||
return nil, fmt.Errorf("缺少上传文件名")
|
||||
}
|
||||
if reader == nil {
|
||||
return nil, fmt.Errorf("缺少上传文件内容")
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取当前服务程序路径失败: %v", err)
|
||||
}
|
||||
if err = verifyExecutableDirectoryWritable(execPath); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tempPath, err := persistUploadedServerBinary(filepath.Dir(execPath), fileName, reader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
detectedVersion, err := detectUploadedServerBinaryVersion(ctx, tempPath)
|
||||
if err != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
currentVersion := strings.TrimSpace(common.Version)
|
||||
uploadedAt := time.Now()
|
||||
info := buildUploadedServerBinaryView(fileName, currentVersion, detectedVersion, uploadedAt)
|
||||
if !info.ReadyToUpgrade {
|
||||
_ = os.Remove(tempPath)
|
||||
return info, nil
|
||||
}
|
||||
|
||||
uploadToken, err := newUpgradeToken()
|
||||
if err != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return nil, fmt.Errorf("生成升级令牌失败: %v", err)
|
||||
}
|
||||
|
||||
manualServerBinaryState.Lock()
|
||||
cleanupManualServerBinaryCandidateLocked()
|
||||
manualServerBinaryState.candidate = &manualServerBinaryCandidate{
|
||||
UploadToken: uploadToken,
|
||||
FileName: fileName,
|
||||
DetectedVersion: detectedVersion,
|
||||
CurrentVersion: currentVersion,
|
||||
TempPath: tempPath,
|
||||
ExecPath: execPath,
|
||||
UploadedAt: uploadedAt,
|
||||
}
|
||||
manualServerBinaryState.Unlock()
|
||||
|
||||
info.UploadToken = uploadToken
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func ConfirmManualServerUpgrade(uploadToken string) (*UploadedServerBinary, error) {
|
||||
uploadToken = strings.TrimSpace(uploadToken)
|
||||
if uploadToken == "" {
|
||||
return nil, fmt.Errorf("缺少升级令牌")
|
||||
}
|
||||
|
||||
serverUpgradeState.Lock()
|
||||
if serverUpgradeState.inProgress {
|
||||
serverUpgradeState.Unlock()
|
||||
return nil, fmt.Errorf("服务升级正在执行中,请稍后再试")
|
||||
}
|
||||
serverUpgradeState.Unlock()
|
||||
|
||||
manualServerBinaryState.Lock()
|
||||
candidate := manualServerBinaryState.candidate
|
||||
if candidate == nil {
|
||||
manualServerBinaryState.Unlock()
|
||||
return nil, fmt.Errorf("未找到待确认的上传升级包,请重新上传")
|
||||
}
|
||||
if candidate.UploadToken != uploadToken {
|
||||
manualServerBinaryState.Unlock()
|
||||
return nil, fmt.Errorf("升级令牌无效或已过期,请重新上传")
|
||||
}
|
||||
manualServerBinaryState.candidate = nil
|
||||
manualServerBinaryState.Unlock()
|
||||
|
||||
info := buildUploadedServerBinaryView(candidate.FileName, candidate.CurrentVersion, candidate.DetectedVersion, candidate.UploadedAt)
|
||||
info.UploadToken = candidate.UploadToken
|
||||
if !info.ReadyToUpgrade {
|
||||
_ = os.Remove(candidate.TempPath)
|
||||
return nil, fmt.Errorf("当前上传的二进制不满足升级条件")
|
||||
}
|
||||
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = true
|
||||
resetServerUpgradeLogsLocked()
|
||||
serverUpgradeState.status = "running"
|
||||
appendServerUpgradeLogLocked("info", fmt.Sprintf("Manual upgrade confirmed for version: %s.", strings.TrimSpace(candidate.DetectedVersion)))
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
|
||||
go func(task *manualServerBinaryCandidate) {
|
||||
time.Sleep(serverUpgradeDispatchDelay)
|
||||
if err := executeServerBinaryCandidateUpgrade(task, "manual"); err != nil {
|
||||
recordServerUpgradeFailure(err)
|
||||
slog.Error("server manual upgrade failed", "error", err)
|
||||
_ = os.Remove(task.TempPath)
|
||||
}
|
||||
}(candidate)
|
||||
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func fetchLatestRelease(ctx context.Context, channel ReleaseChannel) (*githubReleaseResponse, error) {
|
||||
return fetchLatestGitHubRelease(ctx, serverReleaseRepo, channel)
|
||||
}
|
||||
|
||||
func fetchLatestGitHubRelease(ctx context.Context, repo string, channel ReleaseChannel) (*githubReleaseResponse, error) {
|
||||
switch normalizeReleaseChannel(string(channel)) {
|
||||
case ReleaseChannelPreview:
|
||||
return fetchLatestPreviewGitHubRelease(ctx, repo)
|
||||
default:
|
||||
return fetchLatestStableGitHubRelease(ctx, repo)
|
||||
}
|
||||
}
|
||||
|
||||
func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) {
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/latest", strings.TrimSpace(repo))
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
}
|
||||
|
||||
resp, err := updateHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取最新版本失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
|
||||
return decodeGitHubRelease(resp.Body)
|
||||
}
|
||||
|
||||
func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) {
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"?per_page=20", strings.TrimSpace(repo))
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
}
|
||||
|
||||
resp, err := updateHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取 preview 版本失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
|
||||
var releases []githubReleaseResponse
|
||||
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
return nil, fmt.Errorf("解析 preview 版本信息失败")
|
||||
}
|
||||
for _, release := range releases {
|
||||
if release.Draft || !release.Prerelease {
|
||||
continue
|
||||
}
|
||||
releaseCopy := release
|
||||
return &releaseCopy, nil
|
||||
}
|
||||
return nil, fmt.Errorf("当前没有可用的 preview 发布")
|
||||
}
|
||||
|
||||
func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*githubReleaseResponse, error) {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
return nil, fmt.Errorf("缺少发布版本号")
|
||||
}
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/tags/%s", strings.TrimSpace(repo), tag)
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
}
|
||||
|
||||
resp, err := updateHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取指定版本失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, fmt.Errorf("未找到指定版本: %s", tag)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
|
||||
return decodeGitHubRelease(resp.Body)
|
||||
}
|
||||
|
||||
func newGitHubReleaseRequest(ctx context.Context, url string) (*http.Request, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
req.Header.Set("User-Agent", "OpenFlare-Server")
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func decodeGitHubRelease(reader io.Reader) (*githubReleaseResponse, error) {
|
||||
var release githubReleaseResponse
|
||||
if err := json.NewDecoder(reader).Decode(&release); err != nil {
|
||||
return nil, fmt.Errorf("解析版本信息失败")
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func buildLatestServerReleaseView(release *githubReleaseResponse, channel ReleaseChannel) *LatestServerRelease {
|
||||
currentVersion := strings.TrimSpace(common.Version)
|
||||
isDevBuild := currentVersion == "" || strings.EqualFold(currentVersion, "dev")
|
||||
hasUpdate := false
|
||||
if release != nil && !isDevBuild {
|
||||
if channel == ReleaseChannelPreview {
|
||||
// Preview releases use a "major.minor.patch-git-<commit>" scheme that cannot
|
||||
// be meaningfully compared against the running stable version, so we skip the
|
||||
// version check and always allow upgrading when the user explicitly selects
|
||||
// the preview channel.
|
||||
hasUpdate = true
|
||||
} else {
|
||||
hasUpdate = isVersionNewer(currentVersion, release.TagName)
|
||||
}
|
||||
}
|
||||
|
||||
inProgress, upgradeStatus, upgradeLogs := snapshotServerUpgradeState()
|
||||
|
||||
view := &LatestServerRelease{
|
||||
Channel: channel.String(),
|
||||
CurrentVersion: currentVersion,
|
||||
HasUpdate: hasUpdate,
|
||||
UpgradeSupported: !isDevBuild && runtime.GOOS != "windows",
|
||||
InProgress: inProgress,
|
||||
UpgradeStatus: upgradeStatus,
|
||||
UpgradeLogs: upgradeLogs,
|
||||
}
|
||||
if release != nil {
|
||||
view.TagName = release.TagName
|
||||
view.Body = release.Body
|
||||
view.HTMLURL = release.HTMLURL
|
||||
view.PublishedAt = release.PublishedAt
|
||||
view.Prerelease = release.Prerelease
|
||||
}
|
||||
return view
|
||||
}
|
||||
|
||||
func prepareServerUpgrade(ctx context.Context, channel ReleaseChannel) (*preparedServerUpgrade, error) {
|
||||
release, err := fetchLatestRelease(ctx, channel)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
view := buildLatestServerReleaseView(release, channel)
|
||||
if !view.HasUpdate {
|
||||
return nil, fmt.Errorf("当前已经是最新版本")
|
||||
}
|
||||
if !view.UpgradeSupported {
|
||||
return nil, fmt.Errorf("当前平台暂不支持自动升级")
|
||||
}
|
||||
|
||||
assetName := serverAssetName(runtime.GOOS, runtime.GOARCH)
|
||||
recordServerUpgradeLog("info", fmt.Sprintf("Matching release asset: %s.", assetName))
|
||||
|
||||
var downloadURL string
|
||||
for _, asset := range release.Assets {
|
||||
if asset.Name == assetName {
|
||||
downloadURL = asset.BrowserDownloadURL
|
||||
break
|
||||
}
|
||||
}
|
||||
if downloadURL == "" {
|
||||
return nil, fmt.Errorf("最新版本缺少当前平台的服务端二进制: %s", assetName)
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取当前服务程序路径失败: %v", err)
|
||||
}
|
||||
if err = verifyExecutableDirectoryWritable(execPath); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
recordServerUpgradeLog("info", "Verified current executable directory is writable.")
|
||||
|
||||
return &preparedServerUpgrade{
|
||||
release: view,
|
||||
downloadURL: downloadURL,
|
||||
execPath: execPath,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func verifyExecutableDirectoryWritable(execPath string) error {
|
||||
dir := filepath.Dir(execPath)
|
||||
tempFile, err := os.CreateTemp(dir, "openflare-server-upgrade-check-*")
|
||||
if err != nil {
|
||||
return fmt.Errorf("当前服务二进制目录不可写,无法升级: %v", err)
|
||||
}
|
||||
tempPath := tempFile.Name()
|
||||
if closeErr := tempFile.Close(); closeErr != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return fmt.Errorf("校验服务升级目录失败: %v", closeErr)
|
||||
}
|
||||
if err = os.Remove(tempPath); err != nil {
|
||||
return fmt.Errorf("清理升级校验文件失败: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func executeServerUpgrade(task *preparedServerUpgrade) error {
|
||||
recordServerUpgradeLog("info", fmt.Sprintf("Downloading automatic upgrade package for version: %s.", strings.TrimSpace(task.release.TagName)))
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, task.downloadURL, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Accept", "application/octet-stream")
|
||||
req.Header.Set("User-Agent", "OpenFlare-Server")
|
||||
|
||||
resp, err := updateHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("下载服务端升级包失败: %s", resp.Status)
|
||||
}
|
||||
|
||||
recordServerUpgradeLog("info", "Download finished, validating binary version.")
|
||||
candidate, err := persistDownloadedServerBinary(ctx, task.execPath, task.release.TagName, resp.Body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return executeServerBinaryCandidateUpgrade(candidate, "auto")
|
||||
}
|
||||
|
||||
func executeServerBinaryCandidateUpgrade(task *manualServerBinaryCandidate, source string) error {
|
||||
recordServerUpgradeLog("info", fmt.Sprintf("Validated binary version: %s -> %s.", strings.TrimSpace(task.CurrentVersion), strings.TrimSpace(task.DetectedVersion)))
|
||||
recordServerUpgradeLog("info", "Replacing executable and preparing restart.")
|
||||
if source == "manual" {
|
||||
slog.Info("server manual self-update starting", "from", strings.TrimSpace(task.CurrentVersion), "to", strings.TrimSpace(task.DetectedVersion))
|
||||
} else {
|
||||
slog.Info("server self-update starting", "from", strings.TrimSpace(task.CurrentVersion), "to", strings.TrimSpace(task.DetectedVersion))
|
||||
}
|
||||
markServerUpgradeSucceeded()
|
||||
return serverBinaryUpgradeExecutor(task.ExecPath, task.TempPath)
|
||||
}
|
||||
|
||||
func serverAssetName(goos string, goarch string) string {
|
||||
name := fmt.Sprintf("openflare-server-%s-%s", goos, goarch)
|
||||
if goos == "windows" {
|
||||
return name + ".exe"
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func normalizeReleaseChannel(channel string) ReleaseChannel {
|
||||
switch strings.ToLower(strings.TrimSpace(channel)) {
|
||||
case string(ReleaseChannelPreview):
|
||||
return ReleaseChannelPreview
|
||||
default:
|
||||
return ReleaseChannelStable
|
||||
}
|
||||
}
|
||||
|
||||
func (channel ReleaseChannel) String() string {
|
||||
if channel == ReleaseChannelPreview {
|
||||
return string(ReleaseChannelPreview)
|
||||
}
|
||||
return string(ReleaseChannelStable)
|
||||
}
|
||||
|
||||
func isVersionNewer(current string, latest string) bool {
|
||||
return utils.CompareVersions(current, latest) < 0
|
||||
}
|
||||
|
||||
func buildUploadedServerBinaryView(fileName string, currentVersion string, detectedVersion string, uploadedAt time.Time) *UploadedServerBinary {
|
||||
upgradeSupported := isManualServerUpgradeSupported(currentVersion)
|
||||
hasUpdate := false
|
||||
comparisonMessage := ""
|
||||
|
||||
switch {
|
||||
case !upgradeSupported:
|
||||
comparisonMessage = "当前服务版本不支持手动升级确认流程"
|
||||
case normalizeVersion(currentVersion) == normalizeVersion(detectedVersion):
|
||||
comparisonMessage = "上传二进制与当前服务版本一致,无需升级"
|
||||
case isVersionNewer(currentVersion, detectedVersion):
|
||||
hasUpdate = true
|
||||
comparisonMessage = fmt.Sprintf("检测到可升级版本:%s -> %s", strings.TrimSpace(currentVersion), strings.TrimSpace(detectedVersion))
|
||||
default:
|
||||
comparisonMessage = "上传二进制版本不高于当前服务版本,已拒绝升级"
|
||||
}
|
||||
|
||||
return &UploadedServerBinary{
|
||||
FileName: strings.TrimSpace(fileName),
|
||||
DetectedVersion: strings.TrimSpace(detectedVersion),
|
||||
CurrentVersion: strings.TrimSpace(currentVersion),
|
||||
HasUpdate: hasUpdate,
|
||||
UpgradeSupported: upgradeSupported,
|
||||
ReadyToUpgrade: upgradeSupported && hasUpdate,
|
||||
ComparisonMessage: comparisonMessage,
|
||||
UploadedAt: uploadedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func isManualServerUpgradeSupported(currentVersion string) bool {
|
||||
normalized := strings.TrimSpace(strings.TrimPrefix(currentVersion, "v"))
|
||||
return normalized != "" && !strings.EqualFold(normalized, "dev")
|
||||
}
|
||||
|
||||
func persistUploadedServerBinary(tempDir string, fileName string, reader io.Reader) (string, error) {
|
||||
suffix := filepath.Ext(strings.TrimSpace(fileName))
|
||||
if runtime.GOOS == "windows" && suffix == "" {
|
||||
suffix = ".exe"
|
||||
}
|
||||
tempDir = strings.TrimSpace(tempDir)
|
||||
if tempDir == "" {
|
||||
tempDir = os.TempDir()
|
||||
}
|
||||
tempFile, err := os.CreateTemp(tempDir, "openflare-server-manual-upgrade-*"+suffix)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("创建临时升级文件失败: %v", err)
|
||||
}
|
||||
tempPath := tempFile.Name()
|
||||
if _, err = io.Copy(tempFile, reader); err != nil {
|
||||
_ = tempFile.Close()
|
||||
_ = os.Remove(tempPath)
|
||||
return "", fmt.Errorf("写入上传二进制失败: %v", err)
|
||||
}
|
||||
if err = tempFile.Close(); err != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return "", fmt.Errorf("关闭临时升级文件失败: %v", err)
|
||||
}
|
||||
if err = os.Chmod(tempPath, 0o755); err != nil && runtime.GOOS != "windows" {
|
||||
_ = os.Remove(tempPath)
|
||||
return "", fmt.Errorf("设置临时升级文件权限失败: %v", err)
|
||||
}
|
||||
return tempPath, nil
|
||||
}
|
||||
|
||||
func detectUploadedServerBinaryVersion(ctx context.Context, filePath string) (string, error) {
|
||||
commandCtx := ctx
|
||||
if commandCtx == nil {
|
||||
commandCtx = context.Background()
|
||||
}
|
||||
cmd := exec.CommandContext(commandCtx, filePath, "--version")
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("检查上传二进制版本失败: %w: %s", err, strings.TrimSpace(string(output)))
|
||||
}
|
||||
version := strings.TrimSpace(string(output))
|
||||
if version == "" {
|
||||
return "", fmt.Errorf("上传二进制未返回有效版本号")
|
||||
}
|
||||
for _, line := range strings.Split(version, "\n") {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if trimmed != "" {
|
||||
return trimmed, nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("上传二进制未返回有效版本号")
|
||||
}
|
||||
|
||||
func persistDownloadedServerBinary(ctx context.Context, execPath string, releaseTag string, reader io.Reader) (*manualServerBinaryCandidate, error) {
|
||||
fileName := serverAssetName(runtime.GOOS, runtime.GOARCH)
|
||||
tempPath, err := persistUploadedServerBinary(filepath.Dir(execPath), fileName, reader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
detectedVersion, err := detectUploadedServerBinaryVersion(ctx, tempPath)
|
||||
if err != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return nil, err
|
||||
}
|
||||
recordServerUpgradeLog("info", fmt.Sprintf("Detected downloaded binary version: %s.", strings.TrimSpace(detectedVersion)))
|
||||
|
||||
if normalizeVersion(detectedVersion) != normalizeVersion(releaseTag) {
|
||||
_ = os.Remove(tempPath)
|
||||
return nil, fmt.Errorf("下载包版本校验失败:release=%s,binary=%s", strings.TrimSpace(releaseTag), strings.TrimSpace(detectedVersion))
|
||||
}
|
||||
|
||||
info := buildUploadedServerBinaryView(fileName, common.Version, detectedVersion, time.Now())
|
||||
if !info.ReadyToUpgrade {
|
||||
_ = os.Remove(tempPath)
|
||||
return nil, errors.New(info.ComparisonMessage)
|
||||
}
|
||||
|
||||
return &manualServerBinaryCandidate{
|
||||
FileName: fileName,
|
||||
DetectedVersion: detectedVersion,
|
||||
CurrentVersion: strings.TrimSpace(common.Version),
|
||||
TempPath: tempPath,
|
||||
ExecPath: execPath,
|
||||
UploadedAt: time.Now(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func cleanupManualServerBinaryCandidateLocked() {
|
||||
if manualServerBinaryState.candidate == nil {
|
||||
return
|
||||
}
|
||||
_ = os.Remove(manualServerBinaryState.candidate.TempPath)
|
||||
manualServerBinaryState.candidate = nil
|
||||
}
|
||||
|
||||
func newUpgradeToken() (string, error) {
|
||||
buffer := make([]byte, 16)
|
||||
if _, err := rand.Read(buffer); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buffer), nil
|
||||
}
|
||||
|
||||
func normalizeVersion(version string) string {
|
||||
return strings.TrimSpace(strings.TrimPrefix(version, "v"))
|
||||
}
|
||||
|
||||
func snapshotServerUpgradeState() (bool, string, []ServerUpgradeLogRecord) {
|
||||
serverUpgradeState.Lock()
|
||||
defer serverUpgradeState.Unlock()
|
||||
|
||||
status := strings.TrimSpace(serverUpgradeState.status)
|
||||
if status == "" {
|
||||
status = "idle"
|
||||
}
|
||||
logs := make([]ServerUpgradeLogRecord, len(serverUpgradeState.logs))
|
||||
copy(logs, serverUpgradeState.logs)
|
||||
return serverUpgradeState.inProgress, status, logs
|
||||
}
|
||||
|
||||
func snapshotServerUpgradeStream() ServerUpgradeStreamSnapshot {
|
||||
inProgress, status, logs := snapshotServerUpgradeState()
|
||||
return ServerUpgradeStreamSnapshot{
|
||||
InProgress: inProgress,
|
||||
UpgradeStatus: status,
|
||||
UpgradeLogs: logs,
|
||||
}
|
||||
}
|
||||
|
||||
func resetServerUpgradeLogsLocked() {
|
||||
serverUpgradeState.logs = nil
|
||||
}
|
||||
|
||||
func appendServerUpgradeLogLocked(level string, message string) {
|
||||
serverUpgradeState.logs = append(serverUpgradeState.logs, ServerUpgradeLogRecord{
|
||||
Level: strings.TrimSpace(level),
|
||||
Message: strings.TrimSpace(message),
|
||||
CreatedAt: time.Now(),
|
||||
})
|
||||
if len(serverUpgradeState.logs) > 100 {
|
||||
serverUpgradeState.logs = append([]ServerUpgradeLogRecord(nil), serverUpgradeState.logs[len(serverUpgradeState.logs)-100:]...)
|
||||
}
|
||||
}
|
||||
|
||||
func recordServerUpgradeLog(level string, message string) {
|
||||
serverUpgradeState.Lock()
|
||||
appendServerUpgradeLogLocked(level, message)
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
}
|
||||
|
||||
func markServerUpgradeSucceeded() {
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = false
|
||||
serverUpgradeState.status = "succeeded"
|
||||
appendServerUpgradeLogLocked("info", "Upgrade binary is ready; server restart will begin.")
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
}
|
||||
|
||||
func recordServerUpgradeFailure(err error) {
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = false
|
||||
serverUpgradeState.status = "failed"
|
||||
if err != nil {
|
||||
appendServerUpgradeLogLocked("error", err.Error())
|
||||
}
|
||||
serverUpgradeState.Unlock()
|
||||
broadcastServerUpgradeSnapshot()
|
||||
}
|
||||
|
||||
func SubscribeServerUpgradeStream() (<-chan ServerUpgradeStreamSnapshot, func()) {
|
||||
serverUpgradeSubscribers.Lock()
|
||||
if serverUpgradeSubscribers.listeners == nil {
|
||||
serverUpgradeSubscribers.listeners = make(map[int]chan ServerUpgradeStreamSnapshot)
|
||||
}
|
||||
serverUpgradeSubscribers.nextID++
|
||||
listenerID := serverUpgradeSubscribers.nextID
|
||||
listener := make(chan ServerUpgradeStreamSnapshot, 8)
|
||||
serverUpgradeSubscribers.listeners[listenerID] = listener
|
||||
serverUpgradeSubscribers.Unlock()
|
||||
|
||||
listener <- snapshotServerUpgradeStream()
|
||||
|
||||
unsubscribe := func() {
|
||||
serverUpgradeSubscribers.Lock()
|
||||
ch, ok := serverUpgradeSubscribers.listeners[listenerID]
|
||||
if ok {
|
||||
delete(serverUpgradeSubscribers.listeners, listenerID)
|
||||
}
|
||||
serverUpgradeSubscribers.Unlock()
|
||||
if ok {
|
||||
close(ch)
|
||||
}
|
||||
}
|
||||
|
||||
return listener, unsubscribe
|
||||
}
|
||||
|
||||
func broadcastServerUpgradeSnapshot() {
|
||||
snapshot := snapshotServerUpgradeStream()
|
||||
|
||||
serverUpgradeSubscribers.Lock()
|
||||
if len(serverUpgradeSubscribers.listeners) == 0 {
|
||||
serverUpgradeSubscribers.Unlock()
|
||||
return
|
||||
}
|
||||
listeners := make([]chan ServerUpgradeStreamSnapshot, 0, len(serverUpgradeSubscribers.listeners))
|
||||
for _, listener := range serverUpgradeSubscribers.listeners {
|
||||
listeners = append(listeners, listener)
|
||||
}
|
||||
serverUpgradeSubscribers.Unlock()
|
||||
|
||||
for _, listener := range listeners {
|
||||
select {
|
||||
case listener <- snapshot:
|
||||
default:
|
||||
select {
|
||||
case <-listener:
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case listener <- snapshot:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func UpdateHTTPClientForTest() *http.Client {
|
||||
return updateHTTPClient
|
||||
}
|
||||
|
||||
func SetUpdateHTTPClientForTest(client *http.Client) {
|
||||
updateHTTPClient = client
|
||||
}
|
||||
|
||||
func ServerBinaryUpgradeExecutorForTest() func(string, string) error {
|
||||
return serverBinaryUpgradeExecutor
|
||||
}
|
||||
|
||||
func SetServerBinaryUpgradeExecutorForTest(executor func(string, string) error) {
|
||||
if executor == nil {
|
||||
serverBinaryUpgradeExecutor = replaceAndRestartServer
|
||||
return
|
||||
}
|
||||
serverBinaryUpgradeExecutor = executor
|
||||
}
|
||||
|
||||
func ServerUpgradeDispatchDelayForTest() time.Duration {
|
||||
return serverUpgradeDispatchDelay
|
||||
}
|
||||
|
||||
func SetServerUpgradeDispatchDelayForTest(delay time.Duration) {
|
||||
if delay < 0 {
|
||||
delay = 0
|
||||
}
|
||||
serverUpgradeDispatchDelay = delay
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
//go:build !windows
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
var unixRename = os.Rename
|
||||
|
||||
func replaceAndRestartServer(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
_ = os.Remove(backupPath)
|
||||
if err := unixRename(execPath, backupPath); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("备份当前服务端二进制失败: %w", err)
|
||||
}
|
||||
if err := replaceFileUnix(tmpPath, execPath); err != nil {
|
||||
_ = unixRename(backupPath, execPath)
|
||||
return fmt.Errorf("替换服务端二进制失败: %w", err)
|
||||
}
|
||||
_ = os.Remove(backupPath)
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil {
|
||||
return fmt.Errorf("重启服务失败: %w", err)
|
||||
}
|
||||
return fmt.Errorf("unreachable after exec")
|
||||
}
|
||||
|
||||
func replaceFileUnix(srcPath string, dstPath string) error {
|
||||
if err := unixRename(srcPath, dstPath); err == nil {
|
||||
return nil
|
||||
} else if linkErr, ok := err.(*os.LinkError); !ok || linkErr.Err != syscall.EXDEV {
|
||||
return err
|
||||
}
|
||||
|
||||
sourceFile, err := os.Open(srcPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer sourceFile.Close()
|
||||
|
||||
info, err := sourceFile.Stat()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
destinationFile, err := os.OpenFile(dstPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
copyErr := func() error {
|
||||
defer destinationFile.Close()
|
||||
if _, err = io.Copy(destinationFile, sourceFile); err != nil {
|
||||
return err
|
||||
}
|
||||
if err = destinationFile.Sync(); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}()
|
||||
if copyErr != nil {
|
||||
return copyErr
|
||||
}
|
||||
|
||||
if err = os.Chmod(dstPath, info.Mode().Perm()); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Remove(srcPath)
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
//go:build !windows
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestReplaceFileUnixFallsBackOnCrossDeviceRename(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
srcPath := filepath.Join(tempDir, "source.bin")
|
||||
dstPath := filepath.Join(tempDir, "target.bin")
|
||||
|
||||
if err := os.WriteFile(srcPath, []byte("new-binary"), 0o755); err != nil {
|
||||
t.Fatalf("failed to write source file: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(dstPath, []byte("old-binary"), 0o755); err != nil {
|
||||
t.Fatalf("failed to write target file: %v", err)
|
||||
}
|
||||
|
||||
originalRename := unixRename
|
||||
unixRename = func(oldPath string, newPath string) error {
|
||||
if oldPath == srcPath && newPath == dstPath {
|
||||
return &os.LinkError{Op: "rename", Old: oldPath, New: newPath, Err: syscall.EXDEV}
|
||||
}
|
||||
return os.Rename(oldPath, newPath)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
unixRename = originalRename
|
||||
})
|
||||
|
||||
if err := replaceFileUnix(srcPath, dstPath); err != nil {
|
||||
t.Fatalf("expected cross-device fallback to succeed: %v", err)
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(dstPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read target file: %v", err)
|
||||
}
|
||||
if string(content) != "new-binary" {
|
||||
t.Fatalf("unexpected target content: %s", string(content))
|
||||
}
|
||||
if _, err = os.Stat(srcPath); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("expected source file to be removed, got err=%v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
//go:build windows
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func replaceAndRestartServer(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
scriptPath := execPath + ".update.cmd"
|
||||
script := fmt.Sprintf(`@echo off
|
||||
setlocal
|
||||
:waitloop
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 (
|
||||
ping 127.0.0.1 -n 2 >nul
|
||||
goto waitloop
|
||||
)
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 exit /b 1
|
||||
start "" %s
|
||||
del /Q "%s" >nul 2>nul
|
||||
del /Q "%%~f0" >nul 2>nul
|
||||
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
|
||||
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("写入升级重启脚本失败: %w", err)
|
||||
}
|
||||
|
||||
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
|
||||
if err := cmd.Start(); err != nil {
|
||||
_ = os.Remove(scriptPath)
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("调度升级重启失败: %w", err)
|
||||
}
|
||||
|
||||
os.Exit(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildWindowsCommandLine(execPath string, args []string) string {
|
||||
parts := []string{quoteWindowsArg(execPath)}
|
||||
for _, arg := range args {
|
||||
parts = append(parts, quoteWindowsArg(arg))
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func quoteWindowsArg(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -0,0 +1,415 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
)
|
||||
|
||||
type serverUpdateRoundTripFunc func(req *http.Request) (*http.Response, error)
|
||||
|
||||
func (f serverUpdateRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func resetServerUpgradeTestState(t *testing.T) {
|
||||
t.Helper()
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = false
|
||||
serverUpgradeState.status = ""
|
||||
serverUpgradeState.logs = nil
|
||||
serverUpgradeState.Unlock()
|
||||
manualServerBinaryState.Lock()
|
||||
cleanupManualServerBinaryCandidateLocked()
|
||||
manualServerBinaryState.Unlock()
|
||||
}
|
||||
|
||||
func fakeServerBinaryFixture(version string) (string, []byte) {
|
||||
if runtime.GOOS == "windows" {
|
||||
return "openflare-server-test.cmd", []byte("@echo off\r\necho " + version + "\r\n")
|
||||
}
|
||||
return "openflare-server-test.sh", []byte("#!/bin/sh\necho " + version + "\n")
|
||||
}
|
||||
|
||||
func TestIsVersionNewer(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
current string
|
||||
latest string
|
||||
expected bool
|
||||
}{
|
||||
{name: "newer patch", current: "v1.2.3", latest: "v1.2.4", expected: true},
|
||||
{name: "same version", current: "v1.2.3", latest: "v1.2.3", expected: false},
|
||||
{name: "older remote", current: "v1.3.0", latest: "v1.2.9", expected: false},
|
||||
{name: "double digit segment", current: "v1.9.9", latest: "v1.10.0", expected: true},
|
||||
{name: "stable newer than prerelease", current: "v1.2.3-rc.1", latest: "v1.2.3", expected: true},
|
||||
{name: "prerelease not newer than same stable", current: "v1.2.3", latest: "v1.2.3-rc.1", expected: false},
|
||||
{name: "newer prerelease sequence", current: "v1.2.3-rc.1", latest: "v1.2.3-rc.2", expected: true},
|
||||
{name: "git describe newer than same tag", current: "v0.6.3", latest: "v0.6.3-2-gf4d36be", expected: true},
|
||||
{name: "git describe distance compares numerically", current: "v0.6.3-2-gf4d36be", latest: "v0.6.3-5-gabc1234", expected: true},
|
||||
{name: "dev build", current: "dev", latest: "v0.4.0", expected: true},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
actual := isVersionNewer(testCase.current, testCase.latest)
|
||||
if actual != testCase.expected {
|
||||
t.Fatalf("unexpected compare result: current=%s latest=%s actual=%v expected=%v", testCase.current, testCase.latest, actual, testCase.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildLatestServerReleaseView(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "v0.4.0"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = false
|
||||
serverUpgradeState.Unlock()
|
||||
})
|
||||
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = true
|
||||
serverUpgradeState.Unlock()
|
||||
|
||||
view := buildLatestServerReleaseView(&githubReleaseResponse{
|
||||
TagName: "v0.5.0",
|
||||
Body: "release notes",
|
||||
HTMLURL: "https://github.com/Rain-kl/OpenFlare/releases/tag/v0.5.0",
|
||||
PublishedAt: "2026-03-11T00:00:00Z",
|
||||
}, ReleaseChannelStable)
|
||||
|
||||
if view.CurrentVersion != "v0.4.0" {
|
||||
t.Fatalf("unexpected current version: %s", view.CurrentVersion)
|
||||
}
|
||||
if !view.HasUpdate {
|
||||
t.Fatal("expected has_update to be true")
|
||||
}
|
||||
if !view.InProgress {
|
||||
t.Fatal("expected in_progress to reflect upgrade state")
|
||||
}
|
||||
if view.TagName != "v0.5.0" {
|
||||
t.Fatalf("unexpected tag name: %s", view.TagName)
|
||||
}
|
||||
if view.Channel != ReleaseChannelStable.String() {
|
||||
t.Fatalf("unexpected channel: %s", view.Channel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildLatestServerReleaseViewDevBuild(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "dev"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = false
|
||||
serverUpgradeState.Unlock()
|
||||
})
|
||||
|
||||
view := buildLatestServerReleaseView(&githubReleaseResponse{
|
||||
TagName: "v0.5.0",
|
||||
}, ReleaseChannelStable)
|
||||
|
||||
if view.HasUpdate {
|
||||
t.Fatal("expected dev build not to report update availability")
|
||||
}
|
||||
if view.UpgradeSupported {
|
||||
t.Fatal("expected dev build not to support self-upgrade")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildLatestServerReleaseViewPreview(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "v0.5.0-rc.1"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
view := buildLatestServerReleaseView(&githubReleaseResponse{
|
||||
TagName: "v0.5.0-rc.2",
|
||||
Prerelease: true,
|
||||
PublishedAt: "2026-03-12T00:00:00Z",
|
||||
}, ReleaseChannelPreview)
|
||||
|
||||
if !view.HasUpdate {
|
||||
t.Fatal("expected preview release to be newer")
|
||||
}
|
||||
if !view.Prerelease {
|
||||
t.Fatal("expected preview flag to be true")
|
||||
}
|
||||
if view.Channel != ReleaseChannelPreview.String() {
|
||||
t.Fatalf("unexpected channel: %s", view.Channel)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildLatestServerReleaseViewPreviewBypassVersionCheck verifies that switching to
|
||||
// the preview channel always reports has_update=true, even when the preview tag uses a
|
||||
// "major.minor.patch-git-<commit>" scheme that would otherwise compare as equal-or-older
|
||||
// than the currently running stable version.
|
||||
func TestBuildLatestServerReleaseViewPreviewBypassVersionCheck(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "v1.0.0"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
// A typical preview tag: same base version as stable but with a git-commit suffix.
|
||||
// Without the bypass, isVersionNewer("v1.0.0", "v1.0.0-git-abc1234") returns false
|
||||
// because a version without a prerelease identifier is considered higher than one
|
||||
// with a prerelease identifier under semver rules.
|
||||
view := buildLatestServerReleaseView(&githubReleaseResponse{
|
||||
TagName: "v1.0.0-git-abc1234",
|
||||
Prerelease: true,
|
||||
PublishedAt: "2026-03-12T00:00:00Z",
|
||||
}, ReleaseChannelPreview)
|
||||
|
||||
if !view.HasUpdate {
|
||||
t.Fatal("expected preview channel to bypass version comparison and report has_update=true")
|
||||
}
|
||||
if view.Channel != ReleaseChannelPreview.String() {
|
||||
t.Fatalf("unexpected channel: %s", view.Channel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadManualServerBinary(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "v0.4.0"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
fileName, content := fakeServerBinaryFixture("v0.5.0")
|
||||
info, err := UploadManualServerBinary(context.Background(), fileName, bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("expected upload to succeed: %v", err)
|
||||
}
|
||||
if !info.ReadyToUpgrade {
|
||||
t.Fatal("expected uploaded binary to be ready for upgrade")
|
||||
}
|
||||
if info.UploadToken == "" {
|
||||
t.Fatal("expected upload token to be returned")
|
||||
}
|
||||
if info.DetectedVersion != "v0.5.0" {
|
||||
t.Fatalf("unexpected detected version: %s", info.DetectedVersion)
|
||||
}
|
||||
|
||||
manualServerBinaryState.Lock()
|
||||
candidate := manualServerBinaryState.candidate
|
||||
manualServerBinaryState.Unlock()
|
||||
if candidate == nil {
|
||||
t.Fatal("expected manual upgrade candidate to be stored")
|
||||
}
|
||||
if _, err := os.Stat(candidate.TempPath); err != nil {
|
||||
t.Fatalf("expected temporary binary to exist: %v", err)
|
||||
}
|
||||
if candidate.UploadToken != info.UploadToken {
|
||||
t.Fatalf("unexpected stored upload token: %s", candidate.UploadToken)
|
||||
}
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get executable path: %v", err)
|
||||
}
|
||||
if filepath.Dir(candidate.TempPath) != filepath.Dir(execPath) {
|
||||
t.Fatalf("expected temporary binary in executable dir, got %s want %s", filepath.Dir(candidate.TempPath), filepath.Dir(execPath))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildUploadedServerBinaryViewAcceptsGitDescribeNewerThanTag(t *testing.T) {
|
||||
info := buildUploadedServerBinaryView("openflare-server-test", "v0.6.3", "v0.6.3-2-gf4d36be", time.Now())
|
||||
if !info.HasUpdate || !info.ReadyToUpgrade {
|
||||
t.Fatalf("expected git describe binary to be upgradeable: %+v", info)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadManualServerBinaryRejectsSameVersion(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "v0.5.0"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
fileName, content := fakeServerBinaryFixture("v0.5.0")
|
||||
info, err := UploadManualServerBinary(context.Background(), fileName, bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("expected upload to succeed: %v", err)
|
||||
}
|
||||
if info.ReadyToUpgrade {
|
||||
t.Fatal("expected same-version upload not to be upgradeable")
|
||||
}
|
||||
if info.UploadToken != "" {
|
||||
t.Fatal("expected same-version upload not to issue a token")
|
||||
}
|
||||
|
||||
manualServerBinaryState.Lock()
|
||||
defer manualServerBinaryState.Unlock()
|
||||
if manualServerBinaryState.candidate != nil {
|
||||
t.Fatal("expected no pending manual upgrade candidate")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfirmManualServerUpgrade(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
originalExecutor := ServerBinaryUpgradeExecutorForTest()
|
||||
originalDelay := ServerUpgradeDispatchDelayForTest()
|
||||
common.Version = "v0.4.0"
|
||||
called := make(chan string, 1)
|
||||
SetServerBinaryUpgradeExecutorForTest(func(execPath string, tempPath string) error {
|
||||
called <- tempPath
|
||||
return nil
|
||||
})
|
||||
SetServerUpgradeDispatchDelayForTest(0)
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
SetServerBinaryUpgradeExecutorForTest(originalExecutor)
|
||||
SetServerUpgradeDispatchDelayForTest(originalDelay)
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
fileName, content := fakeServerBinaryFixture("v0.5.0")
|
||||
info, err := UploadManualServerBinary(context.Background(), fileName, bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("expected upload to succeed: %v", err)
|
||||
}
|
||||
|
||||
confirmed, err := ConfirmManualServerUpgrade(info.UploadToken)
|
||||
if err != nil {
|
||||
t.Fatalf("expected confirm to succeed: %v", err)
|
||||
}
|
||||
if confirmed.UploadToken != info.UploadToken {
|
||||
t.Fatalf("unexpected confirmed upload token: %s", confirmed.UploadToken)
|
||||
}
|
||||
|
||||
select {
|
||||
case tempPath := <-called:
|
||||
if tempPath == "" {
|
||||
t.Fatal("expected upgrade executor to receive temp path")
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected manual upgrade executor to be called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildLatestServerReleaseViewIncludesUpgradeLogs(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
common.Version = "v0.4.0"
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
serverUpgradeState.Lock()
|
||||
serverUpgradeState.inProgress = true
|
||||
serverUpgradeState.status = "running"
|
||||
serverUpgradeState.logs = []ServerUpgradeLogRecord{
|
||||
{
|
||||
Level: "info",
|
||||
Message: "download started",
|
||||
CreatedAt: time.Now(),
|
||||
},
|
||||
}
|
||||
serverUpgradeState.Unlock()
|
||||
|
||||
view := buildLatestServerReleaseView(&githubReleaseResponse{
|
||||
TagName: "v0.5.0",
|
||||
}, ReleaseChannelStable)
|
||||
|
||||
if view.UpgradeStatus != "running" {
|
||||
t.Fatalf("expected upgrade status to be running, got %s", view.UpgradeStatus)
|
||||
}
|
||||
if len(view.UpgradeLogs) != 1 {
|
||||
t.Fatalf("expected one upgrade log, got %d", len(view.UpgradeLogs))
|
||||
}
|
||||
if view.UpgradeLogs[0].Message != "download started" {
|
||||
t.Fatalf("unexpected upgrade log message: %s", view.UpgradeLogs[0].Message)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScheduleServerUpgradeUsesDownloadedBinaryValidation(t *testing.T) {
|
||||
originalVersion := common.Version
|
||||
originalClient := UpdateHTTPClientForTest()
|
||||
originalExecutor := ServerBinaryUpgradeExecutorForTest()
|
||||
originalDelay := ServerUpgradeDispatchDelayForTest()
|
||||
common.Version = "v0.4.0"
|
||||
called := make(chan string, 1)
|
||||
|
||||
SetUpdateHTTPClientForTest(&http.Client{
|
||||
Transport: serverUpdateRoundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
switch req.URL.String() {
|
||||
case "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest":
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{
|
||||
"tag_name":"v0.5.0",
|
||||
"body":"release notes",
|
||||
"html_url":"https://github.com/Rain-kl/OpenFlare/releases/tag/v0.5.0",
|
||||
"published_at":"2026-03-11T00:00:00Z",
|
||||
"assets":[{"name":"openflare-server-` + runtime.GOOS + `-` + runtime.GOARCH + `","browser_download_url":"https://downloads.example.com/openflare-server"}]
|
||||
}`)),
|
||||
}, nil
|
||||
case "https://downloads.example.com/openflare-server":
|
||||
_, content := fakeServerBinaryFixture("v0.5.0")
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(bytes.NewReader(content)),
|
||||
}, nil
|
||||
default:
|
||||
t.Fatalf("unexpected request url: %s", req.URL.String())
|
||||
return nil, nil
|
||||
}
|
||||
}),
|
||||
})
|
||||
SetServerBinaryUpgradeExecutorForTest(func(execPath string, tempPath string) error {
|
||||
called <- tempPath
|
||||
return nil
|
||||
})
|
||||
SetServerUpgradeDispatchDelayForTest(0)
|
||||
t.Cleanup(func() {
|
||||
common.Version = originalVersion
|
||||
SetUpdateHTTPClientForTest(originalClient)
|
||||
SetServerBinaryUpgradeExecutorForTest(originalExecutor)
|
||||
SetServerUpgradeDispatchDelayForTest(originalDelay)
|
||||
resetServerUpgradeTestState(t)
|
||||
})
|
||||
|
||||
release, err := ScheduleServerUpgrade("stable")
|
||||
if err != nil {
|
||||
t.Fatalf("expected schedule to succeed: %v", err)
|
||||
}
|
||||
if !release.InProgress {
|
||||
t.Fatal("expected release to report in-progress upgrade")
|
||||
}
|
||||
|
||||
select {
|
||||
case tempPath := <-called:
|
||||
if tempPath == "" {
|
||||
t.Fatal("expected upgrade executor to receive temp path")
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected automatic upgrade executor to be called")
|
||||
}
|
||||
|
||||
_, status, logs := snapshotServerUpgradeState()
|
||||
if status != "succeeded" {
|
||||
t.Fatalf("expected succeeded status after executor call, got %s", status)
|
||||
}
|
||||
if len(logs) == 0 {
|
||||
t.Fatal("expected upgrade logs to be recorded")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,305 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/common"
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/uptimekuma"
|
||||
)
|
||||
|
||||
var isSyncing atomic.Bool
|
||||
|
||||
func SyncToUptimeKuma() error {
|
||||
if !common.UptimeKumaEnabled {
|
||||
return fmt.Errorf("Uptime Kuma integration is disabled")
|
||||
}
|
||||
|
||||
if !isSyncing.CompareAndSwap(false, true) {
|
||||
return fmt.Errorf("sync task is already in progress, please try again later")
|
||||
}
|
||||
defer isSyncing.Store(false)
|
||||
|
||||
kumaUrl := strings.TrimSpace(common.UptimeKumaUrl)
|
||||
kumaUsername := strings.TrimSpace(common.UptimeKumaUsername)
|
||||
kumaPassword := strings.TrimSpace(common.UptimeKumaPassword)
|
||||
if kumaUrl == "" || kumaUsername == "" || kumaPassword == "" {
|
||||
return fmt.Errorf("Uptime Kuma URL, username, or password is not configured (URL: %q, Username: %q, PasswordLength: %d)", kumaUrl, kumaUsername, len(kumaPassword))
|
||||
}
|
||||
|
||||
slog.Info("Starting Uptime Kuma sync process", "url", kumaUrl, "username", kumaUsername, "scope", common.UptimeKumaMonitorScope)
|
||||
|
||||
// 1. Fetch expected sites
|
||||
allRoutes, err := model.ListProxyRoutes()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to list local proxy routes: %w", err)
|
||||
}
|
||||
|
||||
var expectedRoutes []*model.ProxyRoute
|
||||
scope := common.UptimeKumaMonitorScope
|
||||
if scope == "selected" {
|
||||
selectedList := strings.Split(common.UptimeKumaSelectedSites, ",")
|
||||
selectedMap := make(map[string]bool)
|
||||
for _, name := range selectedList {
|
||||
trimmedName := strings.TrimSpace(name)
|
||||
if trimmedName != "" {
|
||||
selectedMap[trimmedName] = true
|
||||
}
|
||||
}
|
||||
for _, route := range allRoutes {
|
||||
if route.Enabled && selectedMap[route.SiteName] {
|
||||
expectedRoutes = append(expectedRoutes, route)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for _, route := range allRoutes {
|
||||
if route.Enabled {
|
||||
expectedRoutes = append(expectedRoutes, route)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Connect to Uptime Kuma
|
||||
slog.Debug("Connecting to Uptime Kuma socket endpoint", "url", kumaUrl)
|
||||
client := uptimekuma.NewSocketIOClient(kumaUrl)
|
||||
if err := client.Connect(); err != nil {
|
||||
slog.Error("Failed to connect to Uptime Kuma endpoint", "url", kumaUrl, "error", err)
|
||||
return fmt.Errorf("failed to connect to Uptime Kuma: %w", err)
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
// 3. Login
|
||||
slog.Debug("Sending login request to Uptime Kuma", "username", kumaUsername)
|
||||
var loginAck string
|
||||
loginPayload := map[string]string{
|
||||
"username": kumaUsername,
|
||||
"password": kumaPassword,
|
||||
}
|
||||
loginAck, err = client.Emit("login", loginPayload)
|
||||
if err != nil {
|
||||
slog.Error("Failed to send login request to Uptime Kuma", "username", kumaUsername, "error", err)
|
||||
return fmt.Errorf("login request failed: %w", err)
|
||||
}
|
||||
|
||||
var loginResult struct {
|
||||
Ok bool `json:"ok"`
|
||||
}
|
||||
if err := uptimekuma.ParseAckResponse(loginAck, &loginResult); err != nil || !loginResult.Ok {
|
||||
slog.Error("Uptime Kuma login verification failed", "username", kumaUsername, "error", err)
|
||||
return fmt.Errorf("login failed: %w", err)
|
||||
}
|
||||
slog.Debug("Successfully logged into Uptime Kuma", "username", kumaUsername)
|
||||
|
||||
// 4. Wait for monitor list event
|
||||
slog.Debug("Waiting for monitor list push from Uptime Kuma")
|
||||
select {
|
||||
case <-client.GetMonitorListChan():
|
||||
slog.Debug("Received monitor list from Uptime Kuma")
|
||||
case <-time.After(5 * time.Second):
|
||||
slog.Error("Timeout waiting for Uptime Kuma monitorList push event")
|
||||
return fmt.Errorf("timeout waiting for monitorList event from Uptime Kuma")
|
||||
}
|
||||
|
||||
// 5. Get existing tags to find "OpenFlare"
|
||||
slog.Debug("Fetching tags from Uptime Kuma")
|
||||
tagsAck, err := client.Emit("getTags")
|
||||
if err != nil {
|
||||
slog.Error("Failed to request tags from Uptime Kuma", "error", err)
|
||||
return fmt.Errorf("failed to fetch tags: %w", err)
|
||||
}
|
||||
|
||||
var tagsResult struct {
|
||||
Ok bool `json:"ok"`
|
||||
Tags []uptimekuma.UptimeKumaTagItem `json:"tags"`
|
||||
}
|
||||
if err := uptimekuma.ParseAckResponse(tagsAck, &tagsResult); err != nil {
|
||||
slog.Error("Failed to parse tags response from Uptime Kuma", "error", err)
|
||||
return fmt.Errorf("parse tags response failed: %w", err)
|
||||
}
|
||||
|
||||
var openFlareTagID int
|
||||
for _, t := range tagsResult.Tags {
|
||||
if t.Name == "OpenFlare" {
|
||||
openFlareTagID = t.ID
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Create "OpenFlare" tag if not exists
|
||||
if openFlareTagID == 0 {
|
||||
slog.Debug("OpenFlare tag not found, creating new tag")
|
||||
addTagAck, err := client.Emit("addTag", map[string]string{
|
||||
"name": "OpenFlare",
|
||||
"color": "#4f46e5",
|
||||
})
|
||||
if err != nil {
|
||||
slog.Error("Failed to create OpenFlare tag in Uptime Kuma", "error", err)
|
||||
return fmt.Errorf("failed to create tag: %w", err)
|
||||
}
|
||||
var tagResult struct {
|
||||
Ok bool `json:"ok"`
|
||||
Tag struct {
|
||||
ID int `json:"id"`
|
||||
} `json:"tag"`
|
||||
}
|
||||
if err := uptimekuma.ParseAckResponse(addTagAck, &tagResult); err != nil || tagResult.Tag.ID == 0 {
|
||||
slog.Error("Failed to parse addTag response from Uptime Kuma", "error", err)
|
||||
return fmt.Errorf("parse addTag response failed: %w", err)
|
||||
}
|
||||
openFlareTagID = tagResult.Tag.ID
|
||||
slog.Debug("Successfully created OpenFlare tag", "tag_id", openFlareTagID)
|
||||
} else {
|
||||
slog.Debug("Found existing OpenFlare tag", "tag_id", openFlareTagID)
|
||||
}
|
||||
|
||||
// 6. Filter existing monitors by "OpenFlare" tag
|
||||
existingOpenFlareMonitors := make(map[string]uptimekuma.UptimeKumaMonitor)
|
||||
monitors := client.GetMonitorList()
|
||||
for _, m := range monitors {
|
||||
hasOpenFlareTag := false
|
||||
for _, tag := range m.Tags {
|
||||
if tag.Name == "OpenFlare" || tag.ID == openFlareTagID {
|
||||
hasOpenFlareTag = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if hasOpenFlareTag {
|
||||
existingOpenFlareMonitors[m.Name] = m
|
||||
}
|
||||
}
|
||||
|
||||
// Helper to format route URL
|
||||
getRouteURL := func(route *model.ProxyRoute) string {
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
domain := route.Domain
|
||||
if err == nil && len(domains) > 0 {
|
||||
domain = domains[0]
|
||||
}
|
||||
if route.EnableHTTPS {
|
||||
return "https://" + domain
|
||||
}
|
||||
return "http://" + domain
|
||||
}
|
||||
|
||||
expectedSitesMap := make(map[string]bool)
|
||||
|
||||
// 7. Sync Loop
|
||||
for _, route := range expectedRoutes {
|
||||
expectedSitesMap[route.SiteName] = true
|
||||
targetURL := getRouteURL(route)
|
||||
|
||||
existing, exists := existingOpenFlareMonitors[route.SiteName]
|
||||
if !exists {
|
||||
// Create monitor
|
||||
slog.Info("Creating monitor in Uptime Kuma", "name", route.SiteName, "url", targetURL)
|
||||
monitorPayload := map[string]any{
|
||||
"type": "http",
|
||||
"name": route.SiteName,
|
||||
"url": targetURL,
|
||||
"interval": common.UptimeKumaInterval,
|
||||
"maxretries": common.UptimeKumaRetry,
|
||||
"retryInterval": common.UptimeKumaRetryInterval,
|
||||
"timeout": common.UptimeKumaTimeout,
|
||||
"active": true,
|
||||
"resendInterval": 0,
|
||||
"expiryNotification": false,
|
||||
"ignoreTls": false,
|
||||
"accepted_statuscodes": []string{"200-299"},
|
||||
"dns_resolve_type": "A",
|
||||
"conditions": []any{},
|
||||
}
|
||||
addAck, err := client.Emit("add", monitorPayload)
|
||||
if err != nil {
|
||||
slog.Error("Failed to add monitor to Uptime Kuma", "name", route.SiteName, "error", err)
|
||||
continue
|
||||
}
|
||||
var addResult struct {
|
||||
Ok bool `json:"ok"`
|
||||
MonitorID int `json:"monitorID"`
|
||||
}
|
||||
if err := uptimekuma.ParseAckResponse(addAck, &addResult); err != nil || addResult.MonitorID == 0 {
|
||||
slog.Error("Failed to parse add monitor result", "name", route.SiteName, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Add tag
|
||||
slog.Debug("Adding OpenFlare tag to the new monitor", "name", route.SiteName, "monitor_id", addResult.MonitorID, "tag_id", openFlareTagID)
|
||||
tagAck, err := client.Emit("addMonitorTag", openFlareTagID, addResult.MonitorID, "")
|
||||
if err != nil {
|
||||
slog.Error("Failed to add tag to monitor in Uptime Kuma", "name", route.SiteName, "monitorID", addResult.MonitorID, "error", err)
|
||||
} else {
|
||||
if err := uptimekuma.ParseAckResponse(tagAck, nil); err != nil {
|
||||
slog.Error("Failed to parse add tag result", "name", route.SiteName, "monitorID", addResult.MonitorID, "error", err)
|
||||
} else {
|
||||
slog.Debug("OpenFlare tag successfully added to monitor", "name", route.SiteName, "monitor_id", addResult.MonitorID)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Check if updates are needed
|
||||
needsUpdate := existing.Url != targetURL ||
|
||||
existing.Interval != common.UptimeKumaInterval ||
|
||||
existing.MaxRetries != common.UptimeKumaRetry ||
|
||||
existing.RetryInterval != common.UptimeKumaRetryInterval ||
|
||||
existing.Timeout != common.UptimeKumaTimeout
|
||||
|
||||
if needsUpdate {
|
||||
slog.Info("Updating monitor in Uptime Kuma due to settings mismatch",
|
||||
"name", route.SiteName,
|
||||
"url_changed", existing.Url != targetURL,
|
||||
"interval_changed", existing.Interval != common.UptimeKumaInterval,
|
||||
"max_retries_changed", existing.MaxRetries != common.UptimeKumaRetry,
|
||||
"retry_interval_changed", existing.RetryInterval != common.UptimeKumaRetryInterval,
|
||||
"timeout_changed", existing.Timeout != common.UptimeKumaTimeout,
|
||||
)
|
||||
monitorPayload := map[string]any{
|
||||
"id": existing.ID,
|
||||
"type": "http",
|
||||
"name": route.SiteName,
|
||||
"url": targetURL,
|
||||
"interval": common.UptimeKumaInterval,
|
||||
"maxretries": common.UptimeKumaRetry,
|
||||
"retryInterval": common.UptimeKumaRetryInterval,
|
||||
"timeout": common.UptimeKumaTimeout,
|
||||
"active": true,
|
||||
"resendInterval": 0,
|
||||
"expiryNotification": false,
|
||||
"ignoreTls": false,
|
||||
"accepted_statuscodes": []string{"200-299"},
|
||||
"dns_resolve_type": "A",
|
||||
"conditions": []any{},
|
||||
}
|
||||
editAck, err := client.Emit("editMonitor", monitorPayload)
|
||||
if err != nil {
|
||||
slog.Error("Failed to edit monitor in Uptime Kuma", "name", route.SiteName, "error", err)
|
||||
} else {
|
||||
if err := uptimekuma.ParseAckResponse(editAck, nil); err != nil {
|
||||
slog.Error("Failed to parse edit monitor result", "name", route.SiteName, "error", err)
|
||||
} else {
|
||||
slog.Info("Successfully updated monitor in Uptime Kuma", "name", route.SiteName)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 8. Delete Loop
|
||||
for name, m := range existingOpenFlareMonitors {
|
||||
if !expectedSitesMap[name] {
|
||||
slog.Info("Deleting monitor in Uptime Kuma", "name", name, "monitorID", m.ID)
|
||||
deleteAck, err := client.Emit("deleteMonitor", m.ID)
|
||||
if err != nil {
|
||||
slog.Error("Failed to delete monitor in Uptime Kuma", "name", name, "monitorID", m.ID, "error", err)
|
||||
} else {
|
||||
if err := uptimekuma.ParseAckResponse(deleteAck, nil); err != nil {
|
||||
slog.Error("Failed to parse delete monitor result", "name", name, "monitorID", m.ID, "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,541 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultWAFBlockStatusCode = 418
|
||||
maxWAFBlockBodyBytes = 16 * 1024
|
||||
)
|
||||
|
||||
type WAFRuleGroupInput struct {
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockStatusCode int `json:"block_status_code"`
|
||||
BlockResponseBody string `json:"block_response_body"`
|
||||
IPWhitelist []string `json:"ip_whitelist"`
|
||||
IPBlacklist []string `json:"ip_blacklist"`
|
||||
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
|
||||
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
|
||||
CountryWhitelist []string `json:"country_whitelist"`
|
||||
CountryBlacklist []string `json:"country_blacklist"`
|
||||
RegionWhitelist []string `json:"region_whitelist"`
|
||||
RegionBlacklist []string `json:"region_blacklist"`
|
||||
Remark string `json:"remark"`
|
||||
PoWEnabled bool `json:"pow_enabled"`
|
||||
PoWConfig json.RawMessage `json:"pow_config"`
|
||||
}
|
||||
|
||||
type WAFRuleGroupView struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IsGlobal bool `json:"is_global"`
|
||||
BlockStatusCode int `json:"block_status_code"`
|
||||
BlockResponseBody string `json:"block_response_body"`
|
||||
IPWhitelist []string `json:"ip_whitelist"`
|
||||
IPBlacklist []string `json:"ip_blacklist"`
|
||||
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
|
||||
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
|
||||
CountryWhitelist []string `json:"country_whitelist"`
|
||||
CountryBlacklist []string `json:"country_blacklist"`
|
||||
RegionWhitelist []string `json:"region_whitelist"`
|
||||
RegionBlacklist []string `json:"region_blacklist"`
|
||||
Remark string `json:"remark"`
|
||||
PoWEnabled bool `json:"pow_enabled"`
|
||||
PoWConfig *ProxyRoutePoWConfig `json:"pow_config"`
|
||||
AppliedSiteIDs []uint `json:"applied_site_ids"`
|
||||
AppliedSiteCount int `json:"applied_site_count"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
type WAFSiteRuleGroupsView struct {
|
||||
RouteID uint `json:"route_id"`
|
||||
GlobalRuleGroup *WAFRuleGroupView `json:"global_rule_group"`
|
||||
RuleGroups []WAFRuleGroupView `json:"rule_groups"`
|
||||
AppliedRuleGroups []WAFRuleGroupView `json:"applied_rule_groups"`
|
||||
AppliedIDs []uint `json:"applied_ids"`
|
||||
}
|
||||
|
||||
func ListWAFRuleGroups() ([]WAFRuleGroupView, error) {
|
||||
if err := EnsureDefaultWAFRuleGroup(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups, err := model.ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := loadWAFBindings()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]WAFRuleGroupView, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
view, err := buildWAFRuleGroupView(group, bindings[group.ID])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func GetWAFRuleGroup(id uint) (*WAFRuleGroupView, error) {
|
||||
group, err := model.GetWAFRuleGroupByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := loadWAFBindings()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
view, err := buildWAFRuleGroupView(group, bindings[group.ID])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &view, nil
|
||||
}
|
||||
|
||||
func CreateWAFRuleGroup(input WAFRuleGroupInput) (*WAFRuleGroupView, error) {
|
||||
group, err := buildWAFRuleGroup(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
group.IsGlobal = false
|
||||
if err := group.Insert(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFRuleGroup(group.ID)
|
||||
}
|
||||
|
||||
func UpdateWAFRuleGroup(id uint, input WAFRuleGroupInput) (*WAFRuleGroupView, error) {
|
||||
group, err := model.GetWAFRuleGroupByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
isGlobal := group.IsGlobal
|
||||
group, err = buildWAFRuleGroup(group, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
group.IsGlobal = isGlobal
|
||||
if isGlobal && strings.TrimSpace(group.Name) == "" {
|
||||
group.Name = "全局规则组"
|
||||
}
|
||||
if err := group.Update(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFRuleGroup(group.ID)
|
||||
}
|
||||
|
||||
func DeleteWAFRuleGroup(id uint) error {
|
||||
group, err := model.GetWAFRuleGroupByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return errors.New("全局 WAF 规则组不能删除")
|
||||
}
|
||||
return model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("rule_group_id = ?", group.ID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(group).Error
|
||||
})
|
||||
}
|
||||
|
||||
func ReplaceWAFRuleGroupSites(groupID uint, routeIDs []uint) (*WAFRuleGroupView, error) {
|
||||
group, err := model.GetWAFRuleGroupByID(groupID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return nil, errors.New("全局 WAF 规则组默认应用到所有网站,不能手动绑定")
|
||||
}
|
||||
normalized, err := normalizeWAFRouteIDs(routeIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("rule_group_id = ?", groupID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, routeID := range normalized {
|
||||
binding := model.WAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID}
|
||||
if err := tx.Create(&binding).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFRuleGroup(groupID)
|
||||
}
|
||||
|
||||
func GetWAFSiteRuleGroups(routeID uint) (*WAFSiteRuleGroupsView, error) {
|
||||
if _, err := model.GetProxyRouteByID(routeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups, err := ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
appliedIDs, err := ListWAFSiteRuleGroupIDs(routeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
appliedSet := make(map[uint]struct{}, len(appliedIDs))
|
||||
for _, id := range appliedIDs {
|
||||
appliedSet[id] = struct{}{}
|
||||
}
|
||||
var global *WAFRuleGroupView
|
||||
custom := make([]WAFRuleGroupView, 0, len(groups))
|
||||
applied := make([]WAFRuleGroupView, 0, len(appliedIDs))
|
||||
for index := range groups {
|
||||
group := groups[index]
|
||||
if group.IsGlobal {
|
||||
item := group
|
||||
global = &item
|
||||
continue
|
||||
}
|
||||
custom = append(custom, group)
|
||||
if _, ok := appliedSet[group.ID]; ok {
|
||||
applied = append(applied, group)
|
||||
}
|
||||
}
|
||||
return &WAFSiteRuleGroupsView{
|
||||
RouteID: routeID,
|
||||
GlobalRuleGroup: global,
|
||||
RuleGroups: custom,
|
||||
AppliedRuleGroups: applied,
|
||||
AppliedIDs: appliedIDs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ReplaceWAFSiteRuleGroups(routeID uint, groupIDs []uint) (*WAFSiteRuleGroupsView, error) {
|
||||
if _, err := model.GetProxyRouteByID(routeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
normalized, err := normalizeWAFRuleGroupIDs(groupIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("proxy_route_id = ?", routeID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, groupID := range normalized {
|
||||
binding := model.WAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID}
|
||||
if err := tx.Create(&binding).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFSiteRuleGroups(routeID)
|
||||
}
|
||||
|
||||
func ListWAFSiteRuleGroupIDs(routeID uint) ([]uint, error) {
|
||||
var bindings []model.WAFRuleGroupBinding
|
||||
if err := model.DB.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ids := make([]uint, 0, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
ids = append(ids, binding.RuleGroupID)
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func EnsureDefaultWAFRuleGroup() error {
|
||||
_, err := model.GetGlobalWAFRuleGroup()
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
group := &model.WAFRuleGroup{
|
||||
Name: "全局规则组",
|
||||
Enabled: true,
|
||||
IsGlobal: true,
|
||||
BlockStatusCode: defaultWAFBlockStatusCode,
|
||||
IPWhitelist: "[]",
|
||||
IPBlacklist: "[]",
|
||||
IPWhitelistGroups: "[]",
|
||||
IPBlacklistGroups: "[]",
|
||||
CountryWhitelist: "[]",
|
||||
CountryBlacklist: "[]",
|
||||
RegionWhitelist: "[]",
|
||||
RegionBlacklist: "[]",
|
||||
PoWEnabled: false,
|
||||
PoWConfig: "{}",
|
||||
BlockResponseBody: "",
|
||||
}
|
||||
return group.Insert()
|
||||
}
|
||||
|
||||
func buildWAFRuleGroup(group *model.WAFRuleGroup, input WAFRuleGroupInput) (*model.WAFRuleGroup, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, errors.New("规则组名称不能为空")
|
||||
}
|
||||
statusCode := input.BlockStatusCode
|
||||
if statusCode == 0 {
|
||||
statusCode = defaultWAFBlockStatusCode
|
||||
}
|
||||
if statusCode < 400 || statusCode > 599 {
|
||||
return nil, errors.New("拦截状态码必须在 400-599 之间")
|
||||
}
|
||||
if len([]byte(input.BlockResponseBody)) > maxWAFBlockBodyBytes {
|
||||
return nil, fmt.Errorf("拦截页面内容不能超过 %d 字节", maxWAFBlockBodyBytes)
|
||||
}
|
||||
ipWhitelist, err := normalizeWAFIPList(input.IPWhitelist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("IP 白名单无效: %w", err)
|
||||
}
|
||||
ipBlacklist, err := normalizeWAFIPList(input.IPBlacklist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("IP 黑名单无效: %w", err)
|
||||
}
|
||||
ipWhitelistGroups, err := normalizeWAFIPGroupIDs(input.IPWhitelistGroups)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("IP 白名单引用无效: %w", err)
|
||||
}
|
||||
ipBlacklistGroups, err := normalizeWAFIPGroupIDs(input.IPBlacklistGroups)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("IP 黑名单引用无效: %w", err)
|
||||
}
|
||||
countryWhitelist, err := normalizeWAFCountryList(input.CountryWhitelist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("地域白名单无效: %w", err)
|
||||
}
|
||||
countryBlacklist, err := normalizeWAFCountryList(input.CountryBlacklist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("地域黑名单无效: %w", err)
|
||||
}
|
||||
regionWhitelist := normalizeStringList(input.RegionWhitelist)
|
||||
regionBlacklist := normalizeStringList(input.RegionBlacklist)
|
||||
powConfigRaw := strings.TrimSpace(string(input.PoWConfig))
|
||||
if powConfigRaw == "" {
|
||||
powConfigRaw = "{}"
|
||||
}
|
||||
powConfig, err := normalizePoWConfig(input.PoWEnabled, powConfigRaw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
powConfigJSON, _ := json.Marshal(powConfig)
|
||||
|
||||
ipWhitelistJSON, _ := json.Marshal(ipWhitelist)
|
||||
ipBlacklistJSON, _ := json.Marshal(ipBlacklist)
|
||||
ipWhitelistGroupsJSON, _ := json.Marshal(ipWhitelistGroups)
|
||||
ipBlacklistGroupsJSON, _ := json.Marshal(ipBlacklistGroups)
|
||||
countryWhitelistJSON, _ := json.Marshal(countryWhitelist)
|
||||
countryBlacklistJSON, _ := json.Marshal(countryBlacklist)
|
||||
regionWhitelistJSON, _ := json.Marshal(regionWhitelist)
|
||||
regionBlacklistJSON, _ := json.Marshal(regionBlacklist)
|
||||
|
||||
if group == nil {
|
||||
group = &model.WAFRuleGroup{}
|
||||
}
|
||||
group.Name = name
|
||||
group.Enabled = input.Enabled
|
||||
group.BlockStatusCode = statusCode
|
||||
group.BlockResponseBody = input.BlockResponseBody
|
||||
group.IPWhitelist = string(ipWhitelistJSON)
|
||||
group.IPBlacklist = string(ipBlacklistJSON)
|
||||
group.IPWhitelistGroups = string(ipWhitelistGroupsJSON)
|
||||
group.IPBlacklistGroups = string(ipBlacklistGroupsJSON)
|
||||
group.CountryWhitelist = string(countryWhitelistJSON)
|
||||
group.CountryBlacklist = string(countryBlacklistJSON)
|
||||
group.RegionWhitelist = string(regionWhitelistJSON)
|
||||
group.RegionBlacklist = string(regionBlacklistJSON)
|
||||
group.PoWEnabled = input.PoWEnabled
|
||||
group.PoWConfig = string(powConfigJSON)
|
||||
group.Remark = strings.TrimSpace(input.Remark)
|
||||
return group, nil
|
||||
}
|
||||
|
||||
func buildWAFRuleGroupView(group *model.WAFRuleGroup, appliedSiteIDs []uint) (WAFRuleGroupView, error) {
|
||||
if group == nil {
|
||||
return WAFRuleGroupView{}, errors.New("waf rule group is nil")
|
||||
}
|
||||
sort.Slice(appliedSiteIDs, func(i, j int) bool { return appliedSiteIDs[i] < appliedSiteIDs[j] })
|
||||
view := WAFRuleGroupView{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
Enabled: group.Enabled,
|
||||
IsGlobal: group.IsGlobal,
|
||||
BlockStatusCode: group.BlockStatusCode,
|
||||
BlockResponseBody: group.BlockResponseBody,
|
||||
Remark: group.Remark,
|
||||
PoWEnabled: group.PoWEnabled,
|
||||
AppliedSiteIDs: appliedSiteIDs,
|
||||
AppliedSiteCount: len(appliedSiteIDs),
|
||||
CreatedAt: group.CreatedAt.Format(time.RFC3339),
|
||||
UpdatedAt: group.UpdatedAt.Format(time.RFC3339),
|
||||
}
|
||||
var err error
|
||||
if view.IPWhitelist, err = decodeStringList(group.IPWhitelist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.IPBlacklist, err = decodeStringList(group.IPBlacklist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
view.IPWhitelistGroups = mustDecodeUintList(group.IPWhitelistGroups)
|
||||
view.IPBlacklistGroups = mustDecodeUintList(group.IPBlacklistGroups)
|
||||
if view.CountryWhitelist, err = decodeStringList(group.CountryWhitelist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.CountryBlacklist, err = decodeStringList(group.CountryBlacklist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.RegionWhitelist, err = decodeStringList(group.RegionWhitelist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.RegionBlacklist, err = decodeStringList(group.RegionBlacklist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.PoWConfig, err = decodeStoredPoWConfig(group.PoWEnabled, group.PoWConfig); err != nil {
|
||||
return view, err
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func loadWAFBindings() (map[uint][]uint, error) {
|
||||
var bindings []model.WAFRuleGroupBinding
|
||||
if err := model.DB.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[uint][]uint, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func normalizeWAFIPList(items []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.TrimSpace(raw)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(item, "/") {
|
||||
prefix, err := netip.ParsePrefix(item)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s 不是合法 IP 段", item)
|
||||
}
|
||||
item = prefix.Masked().String()
|
||||
} else {
|
||||
addr, err := netip.ParseAddr(item)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s 不是合法 IP", item)
|
||||
}
|
||||
item = addr.String()
|
||||
}
|
||||
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))
|
||||
for _, raw := range items {
|
||||
item := strings.ToUpper(strings.TrimSpace(raw))
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if len(item) != 2 || !unicode.IsLetter(rune(item[0])) || !unicode.IsLetter(rune(item[1])) {
|
||||
return nil, fmt.Errorf("%s 不是合法国家代码", item)
|
||||
}
|
||||
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))
|
||||
for _, raw := range items {
|
||||
item := strings.TrimSpace(raw)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
sort.Strings(normalized)
|
||||
return normalized
|
||||
}
|
||||
|
||||
func decodeStringList(raw string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(text), &items); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func normalizeWAFRouteIDs(routeIDs []uint) ([]uint, error) {
|
||||
normalized := uniqueUintIDs(routeIDs)
|
||||
for _, routeID := range normalized {
|
||||
if _, err := model.GetProxyRouteByID(routeID); err != nil {
|
||||
return nil, fmt.Errorf("网站 %d 不存在", routeID)
|
||||
}
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeWAFRuleGroupIDs(groupIDs []uint) ([]uint, error) {
|
||||
normalized := uniqueUintIDs(groupIDs)
|
||||
for _, groupID := range normalized {
|
||||
group, err := model.GetWAFRuleGroupByID(groupID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID)
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return nil, errors.New("全局 WAF 规则组不需要手动绑定")
|
||||
}
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func uniqueUintIDs(ids []uint) []uint {
|
||||
normalized := make([]uint, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id == 0 {
|
||||
continue
|
||||
}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
normalized = utils.Unique(normalized)
|
||||
sort.Slice(normalized, func(i, j int) bool { return normalized[i] < normalized[j] })
|
||||
return normalized
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,514 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/model"
|
||||
)
|
||||
|
||||
func TestWAFRuleGroupValidationAndNormalization(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
group, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "edge guard",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 451,
|
||||
IPWhitelist: []string{" 192.0.2.1 ", "192.0.2.1", "198.51.100.0/24"},
|
||||
IPBlacklist: []string{"203.0.113.10"},
|
||||
CountryBlacklist: []string{" cn ", "CN", "us"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if len(group.IPWhitelist) != 2 || group.IPWhitelist[0] != "192.0.2.1" || group.IPWhitelist[1] != "198.51.100.0/24" {
|
||||
t.Fatalf("unexpected normalized ip whitelist: %#v", group.IPWhitelist)
|
||||
}
|
||||
if len(group.CountryBlacklist) != 2 || group.CountryBlacklist[0] != "CN" || group.CountryBlacklist[1] != "US" {
|
||||
t.Fatalf("unexpected normalized countries: %#v", group.CountryBlacklist)
|
||||
}
|
||||
|
||||
if _, err = CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "bad ip",
|
||||
Enabled: true,
|
||||
IPBlacklist: []string{"not-an-ip"},
|
||||
}); err == nil {
|
||||
t.Fatal("expected invalid IP to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFGlobalGroupAndBindings(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
groups, err := ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
t.Fatalf("ListWAFRuleGroups failed: %v", err)
|
||||
}
|
||||
if len(groups) == 0 || !groups[0].IsGlobal {
|
||||
t.Fatalf("expected default global WAF rule group, got %#v", groups)
|
||||
}
|
||||
if err = DeleteWAFRuleGroup(groups[0].ID); err == nil {
|
||||
t.Fatal("expected global WAF rule group delete to be rejected")
|
||||
}
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "waf-site",
|
||||
Domains: []string{"waf.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
custom, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "custom",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 418,
|
||||
IPBlacklist: []string{"203.0.113.10"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if _, err = ReplaceWAFRuleGroupSites(custom.ID, []uint{route.ID}); err != nil {
|
||||
t.Fatalf("ReplaceWAFRuleGroupSites failed: %v", err)
|
||||
}
|
||||
siteGroups, err := GetWAFSiteRuleGroups(route.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetWAFSiteRuleGroups failed: %v", err)
|
||||
}
|
||||
if len(siteGroups.AppliedIDs) != 1 || siteGroups.AppliedIDs[0] != custom.ID {
|
||||
t.Fatalf("unexpected site WAF bindings: %#v", siteGroups.AppliedIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionIncludesWAFSnapshotAndRuntimeConfig(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "waf-publish",
|
||||
Domains: []string{"waf-publish.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
group, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "publish group",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 451,
|
||||
IPBlacklist: []string{"203.0.113.0/24"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if _, err = ReplaceWAFSiteRuleGroups(route.ID, []uint{group.ID}); err != nil {
|
||||
t.Fatalf("ReplaceWAFSiteRuleGroups failed: %v", err)
|
||||
}
|
||||
result, err := PublishConfigVersion("root", false)
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "access_by_lua_file __OPENFLARE_LUA_DIR__/waf/check.lua;") {
|
||||
t.Fatal("expected route config to include WAF lua access hook")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"waf"`) {
|
||||
t.Fatal("expected snapshot to include waf document")
|
||||
}
|
||||
var files []SupportFile
|
||||
if err = json.Unmarshal([]byte(result.Version.SupportFilesJSON), &files); err != nil {
|
||||
t.Fatalf("decode support files failed: %v", err)
|
||||
}
|
||||
found := false
|
||||
for _, file := range files {
|
||||
if file.Path == "waf_config.json" && strings.Contains(file.Content, "203.0.113.0/24") {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("expected waf_config.json support file, got %#v", files)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFIPGroupCRUDAndRuleGroupReference(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
ipGroup, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "bad actors",
|
||||
Type: WAFIPGroupTypeManual,
|
||||
Enabled: true,
|
||||
IPList: []string{"203.0.113.10", "203.0.113.10", "198.51.100.0/24"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
if len(ipGroup.IPList) != 2 {
|
||||
t.Fatalf("unexpected normalized IP group list: %#v", ipGroup.IPList)
|
||||
}
|
||||
|
||||
group, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "referenced",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 403,
|
||||
IPBlacklistGroups: []uint{ipGroup.ID},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if len(group.IPBlacklistGroups) != 1 || group.IPBlacklistGroups[0] != ipGroup.ID {
|
||||
t.Fatalf("unexpected blacklist group refs: %#v", group.IPBlacklistGroups)
|
||||
}
|
||||
if err = DeleteWAFIPGroup(ipGroup.ID); err == nil {
|
||||
t.Fatal("expected referenced IP group delete to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFIPGroupSubscriptionParsers(t *testing.T) {
|
||||
textItems, err := parseWAFIPGroupSubscription([]byte("# comment\n203.0.113.10\n\n198.51.100.0/24\n"), "text", "")
|
||||
if err != nil {
|
||||
t.Fatalf("parse text subscription failed: %v", err)
|
||||
}
|
||||
if len(textItems) != 2 || textItems[0] != "198.51.100.0/24" || textItems[1] != "203.0.113.10" {
|
||||
t.Fatalf("unexpected text subscription items: %#v", textItems)
|
||||
}
|
||||
|
||||
jsonItems, err := parseWAFIPGroupSubscription([]byte(`{"data":{"items":[{"ip":"203.0.113.11"},{"ip":"203.0.113.12"}]}}`), "json", "data.items[].ip")
|
||||
if err != nil {
|
||||
t.Fatalf("parse json subscription failed: %v", err)
|
||||
}
|
||||
if len(jsonItems) != 2 || jsonItems[0] != "203.0.113.11" || jsonItems[1] != "203.0.113.12" {
|
||||
t.Fatalf("unexpected json subscription items: %#v", jsonItems)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncWAFIPGroupDownloadsSubscription(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte("203.0.113.20\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
group, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "subscription",
|
||||
Type: WAFIPGroupTypeSubscription,
|
||||
Enabled: true,
|
||||
SubscriptionURL: server.URL,
|
||||
SubscriptionFormat: WAFIPGroupSubscriptionFormatText,
|
||||
SyncIntervalMinutes: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
result, err := SyncWAFIPGroup(group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("SyncWAFIPGroup failed: %v", err)
|
||||
}
|
||||
if result.IPCount != 1 || result.Group.IPList[0] != "203.0.113.20" {
|
||||
t.Fatalf("unexpected sync result: %+v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncWAFIPGroupAutomaticExprRules(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.10", "app.example.com", 101, 81)
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.11", "198.51.100.10", 60, 0)
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.12", "app.example.com", 120, 10)
|
||||
|
||||
group, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "auto blacklist",
|
||||
Type: WAFIPGroupTypeAutomatic,
|
||||
Enabled: true,
|
||||
AutoConfig: json.RawMessage(`{
|
||||
"lookback_minutes": 60,
|
||||
"rules": [
|
||||
{"name":"单 IP 404 高频扫描","expr":"request_count > 100 && StatusRatio(404) >= 0.8"},
|
||||
{"name":"单 IP 直连访问异常","expr":"ip_host_count > 50 && ip_host_ratio > 0.5"}
|
||||
]
|
||||
}`),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
result, err := SyncWAFIPGroup(group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("SyncWAFIPGroup failed: %v", err)
|
||||
}
|
||||
if result.IPCount != 2 {
|
||||
t.Fatalf("expected two matched IPs, got %#v", result)
|
||||
}
|
||||
want := map[string]bool{"203.0.113.10": true, "203.0.113.11": true}
|
||||
for _, item := range result.Group.IPList {
|
||||
if !want[item] {
|
||||
t.Fatalf("unexpected matched IP %s in %#v", item, result.Group.IPList)
|
||||
}
|
||||
delete(want, item)
|
||||
}
|
||||
if len(want) != 0 {
|
||||
t.Fatalf("missing matched IPs: %#v", want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFIPGroupAutoConfigReturnsMatchedIPs(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.10", "app.example.com", 101, 81)
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.11", "198.51.100.10", 60, 0)
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.12", "app.example.com", 120, 10)
|
||||
|
||||
result, err := TestWAFIPGroupAutoConfig(WAFIPGroupAutoTestInput{
|
||||
AutoConfig: json.RawMessage(`{
|
||||
"lookback_minutes": 60,
|
||||
"rules": [
|
||||
{"name":"单 IP 404 高频扫描","expr":"request_count > 100 && StatusRatio(404) >= 0.8"},
|
||||
{"name":"单 IP 直连访问异常","expr":"ip_host_count > 50 && ip_host_ratio > 0.5"}
|
||||
]
|
||||
}`),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("TestWAFIPGroupAutoConfig failed: %v", err)
|
||||
}
|
||||
if result.MatchedCount != 2 || result.RuleCount != 2 || result.LookbackMinutes != 60 {
|
||||
t.Fatalf("unexpected test result: %+v", result)
|
||||
}
|
||||
want := map[string]bool{"203.0.113.10": true, "203.0.113.11": true}
|
||||
for _, item := range result.MatchedIPs {
|
||||
if !want[item] {
|
||||
t.Fatalf("unexpected matched IP %s in %#v", item, result.MatchedIPs)
|
||||
}
|
||||
delete(want, item)
|
||||
}
|
||||
if len(want) != 0 {
|
||||
t.Fatalf("missing matched IPs: %#v", want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFIPGroupAutomaticRejectsInvalidExpr(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
if _, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "bad auto",
|
||||
Type: WAFIPGroupTypeAutomatic,
|
||||
Enabled: true,
|
||||
AutoConfig: json.RawMessage(`{
|
||||
"rules": [{"name":"bad","expr":"request_count > "}]
|
||||
}`),
|
||||
}); err == nil {
|
||||
t.Fatal("expected invalid Expr to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionKeepsWAFIPGroupReferences(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "waf-ip-groups",
|
||||
Domains: []string{"waf-ip-groups.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
ipGroup, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "publish refs",
|
||||
Type: WAFIPGroupTypeManual,
|
||||
Enabled: true,
|
||||
IPList: []string{"203.0.113.30"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
ruleGroup, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "publish group refs",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 451,
|
||||
IPBlacklistGroups: []uint{ipGroup.ID},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if _, err = ReplaceWAFSiteRuleGroups(route.ID, []uint{ruleGroup.ID}); err != nil {
|
||||
t.Fatalf("ReplaceWAFSiteRuleGroups failed: %v", err)
|
||||
}
|
||||
result, err := PublishConfigVersion("root", false)
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"ip_groups"`) {
|
||||
t.Fatal("expected snapshot to include waf ip groups")
|
||||
}
|
||||
if strings.Contains(result.Version.SnapshotJSON, "203.0.113.30") {
|
||||
t.Fatal("expected snapshot to avoid embedding waf ip group members")
|
||||
}
|
||||
var files []SupportFile
|
||||
if err = json.Unmarshal([]byte(result.Version.SupportFilesJSON), &files); err != nil {
|
||||
t.Fatalf("decode support files failed: %v", err)
|
||||
}
|
||||
foundReference := false
|
||||
for _, file := range files {
|
||||
if file.Path == "waf_config.json" {
|
||||
if strings.Contains(file.Content, "203.0.113.30") {
|
||||
t.Fatalf("expected waf_config.json to avoid expanded IP group members, got %s", file.Content)
|
||||
}
|
||||
if strings.Contains(file.Content, `"ip_blacklist_group_ids":[`) {
|
||||
foundReference = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !foundReference {
|
||||
t.Fatalf("expected IP group reference in waf_config.json, got %#v", files)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFIPGroupAutomaticTTLExpiration(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
// Seed access logs at now
|
||||
seedWAFNodeAccessLogs(t, now, "203.0.113.10", "app.example.com", 120, 100)
|
||||
|
||||
group, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "auto ttl blacklist",
|
||||
Type: WAFIPGroupTypeAutomatic,
|
||||
Enabled: true,
|
||||
AutoConfig: json.RawMessage(`{
|
||||
"lookback_minutes": 60,
|
||||
"ttl": 10,
|
||||
"rules": [
|
||||
{"name":"404 Scan","expr":"request_count > 100 && StatusRatio(404) >= 0.8"}
|
||||
]
|
||||
}`),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
|
||||
groupModel, err := model.GetWAFIPGroupByID(group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetWAFIPGroupByID failed: %v", err)
|
||||
}
|
||||
|
||||
// First Sync (at now): should match 203.0.113.10
|
||||
res1, err := syncWAFIPGroup(groupModel, now)
|
||||
if err != nil {
|
||||
t.Fatalf("First Sync failed: %v", err)
|
||||
}
|
||||
if res1.IPCount != 1 || res1.Group.IPList[0] != "203.0.113.10" {
|
||||
t.Fatalf("expected 203.0.113.10 to be blacklisted, got: %#v", res1.Group.IPList)
|
||||
}
|
||||
if len(res1.Group.ExtIPs) != 1 || res1.Group.ExtIPs[0].IP != "203.0.113.10" {
|
||||
t.Fatalf("expected 203.0.113.10 to be in ExtIPs, got: %#v", res1.Group.ExtIPs)
|
||||
}
|
||||
|
||||
// Second Sync (65 minutes later):
|
||||
// Since 65 minutes is outside the 60 minutes lookback window, the original logs won't match.
|
||||
// And since 65 minutes > 10s TTL, it should be expired and removed!
|
||||
futureTime := now.Add(65 * time.Minute)
|
||||
res2, err := syncWAFIPGroup(groupModel, futureTime)
|
||||
if err != nil {
|
||||
t.Fatalf("Second Sync failed: %v", err)
|
||||
}
|
||||
if res2.IPCount != 0 {
|
||||
t.Fatalf("expected IP to be expired and removed, got: %#v", res2.Group.IPList)
|
||||
}
|
||||
if len(res2.Group.ExtIPs) != 0 {
|
||||
t.Fatalf("expected ExtIPs to be empty after expiration, got: %#v", res2.Group.ExtIPs)
|
||||
}
|
||||
|
||||
// Third Sync: test lease refresh / extension!
|
||||
// Re-run sync at now to get it captured again first
|
||||
_, err = syncWAFIPGroup(groupModel, now)
|
||||
if err != nil {
|
||||
t.Fatalf("Re-sync at now failed: %v", err)
|
||||
}
|
||||
|
||||
// Now run sync at now + 5 seconds (5s < 10s TTL, so not expired, but matched again!):
|
||||
// Since it matches again, it should keep the IP active and extend CapturedAt to now + 5s!
|
||||
futureTime2 := now.Add(5 * time.Second)
|
||||
res3, err := syncWAFIPGroup(groupModel, futureTime2)
|
||||
if err != nil {
|
||||
t.Fatalf("Third Sync failed: %v", err)
|
||||
}
|
||||
if res3.IPCount != 1 || res3.Group.IPList[0] != "203.0.113.10" {
|
||||
t.Fatalf("expected IP to remain active, got: %#v", res3.Group.IPList)
|
||||
}
|
||||
if len(res3.Group.ExtIPs) != 1 || res3.Group.ExtIPs[0].CapturedAt != futureTime2.Format(time.RFC3339) {
|
||||
t.Fatalf("expected CapturedAt to be updated to %v, got %v", futureTime2.Format(time.RFC3339), res3.Group.ExtIPs[0].CapturedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func seedWAFNodeAccessLogs(t *testing.T, loggedAt time.Time, remoteAddr string, host string, total int, notFound int) {
|
||||
t.Helper()
|
||||
for i := 0; i < total; i++ {
|
||||
statusCode := http.StatusOK
|
||||
if i < notFound {
|
||||
statusCode = http.StatusNotFound
|
||||
}
|
||||
if err := model.DB.Create(&model.NodeAccessLog{
|
||||
NodeID: "node-waf-auto",
|
||||
LoggedAt: loggedAt.Add(-time.Duration(i%30) * time.Second),
|
||||
RemoteAddr: remoteAddr,
|
||||
Host: host,
|
||||
Path: "/probe",
|
||||
StatusCode: statusCode,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed access log: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncWAFIPGroupAutomaticCustomStatusRules(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
now := time.Now().UTC()
|
||||
// Seed 10 requests from 203.0.113.50, where 3 return 403, 7 return 200
|
||||
seedWAFNodeAccessLogsWithStatus(t, now, "203.0.113.50", "app.example.com", 7, http.StatusOK)
|
||||
seedWAFNodeAccessLogsWithStatus(t, now, "203.0.113.50", "app.example.com", 3, http.StatusForbidden)
|
||||
|
||||
group, err := CreateWAFIPGroup(WAFIPGroupInput{
|
||||
Name: "custom status code blacklist",
|
||||
Type: WAFIPGroupTypeAutomatic,
|
||||
Enabled: true,
|
||||
AutoConfig: json.RawMessage(`{
|
||||
"lookback_minutes": 60,
|
||||
"rules": [
|
||||
{"name":"高频 403 探测","expr":"StatusCount(403) >= 3 && StatusRatio(403) >= 0.3"}
|
||||
]
|
||||
}`),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFIPGroup failed: %v", err)
|
||||
}
|
||||
result, err := SyncWAFIPGroup(group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("SyncWAFIPGroup failed: %v", err)
|
||||
}
|
||||
if result.IPCount != 1 || result.Group.IPList[0] != "203.0.113.50" {
|
||||
t.Fatalf("expected 203.0.113.50 to be matched, got %#v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func seedWAFNodeAccessLogsWithStatus(t *testing.T, loggedAt time.Time, remoteAddr string, host string, count int, statusCode int) {
|
||||
t.Helper()
|
||||
for i := 0; i < count; i++ {
|
||||
if err := model.DB.Create(&model.NodeAccessLog{
|
||||
NodeID: "node-waf-auto",
|
||||
LoggedAt: loggedAt.Add(-time.Duration(i%30) * time.Second),
|
||||
RemoteAddr: remoteAddr,
|
||||
Host: host,
|
||||
Path: "/probe",
|
||||
StatusCode: statusCode,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("failed to seed access log: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type WSMessage struct {
|
||||
Type string `json:"type"`
|
||||
Payload any `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
type WSClient struct {
|
||||
id string
|
||||
send chan WSMessage
|
||||
done chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (client *WSClient) ID() string {
|
||||
if client == nil {
|
||||
return ""
|
||||
}
|
||||
return client.id
|
||||
}
|
||||
|
||||
func (client *WSClient) Messages() <-chan WSMessage {
|
||||
if client == nil {
|
||||
return nil
|
||||
}
|
||||
return client.send
|
||||
}
|
||||
|
||||
func (client *WSClient) Done() <-chan struct{} {
|
||||
if client == nil {
|
||||
return nil
|
||||
}
|
||||
return client.done
|
||||
}
|
||||
|
||||
func (client *WSClient) Send(message WSMessage) bool {
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-client.done:
|
||||
return false
|
||||
case client.send <- message:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (client *WSClient) Close() {
|
||||
if client == nil {
|
||||
return
|
||||
}
|
||||
client.once.Do(func() {
|
||||
close(client.done)
|
||||
})
|
||||
}
|
||||
|
||||
type WSHub struct {
|
||||
name string
|
||||
mu sync.RWMutex
|
||||
clients map[string]*WSClient
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
func NewWSHub(name string) *WSHub {
|
||||
h := &WSHub{
|
||||
name: name,
|
||||
clients: make(map[string]*WSClient),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
go h.startPingLoop()
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *WSHub) Close() {
|
||||
close(h.done)
|
||||
}
|
||||
|
||||
func (h *WSHub) startPingLoop() {
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-h.done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
h.mu.RLock()
|
||||
if len(h.clients) == 0 {
|
||||
h.mu.RUnlock()
|
||||
continue
|
||||
}
|
||||
clients := make([]*WSClient, 0, len(h.clients))
|
||||
for _, client := range h.clients {
|
||||
clients = append(clients, client)
|
||||
}
|
||||
h.mu.RUnlock()
|
||||
|
||||
for _, client := range clients {
|
||||
if !client.Send(WSMessage{
|
||||
Type: "ping",
|
||||
}) {
|
||||
slog.Warn("ws client send ping failed, queue full, disconnecting", "hub", h.name, "id", client.id)
|
||||
h.Disconnect(client.id)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ShutdownWSHubs() {
|
||||
DefaultAgentWSHub.Close()
|
||||
DefaultFlaredWSHub.Close()
|
||||
DefaultRelayWSHub.Close()
|
||||
}
|
||||
|
||||
func (h *WSHub) Register(id string) *WSClient {
|
||||
client := &WSClient{
|
||||
id: id,
|
||||
send: make(chan WSMessage, 16),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
h.mu.Lock()
|
||||
if existing := h.clients[id]; existing != nil {
|
||||
slog.Debug("ws replacing existing connection", "hub", h.name, "id", id)
|
||||
existing.Close()
|
||||
}
|
||||
h.clients[id] = client
|
||||
count := len(h.clients)
|
||||
h.mu.Unlock()
|
||||
slog.Debug("ws connection registered", "hub", h.name, "id", id, "client_count", count)
|
||||
return client
|
||||
}
|
||||
|
||||
func (h *WSHub) Unregister(client *WSClient) {
|
||||
if client == nil {
|
||||
return
|
||||
}
|
||||
h.mu.Lock()
|
||||
if current := h.clients[client.id]; current == client {
|
||||
delete(h.clients, client.id)
|
||||
}
|
||||
count := len(h.clients)
|
||||
h.mu.Unlock()
|
||||
client.Close()
|
||||
slog.Debug("ws connection unregistered", "hub", h.name, "id", client.id, "client_count", count)
|
||||
}
|
||||
|
||||
func (h *WSHub) Disconnect(id string) {
|
||||
h.mu.Lock()
|
||||
client := h.clients[id]
|
||||
if client != nil {
|
||||
delete(h.clients, id)
|
||||
}
|
||||
count := len(h.clients)
|
||||
h.mu.Unlock()
|
||||
|
||||
if client != nil {
|
||||
client.Close()
|
||||
slog.Debug("ws connection forcefully disconnected", "hub", h.name, "id", id, "client_count", count)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *WSHub) IsConnected(id string) bool {
|
||||
h.mu.RLock()
|
||||
client := h.clients[id]
|
||||
h.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-client.done:
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
func (h *WSHub) SendMessage(id string, message WSMessage) bool {
|
||||
h.mu.RLock()
|
||||
client := h.clients[id]
|
||||
h.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
ok := client.Send(message)
|
||||
if !ok {
|
||||
slog.Debug("ws send queued message failed", "hub", h.name, "id", id, "type", message.Type)
|
||||
}
|
||||
return ok
|
||||
}
|
||||
|
||||
type WSBroadcastResult struct {
|
||||
ClientCount int `json:"client_count"`
|
||||
SuccessCount int `json:"success_count"`
|
||||
FailedIDs []string `json:"failed_ids"`
|
||||
}
|
||||
|
||||
func (h *WSHub) Broadcast(message WSMessage) WSBroadcastResult {
|
||||
h.mu.RLock()
|
||||
clients := make([]*WSClient, 0, len(h.clients))
|
||||
for _, client := range h.clients {
|
||||
clients = append(clients, client)
|
||||
}
|
||||
h.mu.RUnlock()
|
||||
|
||||
var result WSBroadcastResult
|
||||
result.ClientCount = len(clients)
|
||||
for _, client := range clients {
|
||||
if client.Send(message) {
|
||||
result.SuccessCount++
|
||||
continue
|
||||
}
|
||||
result.FailedIDs = append(result.FailedIDs, client.ID())
|
||||
}
|
||||
return result
|
||||
}
|
||||
Reference in New Issue
Block a user