[优化] go 引用调整

This commit is contained in:
ryan
2026-06-06 10:26:20 +08:00
parent ee1110b752
commit 3cfefb4367
552 changed files with 1642 additions and 2185 deletions
+608
View File
@@ -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 ""
}
+320
View File
@@ -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)
}
}
}
+529
View File
@@ -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
}
+522
View File
@@ -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))
}
}
+141
View File
@@ -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
}
+500
View File
@@ -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
+254
View File
@@ -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)
}
}
+54
View File
@@ -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
}
+51
View File
@@ -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
+101
View File
@@ -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)
}
+228
View File
@@ -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")
}
}
+661
View File
@@ -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
+410
View File
@@ -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)
}
}
+242
View File
@@ -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
}
+170
View File
@@ -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
}
+105
View File
@@ -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")
}
}
+834
View File
@@ -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()
}
+420
View File
@@ -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
+569
View File
@@ -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")
}
+529
View File
@@ -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)
}
}
+41
View File
@@ -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,
})
}
+272
View File
@@ -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
}
+365
View File
@@ -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
}
+868
View File
@@ -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, `"`, `""`) + `"`
}
+415
View File
@@ -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")
}
}
+305
View File
@@ -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
}
+541
View File
@@ -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
+514
View File
@@ -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)
}
}
}
+223
View File
@@ -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
}