mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
[优化] 改名
This commit is contained in:
@@ -0,0 +1,111 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"openflare/model"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultAccessLogPageSize = 50
|
||||
maxAccessLogPageSize = 200
|
||||
)
|
||||
|
||||
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"`
|
||||
}
|
||||
|
||||
func ListAccessLogs(nodeID string, page int, pageSize int) (*AccessLogList, error) {
|
||||
normalizedPage := normalizeAccessLogPage(page)
|
||||
normalizedPageSize := normalizeAccessLogPageSize(pageSize)
|
||||
offset := normalizedPage * normalizedPageSize
|
||||
trimmedNodeID := strings.TrimSpace(nodeID)
|
||||
since := time.Now().Add(-nodeAccessLogRetentionWindow)
|
||||
logs, err := model.ListNodeAccessLogs(
|
||||
trimmedNodeID,
|
||||
since,
|
||||
offset,
|
||||
normalizedPageSize+1,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalRecords, totalIPs, err := model.CountNodeAccessLogs(trimmedNodeID, since)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodes, err := model.ListNodes()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeNames := make(map[string]string, len(nodes))
|
||||
for _, node := range nodes {
|
||||
if node == nil {
|
||||
continue
|
||||
}
|
||||
nodeNames[node.NodeID] = node.Name
|
||||
}
|
||||
hasMore := len(logs) > normalizedPageSize
|
||||
if hasMore {
|
||||
logs = logs[:normalizedPageSize]
|
||||
}
|
||||
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: normalizedPage,
|
||||
PageSize: normalizedPageSize,
|
||||
HasMore: hasMore,
|
||||
TotalRecord: totalRecords,
|
||||
TotalIP: totalIPs,
|
||||
}, 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
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"openflare/utils/geoip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var accessLogGeoProviderFactory = func() (geoip.GeoIPService, error) {
|
||||
return geoip.NewMaxMindGeoIPService()
|
||||
}
|
||||
|
||||
type accessLogRegionResolver struct {
|
||||
provider geoip.GeoIPService
|
||||
cache map[string]string
|
||||
}
|
||||
|
||||
func newAccessLogRegionResolver() (*accessLogRegionResolver, error) {
|
||||
provider, err := accessLogGeoProviderFactory()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &accessLogRegionResolver{
|
||||
provider: provider,
|
||||
cache: make(map[string]string),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Close() {
|
||||
if r == nil || r.provider == nil {
|
||||
return
|
||||
}
|
||||
if err := r.provider.Close(); err != nil {
|
||||
slog.Warn("close access log geo provider failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Resolve(rawIP string) string {
|
||||
if r == nil || r.provider == nil {
|
||||
return ""
|
||||
}
|
||||
normalizedIP := normalizeAccessLogIP(rawIP)
|
||||
if normalizedIP == "" {
|
||||
return ""
|
||||
}
|
||||
if cached, ok := r.cache[normalizedIP]; ok {
|
||||
return cached
|
||||
}
|
||||
|
||||
info, err := r.provider.GetGeoInfo(net.ParseIP(normalizedIP))
|
||||
if err != nil || info == nil {
|
||||
r.cache[normalizedIP] = ""
|
||||
return ""
|
||||
}
|
||||
|
||||
region := strings.TrimSpace(info.Name)
|
||||
if region == "" {
|
||||
region = strings.TrimSpace(info.ISOCode)
|
||||
}
|
||||
r.cache[normalizedIP] = region
|
||||
return region
|
||||
}
|
||||
|
||||
func normalizeAccessLogIP(raw string) string {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
trimmed = strings.TrimPrefix(trimmed, "[")
|
||||
trimmed = strings.TrimSuffix(trimmed, "]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(raw))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
host = strings.TrimPrefix(host, "[")
|
||||
host = strings.TrimSuffix(host, "]")
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"openflare/model"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
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,
|
||||
},
|
||||
}
|
||||
if err := model.DB.Create(&logs).Error; err != nil {
|
||||
t.Fatalf("failed to seed access logs: %v", err)
|
||||
}
|
||||
|
||||
result, err := ListAccessLogs("", 0, 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("node-a", 0, 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))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,376 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"openflare/common"
|
||||
"openflare/model"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
NodeStatusOnline = "online"
|
||||
NodeStatusOffline = "offline"
|
||||
NodeStatusPending = "pending"
|
||||
ApplyResultOK = "success"
|
||||
ApplyResultFailed = "failed"
|
||||
OpenrestyStatusHealthy = "healthy"
|
||||
OpenrestyStatusUnhealthy = "unhealthy"
|
||||
OpenrestyStatusUnknown = "unknown"
|
||||
)
|
||||
|
||||
type AgentNodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
AgentVersion string `json:"agent_version"`
|
||||
NginxVersion string `json:"nginx_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"`
|
||||
TrafficReport *AgentNodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []AgentNodeAccessLog `json:"access_logs,omitempty"`
|
||||
BufferedObservability []AgentBufferedObservabilityRecord `json:"buffered_observability,omitempty"`
|
||||
HealthEvents []AgentNodeHealthEvent `json:"health_events"`
|
||||
}
|
||||
|
||||
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 AgentConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
MainConfig string `json:"main_config"`
|
||||
RouteConfig string `json:"route_config"`
|
||||
RenderedConfig string `json:"rendered_config"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type AgentSettings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
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"`
|
||||
}
|
||||
|
||||
type NodeView struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
GeoManualOverride bool `json:"geo_manual_override"`
|
||||
AgentToken string `json:"agent_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"`
|
||||
AgentVersion string `json:"agent_version"`
|
||||
NginxVersion string `json:"nginx_version"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
Status string `json:"status"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastSeenAt time.Time `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"`
|
||||
}
|
||||
|
||||
func RegisterNode(node *model.Node, payload AgentNodePayload) (*AgentRegistrationResponse, error) {
|
||||
return RegisterNodeWithAgentToken(node, payload)
|
||||
}
|
||||
|
||||
func HeartbeatNode(node *model.Node, payload AgentNodePayload) (*HeartbeatResponse, error) {
|
||||
slog.Debug("agent heartbeat received", "node_id", node.NodeID, "current_version", strings.TrimSpace(payload.CurrentVersion))
|
||||
payload.NodeID = node.NodeID
|
||||
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
|
||||
}
|
||||
}
|
||||
refreshAgentTokenCache(node)
|
||||
persistHeartbeatObservability(node.NodeID, payload, node.LastSeenAt)
|
||||
activeConfig, err := GetActiveConfigMetaForAgent()
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
return &HeartbeatResponse{
|
||||
Node: node,
|
||||
AgentSettings: &AgentSettings{
|
||||
HeartbeatInterval: common.AgentHeartbeatInterval,
|
||||
AutoUpdate: node.AutoUpdateEnabled,
|
||||
UpdateRepo: common.AgentUpdateRepo,
|
||||
UpdateNow: updateNow,
|
||||
UpdateChannel: updateChannel.String(),
|
||||
UpdateTag: updateTag,
|
||||
RestartOpenrestyNow: restartOpenrestyNow,
|
||||
},
|
||||
ActiveConfig: activeConfig,
|
||||
}, nil
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
supportFiles = filterCertificateSupportFiles(supportFiles)
|
||||
slog.Debug("agent fetched active config", "version", version.Version, "checksum", version.Checksum)
|
||||
return &AgentConfigResponse{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
MainConfig: version.MainConfig,
|
||||
RouteConfig: version.RenderedConfig,
|
||||
RenderedConfig: version.RenderedConfig,
|
||||
SupportFiles: supportFiles,
|
||||
CreatedAt: version.CreatedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func filterCertificateSupportFiles(files []SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
filtered := make([]SupportFile, 0, len(files))
|
||||
for _, file := range files {
|
||||
path := strings.ToLower(strings.TrimSpace(file.Path))
|
||||
switch {
|
||||
case strings.HasSuffix(path, ".crt"), strings.HasSuffix(path, ".key"), strings.HasSuffix(path, ".pem"):
|
||||
filtered = append(filtered, file)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func ReportApplyLog(payload ApplyLogPayload) (*model.ApplyLog, error) {
|
||||
now := time.Now()
|
||||
payload.NodeID = strings.TrimSpace(payload.NodeID)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.Result = strings.TrimSpace(strings.ToLower(payload.Result))
|
||||
payload.Message = strings.TrimSpace(payload.Message)
|
||||
payload.Checksum = strings.TrimSpace(payload.Checksum)
|
||||
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
|
||||
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
|
||||
if payload.NodeID == "" {
|
||||
return nil, errors.New("node_id 不能为空")
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return nil, errors.New("version 不能为空")
|
||||
}
|
||||
if payload.Result != ApplyResultOK && payload.Result != ApplyResultFailed {
|
||||
return nil, errors.New("result 仅支持 success 或 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 {
|
||||
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 ListApplyLogs(nodeID string) ([]*model.ApplyLog, error) {
|
||||
return model.ListApplyLogs(strings.TrimSpace(nodeID))
|
||||
}
|
||||
|
||||
func upsertNode(payload AgentNodePayload) (*model.Node, error) {
|
||||
return nil, errors.New("不再支持匿名自动注册")
|
||||
}
|
||||
|
||||
func computeNodeStatus(node *model.Node) string {
|
||||
if node == nil {
|
||||
return NodeStatusOffline
|
||||
}
|
||||
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("agent_version", previous.AgentVersion, current.AgentVersion)
|
||||
appendIfChanged("nginx_version", previous.NginxVersion, current.NginxVersion)
|
||||
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
|
||||
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
|
||||
appendIfChanged("status", previous.Status, current.Status)
|
||||
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
|
||||
appendIfChanged("last_error", previous.LastError, current.LastError)
|
||||
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
|
||||
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
|
||||
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
|
||||
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
|
||||
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
|
||||
changes["geo_latitude"] = current.GeoLatitude
|
||||
}
|
||||
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
|
||||
changes["geo_longitude"] = current.GeoLongitude
|
||||
}
|
||||
if !previous.LastSeenAt.Equal(current.LastSeenAt) {
|
||||
changes["last_seen_at"] = current.LastSeenAt
|
||||
}
|
||||
return changes
|
||||
}
|
||||
|
||||
func coordinatesEqual(before *float64, after *float64) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return *before == *after
|
||||
}
|
||||
@@ -0,0 +1,814 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"openflare/common"
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type ReleaseResult struct {
|
||||
Version *model.ConfigVersion `json:"version"`
|
||||
Routes []*model.ProxyRoute `json:"routes"`
|
||||
}
|
||||
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
type ConfigPreviewResult struct {
|
||||
SnapshotJSON string `json:"snapshot_json"`
|
||||
MainConfig string `json:"main_config"`
|
||||
RouteConfig string `json:"route_config"`
|
||||
RenderedConfig string `json:"rendered_config"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
Checksum string `json:"checksum"`
|
||||
RouteCount int `json:"route_count"`
|
||||
}
|
||||
|
||||
type ConfigDiffResult struct {
|
||||
ActiveVersion string `json:"active_version,omitempty"`
|
||||
AddedDomains []string `json:"added_domains"`
|
||||
RemovedDomains []string `json:"removed_domains"`
|
||||
ModifiedDomains []string `json:"modified_domains"`
|
||||
MainConfigChanged bool `json:"main_config_changed"`
|
||||
ChangedOptionKeys []string `json:"changed_option_keys"`
|
||||
ChangedOptionDetails []ConfigOptionDiffItem `json:"changed_option_details"`
|
||||
}
|
||||
|
||||
type ConfigOptionDiffItem struct {
|
||||
Key string `json:"key"`
|
||||
PreviousValue string `json:"previous_value"`
|
||||
CurrentValue string `json:"current_value"`
|
||||
}
|
||||
|
||||
type snapshotRoute struct {
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id,omitempty"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
CustomHeaders []ProxyRouteCustomHeaderInput `json:"custom_headers,omitempty"`
|
||||
Remark string `json:"remark,omitempty"`
|
||||
}
|
||||
|
||||
type openRestyConfigSnapshot struct {
|
||||
WorkerProcesses string `json:"worker_processes"`
|
||||
WorkerConnections int `json:"worker_connections"`
|
||||
WorkerRlimitNofile int `json:"worker_rlimit_nofile"`
|
||||
EventsUse string `json:"events_use,omitempty"`
|
||||
EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"`
|
||||
KeepaliveTimeout int `json:"keepalive_timeout"`
|
||||
KeepaliveRequests int `json:"keepalive_requests"`
|
||||
ClientHeaderTimeout int `json:"client_header_timeout"`
|
||||
ClientBodyTimeout int `json:"client_body_timeout"`
|
||||
ClientMaxBodySize string `json:"client_max_body_size"`
|
||||
LargeClientHeaderBuffers string `json:"large_client_header_buffers"`
|
||||
SendTimeout int `json:"send_timeout"`
|
||||
ProxyConnectTimeout int `json:"proxy_connect_timeout"`
|
||||
ProxySendTimeout int `json:"proxy_send_timeout"`
|
||||
ProxyReadTimeout int `json:"proxy_read_timeout"`
|
||||
WebsocketEnabled bool `json:"websocket_enabled"`
|
||||
ProxyRequestBuffering bool `json:"proxy_request_buffering"`
|
||||
ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"`
|
||||
ProxyBuffers string `json:"proxy_buffers"`
|
||||
ProxyBufferSize string `json:"proxy_buffer_size"`
|
||||
ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"`
|
||||
GzipEnabled bool `json:"gzip_enabled"`
|
||||
GzipMinLength int `json:"gzip_min_length"`
|
||||
GzipCompLevel int `json:"gzip_comp_level"`
|
||||
CacheEnabled bool `json:"cache_enabled"`
|
||||
CachePath string `json:"cache_path,omitempty"`
|
||||
CacheLevels string `json:"cache_levels"`
|
||||
CacheInactive string `json:"cache_inactive"`
|
||||
CacheMaxSize string `json:"cache_max_size"`
|
||||
CacheKeyTemplate string `json:"cache_key_template"`
|
||||
CacheLockEnabled bool `json:"cache_lock_enabled"`
|
||||
CacheLockTimeout string `json:"cache_lock_timeout"`
|
||||
CacheUseStale string `json:"cache_use_stale"`
|
||||
}
|
||||
|
||||
type snapshotDocument struct {
|
||||
Routes []snapshotRoute `json:"routes"`
|
||||
OpenRestyConfig openRestyConfigSnapshot `json:"openresty_config"`
|
||||
}
|
||||
|
||||
type configBundle struct {
|
||||
Routes []*model.ProxyRoute
|
||||
SnapshotRoutes []snapshotRoute
|
||||
OpenRestyConfig openRestyConfigSnapshot
|
||||
SnapshotJSON string
|
||||
MainConfig string
|
||||
RouteConfig string
|
||||
SupportFiles []SupportFile
|
||||
Checksum string
|
||||
ChangedOptionKeys []string
|
||||
}
|
||||
|
||||
const (
|
||||
nginxCertDirPlaceholder = "__OPENFLARE_CERT_DIR__"
|
||||
nginxRouteConfigPlaceholder = "__OPENFLARE_ROUTE_CONFIG__"
|
||||
nginxAccessLogPlaceholder = "__OPENFLARE_ACCESS_LOG__"
|
||||
nginxLuaDirPlaceholder = "__OPENFLARE_LUA_DIR__"
|
||||
nginxObservabilityListenPlaceholder = "__OPENFLARE_OBSERVABILITY_LISTEN__"
|
||||
nginxObservabilityPortPlaceholder = "__OPENFLARE_OBSERVABILITY_PORT__"
|
||||
)
|
||||
|
||||
var requiredMainConfigTemplatePlaceholders = []string{
|
||||
"{{OpenRestyWorkerProcesses}}",
|
||||
"{{OpenRestyWorkerConnections}}",
|
||||
"{{OpenRestyWorkerRlimitNofile}}",
|
||||
"{{OpenRestyAccessLogPath}}",
|
||||
"{{OpenRestyEventsUseDirective}}",
|
||||
"{{OpenRestyEventsMultiAcceptDirective}}",
|
||||
"{{OpenRestyKeepaliveTimeout}}",
|
||||
"{{OpenRestyKeepaliveRequests}}",
|
||||
"{{OpenRestyClientHeaderTimeout}}",
|
||||
"{{OpenRestyClientBodyTimeout}}",
|
||||
"{{OpenRestyClientMaxBodySize}}",
|
||||
"{{OpenRestyLargeClientHeaderBuffers}}",
|
||||
"{{OpenRestySendTimeout}}",
|
||||
"{{OpenRestyProxyConnectTimeout}}",
|
||||
"{{OpenRestyProxySendTimeout}}",
|
||||
"{{OpenRestyProxyReadTimeout}}",
|
||||
"{{OpenRestyProxyRequestBuffering}}",
|
||||
"{{OpenRestyProxyBuffering}}",
|
||||
"{{OpenRestyProxyBuffers}}",
|
||||
"{{OpenRestyProxyBufferSize}}",
|
||||
"{{OpenRestyProxyBusyBuffersSize}}",
|
||||
"{{OpenRestyGzip}}",
|
||||
"{{OpenRestyGzipMinLength}}",
|
||||
"{{OpenRestyGzipCompLevel}}",
|
||||
"{{OpenRestyCacheBlock}}",
|
||||
"{{OpenRestyRouteConfigInclude}}",
|
||||
}
|
||||
|
||||
func ListConfigVersions() ([]*model.ConfigVersion, error) {
|
||||
return model.ListConfigVersions()
|
||||
}
|
||||
|
||||
func GetActiveConfigVersion() (*model.ConfigVersion, error) {
|
||||
return model.GetActiveConfigVersion()
|
||||
}
|
||||
|
||||
func PreviewConfigVersion() (*ConfigPreviewResult, error) {
|
||||
bundle, err := buildCurrentConfigBundle(false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ConfigPreviewResult{
|
||||
SnapshotJSON: bundle.SnapshotJSON,
|
||||
MainConfig: bundle.MainConfig,
|
||||
RouteConfig: bundle.RouteConfig,
|
||||
RenderedConfig: bundle.RouteConfig,
|
||||
SupportFiles: bundle.SupportFiles,
|
||||
Checksum: bundle.Checksum,
|
||||
RouteCount: len(bundle.Routes),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func DiffConfigVersion() (*ConfigDiffResult, error) {
|
||||
bundle, err := buildCurrentConfigBundle(false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := &ConfigDiffResult{
|
||||
AddedDomains: []string{},
|
||||
RemovedDomains: []string{},
|
||||
ModifiedDomains: []string{},
|
||||
ChangedOptionKeys: []string{},
|
||||
ChangedOptionDetails: []ConfigOptionDiffItem{},
|
||||
}
|
||||
activeVersion, err := model.GetActiveConfigVersion()
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
for _, route := range bundle.SnapshotRoutes {
|
||||
result.AddedDomains = append(result.AddedDomains, route.Domain)
|
||||
}
|
||||
result.MainConfigChanged = true
|
||||
result.ChangedOptionKeys = openRestyOptionKeys()
|
||||
result.ChangedOptionDetails = buildInitialOpenRestyOptionDiffs(bundle.OpenRestyConfig)
|
||||
return result, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
result.ActiveVersion = activeVersion.Version
|
||||
activeSnapshot, err := parseSnapshotDocument(activeVersion.SnapshotJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
currentMap := make(map[string]snapshotRoute, len(bundle.SnapshotRoutes))
|
||||
for _, route := range bundle.SnapshotRoutes {
|
||||
currentMap[route.Domain] = route
|
||||
}
|
||||
activeMap := make(map[string]snapshotRoute, len(activeSnapshot.Routes))
|
||||
for _, route := range activeSnapshot.Routes {
|
||||
activeMap[route.Domain] = route
|
||||
}
|
||||
for domain, currentRoute := range currentMap {
|
||||
activeRoute, ok := activeMap[domain]
|
||||
if !ok {
|
||||
result.AddedDomains = append(result.AddedDomains, domain)
|
||||
continue
|
||||
}
|
||||
if !snapshotRouteConfigEqual(activeRoute, currentRoute) {
|
||||
result.ModifiedDomains = append(result.ModifiedDomains, domain)
|
||||
}
|
||||
}
|
||||
for domain := range activeMap {
|
||||
if _, ok := currentMap[domain]; !ok {
|
||||
result.RemovedDomains = append(result.RemovedDomains, domain)
|
||||
}
|
||||
}
|
||||
result.MainConfigChanged = activeVersion.MainConfig != bundle.MainConfig
|
||||
result.ChangedOptionDetails = diffOpenRestyOptionDetails(activeSnapshot.OpenRestyConfig, bundle.OpenRestyConfig)
|
||||
result.ChangedOptionKeys = extractOptionDiffKeys(result.ChangedOptionDetails)
|
||||
sort.Strings(result.AddedDomains)
|
||||
sort.Strings(result.RemovedDomains)
|
||||
sort.Strings(result.ModifiedDomains)
|
||||
sort.Strings(result.ChangedOptionKeys)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func HasConfigChanges() (bool, error) {
|
||||
bundle, err := buildCurrentConfigBundle(false)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
activeVersion, err := model.GetActiveConfigVersion()
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return len(bundle.Routes) > 0, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
return activeVersion.Checksum != bundle.Checksum, nil
|
||||
}
|
||||
|
||||
func PublishConfigVersion(createdBy string) (*ReleaseResult, error) {
|
||||
bundle, err := buildCurrentConfigBundle(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(bundle.Routes) == 0 {
|
||||
return nil, errors.New("没有可发布的启用规则")
|
||||
}
|
||||
activeVersion, err := model.GetActiveConfigVersion()
|
||||
if err == nil && activeVersion.Checksum == bundle.Checksum {
|
||||
return nil, errors.New("当前规则没有变更,不能重复发布")
|
||||
}
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
supportFilesJSON, err := json.Marshal(bundle.SupportFiles)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
version, err := nextVersionNumber(time.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record := &model.ConfigVersion{
|
||||
Version: version,
|
||||
SnapshotJSON: bundle.SnapshotJSON,
|
||||
MainConfig: bundle.MainConfig,
|
||||
RenderedConfig: bundle.RouteConfig,
|
||||
SupportFilesJSON: string(supportFilesJSON),
|
||||
Checksum: bundle.Checksum,
|
||||
IsActive: true,
|
||||
CreatedBy: createdBy,
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&model.ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Create(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("版本号生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &ReleaseResult{
|
||||
Version: record,
|
||||
Routes: bundle.Routes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ActivateConfigVersion(id uint) (*model.ConfigVersion, error) {
|
||||
version, err := model.GetConfigVersionByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&model.ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(version).Update("is_active", true).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
version.IsActive = true
|
||||
return version, nil
|
||||
}
|
||||
|
||||
func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) {
|
||||
routes, err := model.GetEnabledProxyRoutes()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if requireRoutes && len(routes) == 0 {
|
||||
return nil, errors.New("没有可发布的启用规则")
|
||||
}
|
||||
snapshotRoutes, err := buildSnapshotRoutes(routes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
openRestyConfig := buildOpenRestyConfigSnapshot()
|
||||
snapshotDoc := snapshotDocument{
|
||||
Routes: snapshotRoutes,
|
||||
OpenRestyConfig: openRestyConfig,
|
||||
}
|
||||
snapshotJSON, err := json.Marshal(snapshotDoc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
routeConfig, supportFiles, err := renderRouteConfig(routes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
mainConfig := renderMainConfig(openRestyConfig)
|
||||
return &configBundle{
|
||||
Routes: routes,
|
||||
SnapshotRoutes: snapshotRoutes,
|
||||
OpenRestyConfig: openRestyConfig,
|
||||
SnapshotJSON: string(snapshotJSON),
|
||||
MainConfig: mainConfig,
|
||||
RouteConfig: routeConfig,
|
||||
SupportFiles: supportFiles,
|
||||
Checksum: checksumBundle(mainConfig, routeConfig, supportFiles),
|
||||
ChangedOptionKeys: openRestyOptionKeys(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
items := make([]snapshotRoute, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
}
|
||||
items = append(items, snapshotRoute{
|
||||
Domain: route.Domain,
|
||||
OriginURL: route.OriginURL,
|
||||
Enabled: route.Enabled,
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: route.CertID,
|
||||
RedirectHTTP: route.RedirectHTTP,
|
||||
CustomHeaders: customHeaders,
|
||||
Remark: route.Remark,
|
||||
})
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
return &snapshotDocument{Routes: []snapshotRoute{}}, nil
|
||||
}
|
||||
if strings.HasPrefix(text, "[") {
|
||||
var routes []snapshotRoute
|
||||
if err := json.Unmarshal([]byte(text), &routes); err != nil {
|
||||
return nil, errors.New("历史版本快照格式不合法")
|
||||
}
|
||||
return &snapshotDocument{Routes: normalizeSnapshotRoutes(routes)}, nil
|
||||
}
|
||||
var snapshot snapshotDocument
|
||||
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
|
||||
return nil, errors.New("历史版本快照格式不合法")
|
||||
}
|
||||
snapshot.Routes = normalizeSnapshotRoutes(snapshot.Routes)
|
||||
return &snapshot, nil
|
||||
}
|
||||
|
||||
func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
|
||||
if len(routes) == 0 {
|
||||
return []snapshotRoute{}
|
||||
}
|
||||
for index := range routes {
|
||||
normalizedHeaders, err := normalizeCustomHeaders(routes[index].CustomHeaders)
|
||||
if err == nil {
|
||||
routes[index].CustomHeaders = normalizedHeaders
|
||||
}
|
||||
}
|
||||
return routes
|
||||
}
|
||||
|
||||
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
if left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || !uintPointerEqual(left.CertID, right.CertID) {
|
||||
return false
|
||||
}
|
||||
if len(left.CustomHeaders) != len(right.CustomHeaders) {
|
||||
return false
|
||||
}
|
||||
for index := range left.CustomHeaders {
|
||||
if left.CustomHeaders[index] != right.CustomHeaders[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func buildOpenRestyConfigSnapshot() openRestyConfigSnapshot {
|
||||
return openRestyConfigSnapshot{
|
||||
WorkerProcesses: common.OpenRestyWorkerProcesses,
|
||||
WorkerConnections: common.OpenRestyWorkerConnections,
|
||||
WorkerRlimitNofile: common.OpenRestyWorkerRlimitNofile,
|
||||
EventsUse: common.OpenRestyEventsUse,
|
||||
EventsMultiAcceptEnabled: common.OpenRestyEventsMultiAcceptEnabled,
|
||||
KeepaliveTimeout: common.OpenRestyKeepaliveTimeout,
|
||||
KeepaliveRequests: common.OpenRestyKeepaliveRequests,
|
||||
ClientHeaderTimeout: common.OpenRestyClientHeaderTimeout,
|
||||
ClientBodyTimeout: common.OpenRestyClientBodyTimeout,
|
||||
ClientMaxBodySize: common.OpenRestyClientMaxBodySize,
|
||||
LargeClientHeaderBuffers: common.OpenRestyLargeClientHeaderBuffers,
|
||||
SendTimeout: common.OpenRestySendTimeout,
|
||||
ProxyConnectTimeout: common.OpenRestyProxyConnectTimeout,
|
||||
ProxySendTimeout: common.OpenRestyProxySendTimeout,
|
||||
ProxyReadTimeout: common.OpenRestyProxyReadTimeout,
|
||||
WebsocketEnabled: common.OpenRestyWebsocketEnabled,
|
||||
ProxyRequestBuffering: common.OpenRestyProxyRequestBufferingEnabled,
|
||||
ProxyBufferingEnabled: common.OpenRestyProxyBufferingEnabled,
|
||||
ProxyBuffers: common.OpenRestyProxyBuffers,
|
||||
ProxyBufferSize: common.OpenRestyProxyBufferSize,
|
||||
ProxyBusyBuffersSize: common.OpenRestyProxyBusyBuffersSize,
|
||||
GzipEnabled: common.OpenRestyGzipEnabled,
|
||||
GzipMinLength: common.OpenRestyGzipMinLength,
|
||||
GzipCompLevel: common.OpenRestyGzipCompLevel,
|
||||
CacheEnabled: common.OpenRestyCacheEnabled,
|
||||
CachePath: common.OpenRestyCachePath,
|
||||
CacheLevels: common.OpenRestyCacheLevels,
|
||||
CacheInactive: common.OpenRestyCacheInactive,
|
||||
CacheMaxSize: common.OpenRestyCacheMaxSize,
|
||||
CacheKeyTemplate: common.OpenRestyCacheKeyTemplate,
|
||||
CacheLockEnabled: common.OpenRestyCacheLockEnabled,
|
||||
CacheLockTimeout: common.OpenRestyCacheLockTimeout,
|
||||
CacheUseStale: common.OpenRestyCacheUseStale,
|
||||
}
|
||||
}
|
||||
|
||||
func diffOpenRestyOptionKeys(left openRestyConfigSnapshot, right openRestyConfigSnapshot) []string {
|
||||
details := diffOpenRestyOptionDetails(left, right)
|
||||
return extractOptionDiffKeys(details)
|
||||
}
|
||||
|
||||
func buildInitialOpenRestyOptionDiffs(current openRestyConfigSnapshot) []ConfigOptionDiffItem {
|
||||
details := diffOpenRestyOptionDetails(openRestyConfigSnapshot{}, current)
|
||||
for index := range details {
|
||||
details[index].PreviousValue = ""
|
||||
}
|
||||
return details
|
||||
}
|
||||
|
||||
func diffOpenRestyOptionDetails(left openRestyConfigSnapshot, right openRestyConfigSnapshot) []ConfigOptionDiffItem {
|
||||
changes := make([]ConfigOptionDiffItem, 0)
|
||||
appendIfChanged := func(key string, previous string, current string) {
|
||||
if previous == current {
|
||||
return
|
||||
}
|
||||
changes = append(changes, ConfigOptionDiffItem{
|
||||
Key: key,
|
||||
PreviousValue: previous,
|
||||
CurrentValue: current,
|
||||
})
|
||||
}
|
||||
appendIfChanged("OpenRestyWorkerProcesses", left.WorkerProcesses, right.WorkerProcesses)
|
||||
appendIfChanged("OpenRestyWorkerConnections", fmt.Sprintf("%d", left.WorkerConnections), fmt.Sprintf("%d", right.WorkerConnections))
|
||||
appendIfChanged("OpenRestyWorkerRlimitNofile", fmt.Sprintf("%d", left.WorkerRlimitNofile), fmt.Sprintf("%d", right.WorkerRlimitNofile))
|
||||
appendIfChanged("OpenRestyEventsUse", left.EventsUse, right.EventsUse)
|
||||
appendIfChanged("OpenRestyEventsMultiAcceptEnabled", fmt.Sprintf("%t", left.EventsMultiAcceptEnabled), fmt.Sprintf("%t", right.EventsMultiAcceptEnabled))
|
||||
appendIfChanged("OpenRestyKeepaliveTimeout", fmt.Sprintf("%d", left.KeepaliveTimeout), fmt.Sprintf("%d", right.KeepaliveTimeout))
|
||||
appendIfChanged("OpenRestyKeepaliveRequests", fmt.Sprintf("%d", left.KeepaliveRequests), fmt.Sprintf("%d", right.KeepaliveRequests))
|
||||
appendIfChanged("OpenRestyClientHeaderTimeout", fmt.Sprintf("%d", left.ClientHeaderTimeout), fmt.Sprintf("%d", right.ClientHeaderTimeout))
|
||||
appendIfChanged("OpenRestyClientBodyTimeout", fmt.Sprintf("%d", left.ClientBodyTimeout), fmt.Sprintf("%d", right.ClientBodyTimeout))
|
||||
appendIfChanged("OpenRestyClientMaxBodySize", left.ClientMaxBodySize, right.ClientMaxBodySize)
|
||||
appendIfChanged("OpenRestyLargeClientHeaderBuffers", left.LargeClientHeaderBuffers, right.LargeClientHeaderBuffers)
|
||||
appendIfChanged("OpenRestySendTimeout", fmt.Sprintf("%d", left.SendTimeout), fmt.Sprintf("%d", right.SendTimeout))
|
||||
appendIfChanged("OpenRestyProxyConnectTimeout", fmt.Sprintf("%d", left.ProxyConnectTimeout), fmt.Sprintf("%d", right.ProxyConnectTimeout))
|
||||
appendIfChanged("OpenRestyProxySendTimeout", fmt.Sprintf("%d", left.ProxySendTimeout), fmt.Sprintf("%d", right.ProxySendTimeout))
|
||||
appendIfChanged("OpenRestyProxyReadTimeout", fmt.Sprintf("%d", left.ProxyReadTimeout), fmt.Sprintf("%d", right.ProxyReadTimeout))
|
||||
appendIfChanged("OpenRestyWebsocketEnabled", fmt.Sprintf("%t", left.WebsocketEnabled), fmt.Sprintf("%t", right.WebsocketEnabled))
|
||||
appendIfChanged("OpenRestyProxyRequestBufferingEnabled", fmt.Sprintf("%t", left.ProxyRequestBuffering), fmt.Sprintf("%t", right.ProxyRequestBuffering))
|
||||
appendIfChanged("OpenRestyProxyBufferingEnabled", fmt.Sprintf("%t", left.ProxyBufferingEnabled), fmt.Sprintf("%t", right.ProxyBufferingEnabled))
|
||||
appendIfChanged("OpenRestyProxyBuffers", left.ProxyBuffers, right.ProxyBuffers)
|
||||
appendIfChanged("OpenRestyProxyBufferSize", left.ProxyBufferSize, right.ProxyBufferSize)
|
||||
appendIfChanged("OpenRestyProxyBusyBuffersSize", left.ProxyBusyBuffersSize, right.ProxyBusyBuffersSize)
|
||||
appendIfChanged("OpenRestyGzipEnabled", fmt.Sprintf("%t", left.GzipEnabled), fmt.Sprintf("%t", right.GzipEnabled))
|
||||
appendIfChanged("OpenRestyGzipMinLength", fmt.Sprintf("%d", left.GzipMinLength), fmt.Sprintf("%d", right.GzipMinLength))
|
||||
appendIfChanged("OpenRestyGzipCompLevel", fmt.Sprintf("%d", left.GzipCompLevel), fmt.Sprintf("%d", right.GzipCompLevel))
|
||||
appendIfChanged("OpenRestyCacheEnabled", fmt.Sprintf("%t", left.CacheEnabled), fmt.Sprintf("%t", right.CacheEnabled))
|
||||
appendIfChanged("OpenRestyCachePath", left.CachePath, right.CachePath)
|
||||
appendIfChanged("OpenRestyCacheLevels", left.CacheLevels, right.CacheLevels)
|
||||
appendIfChanged("OpenRestyCacheInactive", left.CacheInactive, right.CacheInactive)
|
||||
appendIfChanged("OpenRestyCacheMaxSize", left.CacheMaxSize, right.CacheMaxSize)
|
||||
appendIfChanged("OpenRestyCacheKeyTemplate", left.CacheKeyTemplate, right.CacheKeyTemplate)
|
||||
appendIfChanged("OpenRestyCacheLockEnabled", fmt.Sprintf("%t", left.CacheLockEnabled), fmt.Sprintf("%t", right.CacheLockEnabled))
|
||||
appendIfChanged("OpenRestyCacheLockTimeout", left.CacheLockTimeout, right.CacheLockTimeout)
|
||||
appendIfChanged("OpenRestyCacheUseStale", left.CacheUseStale, right.CacheUseStale)
|
||||
return changes
|
||||
}
|
||||
|
||||
func extractOptionDiffKeys(details []ConfigOptionDiffItem) []string {
|
||||
keys := make([]string, 0, len(details))
|
||||
for _, item := range details {
|
||||
keys = append(keys, item.Key)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func openRestyOptionKeys() []string {
|
||||
return []string{
|
||||
"OpenRestyWorkerProcesses",
|
||||
"OpenRestyWorkerConnections",
|
||||
"OpenRestyWorkerRlimitNofile",
|
||||
"OpenRestyEventsUse",
|
||||
"OpenRestyEventsMultiAcceptEnabled",
|
||||
"OpenRestyKeepaliveTimeout",
|
||||
"OpenRestyKeepaliveRequests",
|
||||
"OpenRestyClientHeaderTimeout",
|
||||
"OpenRestyClientBodyTimeout",
|
||||
"OpenRestyClientMaxBodySize",
|
||||
"OpenRestyLargeClientHeaderBuffers",
|
||||
"OpenRestySendTimeout",
|
||||
"OpenRestyProxyConnectTimeout",
|
||||
"OpenRestyProxySendTimeout",
|
||||
"OpenRestyProxyReadTimeout",
|
||||
"OpenRestyWebsocketEnabled",
|
||||
"OpenRestyProxyRequestBufferingEnabled",
|
||||
"OpenRestyProxyBufferingEnabled",
|
||||
"OpenRestyProxyBuffers",
|
||||
"OpenRestyProxyBufferSize",
|
||||
"OpenRestyProxyBusyBuffersSize",
|
||||
"OpenRestyGzipEnabled",
|
||||
"OpenRestyGzipMinLength",
|
||||
"OpenRestyGzipCompLevel",
|
||||
"OpenRestyCacheEnabled",
|
||||
"OpenRestyCachePath",
|
||||
"OpenRestyCacheLevels",
|
||||
"OpenRestyCacheInactive",
|
||||
"OpenRestyCacheMaxSize",
|
||||
"OpenRestyCacheKeyTemplate",
|
||||
"OpenRestyCacheLockEnabled",
|
||||
"OpenRestyCacheLockTimeout",
|
||||
"OpenRestyCacheUseStale",
|
||||
}
|
||||
}
|
||||
|
||||
func renderRouteConfig(routes []*model.ProxyRoute) (string, []SupportFile, error) {
|
||||
var builder strings.Builder
|
||||
builder.WriteString("# This file is generated by OpenFlare. Do not edit manually.\n")
|
||||
supportFiles := make([]SupportFile, 0)
|
||||
for _, route := range routes {
|
||||
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPProxyServer(route.Domain, route.OriginURL, customHeaders))
|
||||
continue
|
||||
}
|
||||
if route.CertID == nil || *route.CertID == 0 {
|
||||
return "", nil, fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
certificate, err := model.GetTLSCertificateByID(*route.CertID)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("路由 %s 关联证书不存在", route.Domain)
|
||||
}
|
||||
supportFiles = append(supportFiles,
|
||||
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
|
||||
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)},
|
||||
)
|
||||
if route.RedirectHTTP {
|
||||
builder.WriteString(renderHTTPRedirectServer(route.Domain))
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(route.Domain, route.OriginURL, customHeaders))
|
||||
}
|
||||
builder.WriteString(renderHTTPSServer(route.Domain, route.OriginURL, certificate.ID, customHeaders))
|
||||
}
|
||||
return builder.String(), dedupeSupportFiles(supportFiles), nil
|
||||
}
|
||||
|
||||
func renderMainConfig(cfg openRestyConfigSnapshot) string {
|
||||
templateText := common.OpenRestyMainConfigTemplate
|
||||
if strings.TrimSpace(templateText) == "" {
|
||||
templateText = defaultOpenRestyMainConfigTemplate()
|
||||
}
|
||||
return renderMainConfigTemplate(templateText, cfg)
|
||||
}
|
||||
|
||||
func ValidateOpenRestyMainConfigTemplate(templateText string) error {
|
||||
trimmed := strings.TrimSpace(templateText)
|
||||
if trimmed == "" {
|
||||
return errors.New("OpenRestyMainConfigTemplate 不能为空")
|
||||
}
|
||||
for _, placeholder := range requiredMainConfigTemplatePlaceholders {
|
||||
if !strings.Contains(trimmed, placeholder) {
|
||||
return fmt.Errorf("OpenRestyMainConfigTemplate 必须保留占位符 %s", placeholder)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func defaultOpenRestyMainConfigTemplate() string {
|
||||
return common.OpenRestyMainConfigTemplate
|
||||
}
|
||||
|
||||
func renderMainConfigTemplate(templateText string, cfg openRestyConfigSnapshot) string {
|
||||
replacer := strings.NewReplacer(
|
||||
"{{OpenRestyWorkerProcesses}}", cfg.WorkerProcesses,
|
||||
"{{OpenRestyWorkerConnections}}", fmt.Sprintf("%d", cfg.WorkerConnections),
|
||||
"{{OpenRestyWorkerRlimitNofile}}", fmt.Sprintf("%d", cfg.WorkerRlimitNofile),
|
||||
"{{OpenRestyAccessLogPath}}", nginxAccessLogPlaceholder,
|
||||
"{{OpenRestyEventsUseDirective}}", renderTemplateDirective(cfg.EventsUse != "", fmt.Sprintf("use %s;", cfg.EventsUse)),
|
||||
"{{OpenRestyEventsMultiAcceptDirective}}", renderTemplateDirective(cfg.EventsMultiAcceptEnabled, "multi_accept on;"),
|
||||
"{{OpenRestyKeepaliveTimeout}}", fmt.Sprintf("%d", cfg.KeepaliveTimeout),
|
||||
"{{OpenRestyKeepaliveRequests}}", fmt.Sprintf("%d", cfg.KeepaliveRequests),
|
||||
"{{OpenRestyClientHeaderTimeout}}", fmt.Sprintf("%d", cfg.ClientHeaderTimeout),
|
||||
"{{OpenRestyClientBodyTimeout}}", fmt.Sprintf("%d", cfg.ClientBodyTimeout),
|
||||
"{{OpenRestyClientMaxBodySize}}", cfg.ClientMaxBodySize,
|
||||
"{{OpenRestyLargeClientHeaderBuffers}}", cfg.LargeClientHeaderBuffers,
|
||||
"{{OpenRestySendTimeout}}", fmt.Sprintf("%d", cfg.SendTimeout),
|
||||
"{{OpenRestyProxyConnectTimeout}}", fmt.Sprintf("%d", cfg.ProxyConnectTimeout),
|
||||
"{{OpenRestyProxySendTimeout}}", fmt.Sprintf("%d", cfg.ProxySendTimeout),
|
||||
"{{OpenRestyProxyReadTimeout}}", fmt.Sprintf("%d", cfg.ProxyReadTimeout),
|
||||
"{{OpenRestyProxyRequestBuffering}}", onOff(cfg.ProxyRequestBuffering),
|
||||
"{{OpenRestyProxyBuffering}}", onOff(cfg.ProxyBufferingEnabled),
|
||||
"{{OpenRestyProxyBuffers}}", cfg.ProxyBuffers,
|
||||
"{{OpenRestyProxyBufferSize}}", cfg.ProxyBufferSize,
|
||||
"{{OpenRestyProxyBusyBuffersSize}}", cfg.ProxyBusyBuffersSize,
|
||||
"{{OpenRestyGzip}}", onOff(cfg.GzipEnabled),
|
||||
"{{OpenRestyGzipMinLength}}", fmt.Sprintf("%d", cfg.GzipMinLength),
|
||||
"{{OpenRestyGzipCompLevel}}", fmt.Sprintf("%d", cfg.GzipCompLevel),
|
||||
"{{OpenRestyCacheBlock}}", renderOpenRestyCacheTemplateBlock(cfg),
|
||||
"{{OpenRestyRouteConfigInclude}}", nginxRouteConfigPlaceholder,
|
||||
)
|
||||
return replacer.Replace(templateText)
|
||||
}
|
||||
|
||||
func renderTemplateDirective(enabled bool, statement string) string {
|
||||
if !enabled {
|
||||
return ""
|
||||
}
|
||||
return fmt.Sprintf(" %s\n", statement)
|
||||
}
|
||||
|
||||
func renderOpenRestyCacheTemplateBlock(cfg openRestyConfigSnapshot) string {
|
||||
lines := make([]string, 0, 8)
|
||||
if !cfg.CacheEnabled {
|
||||
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
|
||||
return strings.Join(lines, "")
|
||||
}
|
||||
lines = append(lines, strings.Join([]string{
|
||||
fmt.Sprintf(" proxy_cache_path %s levels=%s keys_zone=openflare_cache:10m inactive=%s max_size=%s;", cfg.CachePath, cfg.CacheLevels, cfg.CacheInactive, cfg.CacheMaxSize),
|
||||
fmt.Sprintf(" proxy_cache_key \"%s\";", cfg.CacheKeyTemplate),
|
||||
fmt.Sprintf(" proxy_cache_lock %s;", onOff(cfg.CacheLockEnabled)),
|
||||
fmt.Sprintf(" proxy_cache_lock_timeout %s;", cfg.CacheLockTimeout),
|
||||
fmt.Sprintf(" proxy_cache_use_stale %s;", cfg.CacheUseStale),
|
||||
"",
|
||||
}, "\n"))
|
||||
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
|
||||
return strings.Join(lines, "")
|
||||
}
|
||||
|
||||
func onOff(value bool) string {
|
||||
if value {
|
||||
return "on"
|
||||
}
|
||||
return "off"
|
||||
}
|
||||
|
||||
func uintPointerEqual(left *uint, right *uint) bool {
|
||||
if left == nil || right == nil {
|
||||
return left == nil && right == nil
|
||||
}
|
||||
return *left == *right
|
||||
}
|
||||
|
||||
func checksum(content string) string {
|
||||
sum := sha256.Sum256([]byte(content))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func checksumBundle(mainConfig string, routeConfig string, supportFiles []SupportFile) string {
|
||||
var builder strings.Builder
|
||||
builder.WriteString(mainConfig)
|
||||
builder.WriteString("\n--route-config--\n")
|
||||
builder.WriteString(routeConfig)
|
||||
builder.WriteString("\n--support-files--\n")
|
||||
files := dedupeSupportFiles(supportFiles)
|
||||
sort.Slice(files, func(i int, j int) bool {
|
||||
return files[i].Path < files[j].Path
|
||||
})
|
||||
for _, file := range files {
|
||||
builder.WriteString(file.Path)
|
||||
builder.WriteString("\n")
|
||||
builder.WriteString(file.Content)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
return checksum(builder.String())
|
||||
}
|
||||
|
||||
func nextVersionNumber(now time.Time) (string, error) {
|
||||
prefix := now.Format("20060102")
|
||||
var count int64
|
||||
if err := model.DB.Model(&model.ConfigVersion{}).Where("version LIKE ?", prefix+"-%").Count(&count).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
return fmt.Sprintf("%s-%03d", prefix, count+1), nil
|
||||
}
|
||||
|
||||
func renderHTTPProxyServer(domain string, originURL string, customHeaders []ProxyRouteCustomHeaderInput) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n location / {\n%s proxy_pass %s;\n }\n}\n\n", domain, renderProxyHeaderBlock(customHeaders), originURL)
|
||||
}
|
||||
|
||||
func renderHTTPRedirectServer(domain string) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", domain)
|
||||
}
|
||||
|
||||
func renderHTTPSServer(domain string, originURL string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput) string {
|
||||
certPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateCertFileName(certificateID))
|
||||
keyPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateKeyFileName(certificateID))
|
||||
return fmt.Sprintf("server {\n listen 443 ssl;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n\n location / {\n%s proxy_pass %s;\n }\n}\n\n", domain, certPath, keyPath, renderProxyHeaderBlock(customHeaders), originURL)
|
||||
}
|
||||
|
||||
func renderProxyHeaderBlock(customHeaders []ProxyRouteCustomHeaderInput) string {
|
||||
var builder strings.Builder
|
||||
builder.WriteString(" proxy_set_header Host $host;\n")
|
||||
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
|
||||
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
|
||||
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
|
||||
if common.OpenRestyWebsocketEnabled {
|
||||
builder.WriteString(" proxy_http_version 1.1;\n")
|
||||
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
|
||||
builder.WriteString(" proxy_set_header Connection $http_connection;\n")
|
||||
}
|
||||
for _, header := range customHeaders {
|
||||
builder.WriteString(fmt.Sprintf(" proxy_set_header %s %s;\n", header.Key, quoteNginxHeaderValue(header.Value)))
|
||||
}
|
||||
if common.OpenRestyCacheEnabled {
|
||||
builder.WriteString(" proxy_cache openflare_cache;\n")
|
||||
}
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func quoteNginxHeaderValue(value string) string {
|
||||
escaped := strings.ReplaceAll(value, `\`, `\\`)
|
||||
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
|
||||
return fmt.Sprintf(`"%s"`, escaped)
|
||||
}
|
||||
|
||||
func certificateCertFileName(id uint) string {
|
||||
return fmt.Sprintf("%d.crt", id)
|
||||
}
|
||||
|
||||
func certificateKeyFileName(id uint) string {
|
||||
return fmt.Sprintf("%d.key", id)
|
||||
}
|
||||
|
||||
func normalizePEM(content string) string {
|
||||
return strings.TrimSpace(content) + "\n"
|
||||
}
|
||||
|
||||
func dedupeSupportFiles(files []SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
unique := make(map[string]SupportFile, len(files))
|
||||
for _, file := range files {
|
||||
unique[file.Path] = file
|
||||
}
|
||||
result := make([]SupportFile, 0, len(unique))
|
||||
for _, file := range unique {
|
||||
result = append(result, file)
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"time"
|
||||
)
|
||||
|
||||
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 time.Time `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
|
||||
}
|
||||
|
||||
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),
|
||||
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: node.LastSeenAt,
|
||||
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,50 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"openflare/utils/geoip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
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,64 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"net"
|
||||
"openflare/utils/geoip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,458 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"openflare/common"
|
||||
"openflare/model"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCreateTLSCertificateAndRenderHTTPSConfig(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
certPEM, keyPEM := generateCertificatePair(t, []string{"app.example.com"})
|
||||
certificate, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "app-example",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
Remark: "test cert",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTLSCertificate failed: %v", err)
|
||||
}
|
||||
if certificate.NotAfter.Before(certificate.NotBefore) {
|
||||
t.Fatal("expected certificate validity period to be parsed")
|
||||
}
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "app.example.com",
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
EnableHTTPS: true,
|
||||
CertID: &certificate.ID,
|
||||
RedirectHTTP: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if !route.EnableHTTPS || route.CertID == nil {
|
||||
t.Fatal("expected https fields to be persisted")
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root")
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, "include __OPENFLARE_ROUTE_CONFIG__;") {
|
||||
t.Fatal("expected main config to include managed route config placeholder")
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, "access_log __OPENFLARE_ACCESS_LOG__ openflare_json;") {
|
||||
t.Fatal("expected main config to include managed access log placeholder")
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, "log_by_lua_file __OPENFLARE_LUA_DIR__/log.lua;") {
|
||||
t.Fatal("expected main config to include managed openresty lua log hook")
|
||||
}
|
||||
if !strings.Contains(result.Version.MainConfig, "listen __OPENFLARE_OBSERVABILITY_LISTEN__;") {
|
||||
t.Fatal("expected main config to include managed openresty observability listen placeholder")
|
||||
}
|
||||
if strings.Contains(result.Version.MainConfig, "allow 127.0.0.1;") {
|
||||
t.Fatal("expected main config to avoid hard-coded allow rules on observability server")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "listen 443 ssl;") {
|
||||
t.Fatal("expected rendered config to include https server block")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "return 301 https://$host$request_uri;") {
|
||||
t.Fatal("expected rendered config to include http redirect")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "__OPENFLARE_CERT_DIR__/") {
|
||||
t.Fatal("expected rendered config to keep cert dir placeholder for certificates")
|
||||
}
|
||||
if !strings.Contains(result.Version.SupportFilesJSON, ".crt") || !strings.Contains(result.Version.SupportFilesJSON, ".key") {
|
||||
t.Fatal("expected support files to contain certificate and key")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProxyRouteRejectsHTTPSWithoutCertificate(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "secure.example.com",
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
EnableHTTPS: true,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "必须选择证书") {
|
||||
t.Fatalf("expected certificate validation error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionRendersCustomHeaders(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
if err := model.UpdateOption("OpenRestyWebsocketEnabled", "true"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyWebsocketEnabled failed: %v", err)
|
||||
}
|
||||
|
||||
_, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "custom.example.com",
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
CustomHeaders: []ProxyRouteCustomHeaderInput{
|
||||
{Key: "X-Trace-Id", Value: "$request_id"},
|
||||
{Key: "X-Env", Value: "staging edge"},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
|
||||
result, err := PublishConfigVersion("root")
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, `proxy_set_header X-Trace-Id "$request_id";`) {
|
||||
t.Fatal("expected rendered config to include custom header")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, `proxy_set_header X-Env "staging edge";`) {
|
||||
t.Fatal("expected rendered config to include quoted custom header value")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, "custom_headers") {
|
||||
t.Fatal("expected snapshot to include custom headers")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "proxy_http_version 1.1;") {
|
||||
t.Fatal("expected rendered config to enable HTTP/1.1 proxying for websocket upgrades")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "proxy_set_header Upgrade $http_upgrade;") {
|
||||
t.Fatal("expected rendered config to forward websocket upgrade header")
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "proxy_set_header Connection $http_connection;") {
|
||||
t.Fatal("expected rendered config to forward websocket connection header")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreviewConfigVersionCanDisableWebsocketHeaders(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "ws-off.example.com",
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
if err := model.UpdateOption("OpenRestyWebsocketEnabled", "false"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyWebsocketEnabled failed: %v", err)
|
||||
}
|
||||
|
||||
preview, err := PreviewConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("PreviewConfigVersion failed: %v", err)
|
||||
}
|
||||
if strings.Contains(preview.RenderedConfig, "proxy_http_version 1.1;") {
|
||||
t.Fatal("expected preview config to omit websocket proxy_http_version when disabled")
|
||||
}
|
||||
if strings.Contains(preview.RenderedConfig, "proxy_set_header Upgrade $http_upgrade;") {
|
||||
t.Fatal("expected preview config to omit websocket upgrade header when disabled")
|
||||
}
|
||||
if strings.Contains(preview.RenderedConfig, "proxy_set_header Connection $http_connection;") {
|
||||
t.Fatal("expected preview config to omit websocket connection header when disabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreviewAndDiffConfigVersion(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
if err := model.UpdateOption("OpenRestyWebsocketEnabled", "true"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyWebsocketEnabled failed: %v", err)
|
||||
}
|
||||
|
||||
stableRoute, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "stable.example.com",
|
||||
OriginURL: "https://origin-a.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute stable failed: %v", err)
|
||||
}
|
||||
modifiedRoute, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "api.example.com",
|
||||
OriginURL: "https://origin-api-a.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute modified failed: %v", err)
|
||||
}
|
||||
removedRoute, err := CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "old.example.com",
|
||||
OriginURL: "https://origin-old.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute removed failed: %v", err)
|
||||
}
|
||||
if _, err = PublishConfigVersion("root"); err != nil {
|
||||
t.Fatalf("initial PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
|
||||
if _, err = UpdateProxyRoute(modifiedRoute.ID, ProxyRouteInput{
|
||||
Domain: "api.example.com",
|
||||
OriginURL: "https://origin-api-b.internal",
|
||||
Enabled: true,
|
||||
CustomHeaders: []ProxyRouteCustomHeaderInput{
|
||||
{Key: "X-Release", Value: "candidate"},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateProxyRoute failed: %v", err)
|
||||
}
|
||||
if _, err = UpdateProxyRoute(removedRoute.ID, ProxyRouteInput{
|
||||
Domain: "old.example.com",
|
||||
OriginURL: "https://origin-old.internal",
|
||||
Enabled: false,
|
||||
}); err != nil {
|
||||
t.Fatalf("disable removed route failed: %v", err)
|
||||
}
|
||||
if _, err = CreateProxyRoute(ProxyRouteInput{
|
||||
Domain: "new.example.com",
|
||||
OriginURL: "https://origin-new.internal",
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("CreateProxyRoute new failed: %v", err)
|
||||
}
|
||||
if _, err = UpdateProxyRoute(stableRoute.ID, ProxyRouteInput{
|
||||
Domain: stableRoute.Domain,
|
||||
OriginURL: stableRoute.OriginURL,
|
||||
Enabled: true,
|
||||
Remark: "remark only change",
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateProxyRoute stable failed: %v", err)
|
||||
}
|
||||
|
||||
preview, err := PreviewConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("PreviewConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "include __OPENFLARE_ROUTE_CONFIG__;") {
|
||||
t.Fatal("expected preview main config to include managed route config placeholder")
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "log_by_lua_file __OPENFLARE_LUA_DIR__/log.lua;") {
|
||||
t.Fatal("expected preview main config to include managed openresty lua log hook")
|
||||
}
|
||||
if !strings.Contains(preview.RenderedConfig, `proxy_set_header X-Release "candidate";`) {
|
||||
t.Fatal("expected preview config to include modified custom header")
|
||||
}
|
||||
if preview.RouteCount != 3 {
|
||||
t.Fatalf("expected 3 enabled routes in preview, got %d", preview.RouteCount)
|
||||
}
|
||||
|
||||
diff, err := DiffConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("DiffConfigVersion failed: %v", err)
|
||||
}
|
||||
if len(diff.AddedDomains) != 1 || diff.AddedDomains[0] != "new.example.com" {
|
||||
t.Fatalf("unexpected added domains: %#v", diff.AddedDomains)
|
||||
}
|
||||
if len(diff.RemovedDomains) != 1 || diff.RemovedDomains[0] != "old.example.com" {
|
||||
t.Fatalf("unexpected removed domains: %#v", diff.RemovedDomains)
|
||||
}
|
||||
if len(diff.ModifiedDomains) != 1 || diff.ModifiedDomains[0] != "api.example.com" {
|
||||
t.Fatalf("unexpected modified domains: %#v", diff.ModifiedDomains)
|
||||
}
|
||||
if diff.MainConfigChanged {
|
||||
t.Fatal("expected main config to remain unchanged when only routes change")
|
||||
}
|
||||
|
||||
if err = model.UpdateOption("OpenRestyProxyReadTimeout", "120"); err != nil {
|
||||
t.Fatalf("UpdateOption failed: %v", err)
|
||||
}
|
||||
if err = model.UpdateOption("OpenRestyWebsocketEnabled", "false"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyWebsocketEnabled failed: %v", err)
|
||||
}
|
||||
diff, err = DiffConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("DiffConfigVersion after option change failed: %v", err)
|
||||
}
|
||||
if !diff.MainConfigChanged {
|
||||
t.Fatal("expected main config change after OpenResty option update")
|
||||
}
|
||||
if len(diff.ChangedOptionKeys) == 0 || diff.ChangedOptionKeys[0] == "" {
|
||||
t.Fatal("expected changed OpenResty option keys to be reported")
|
||||
}
|
||||
if len(diff.ChangedOptionDetails) == 0 {
|
||||
t.Fatal("expected changed OpenResty option details to be reported")
|
||||
}
|
||||
found := false
|
||||
foundWebsocket := false
|
||||
for _, item := range diff.ChangedOptionDetails {
|
||||
if item.Key == "OpenRestyProxyReadTimeout" {
|
||||
found = true
|
||||
if item.PreviousValue != "60" || item.CurrentValue != "120" {
|
||||
t.Fatalf("unexpected option diff values: %+v", item)
|
||||
}
|
||||
}
|
||||
if item.Key == "OpenRestyWebsocketEnabled" {
|
||||
foundWebsocket = true
|
||||
if item.PreviousValue != "true" || item.CurrentValue != "false" {
|
||||
t.Fatalf("unexpected websocket option diff values: %+v", item)
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("expected OpenRestyProxyReadTimeout diff detail")
|
||||
}
|
||||
if !foundWebsocket {
|
||||
t.Fatal("expected OpenRestyWebsocketEnabled diff detail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateTLSCertificateRejectsInvalidPEM(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "broken-cert",
|
||||
CertPEM: "invalid",
|
||||
KeyPEM: "invalid",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid pem to fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenRestyMainConfigTemplateRenderAndValidate(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
customTemplate := strings.ReplaceAll(
|
||||
common.OpenRestyMainConfigTemplate,
|
||||
"pid logs/nginx.pid;",
|
||||
"pid logs/nginx.pid;\nworker_shutdown_timeout 10s;",
|
||||
)
|
||||
if err := ValidateOpenRestyMainConfigTemplate(customTemplate); err != nil {
|
||||
t.Fatalf("ValidateOpenRestyMainConfigTemplate failed: %v", err)
|
||||
}
|
||||
if err := model.UpdateOption("OpenRestyMainConfigTemplate", customTemplate); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyMainConfigTemplate failed: %v", err)
|
||||
}
|
||||
|
||||
preview, err := PreviewConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("PreviewConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "worker_shutdown_timeout 10s;") {
|
||||
t.Fatal("expected preview main config to include custom template content")
|
||||
}
|
||||
if strings.Contains(preview.MainConfig, "{{OpenRestyWorkerProcesses}}") {
|
||||
t.Fatal("expected preview main config placeholders to be rendered")
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "include __OPENFLARE_ROUTE_CONFIG__;") {
|
||||
t.Fatal("expected preview main config to preserve managed route include")
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "access_log __OPENFLARE_ACCESS_LOG__ openflare_json;") {
|
||||
t.Fatal("expected preview main config to preserve managed access log placeholder")
|
||||
}
|
||||
|
||||
invalidTemplate := strings.ReplaceAll(
|
||||
common.OpenRestyMainConfigTemplate,
|
||||
"{{OpenRestyRouteConfigInclude}}",
|
||||
"",
|
||||
)
|
||||
if err := ValidateOpenRestyMainConfigTemplate(invalidTemplate); err == nil {
|
||||
t.Fatal("expected template without managed route placeholder to fail validation")
|
||||
}
|
||||
|
||||
invalidTemplate = strings.ReplaceAll(
|
||||
common.OpenRestyMainConfigTemplate,
|
||||
"{{OpenRestyAccessLogPath}}",
|
||||
"",
|
||||
)
|
||||
if err := ValidateOpenRestyMainConfigTemplate(invalidTemplate); err == nil {
|
||||
t.Fatal("expected template without managed access log placeholder to fail validation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenRestyCommonRequestOptionsRender(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
if err := model.UpdateOption("OpenRestyClientMaxBodySize", "128m"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyClientMaxBodySize failed: %v", err)
|
||||
}
|
||||
if err := model.UpdateOption("OpenRestyLargeClientHeaderBuffers", "8 32k"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyLargeClientHeaderBuffers failed: %v", err)
|
||||
}
|
||||
if err := model.UpdateOption("OpenRestyProxyRequestBufferingEnabled", "false"); err != nil {
|
||||
t.Fatalf("UpdateOption OpenRestyProxyRequestBufferingEnabled failed: %v", err)
|
||||
}
|
||||
|
||||
preview, err := PreviewConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("PreviewConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "client_max_body_size 128m;") {
|
||||
t.Fatal("expected preview main config to include client_max_body_size")
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "large_client_header_buffers 8 32k;") {
|
||||
t.Fatal("expected preview main config to include large_client_header_buffers")
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "proxy_request_buffering off;") {
|
||||
t.Fatal("expected preview main config to include proxy_request_buffering off")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenRestyProxyRequestBufferingDefaultsToOff(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
preview, err := PreviewConfigVersion()
|
||||
if err != nil {
|
||||
t.Fatalf("PreviewConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(preview.MainConfig, "proxy_request_buffering off;") {
|
||||
t.Fatal("expected preview main config to default proxy_request_buffering to off")
|
||||
}
|
||||
}
|
||||
|
||||
func setupServiceTestDB(t *testing.T) {
|
||||
t.Helper()
|
||||
nodeAgentTokenCache.reset()
|
||||
common.SQLitePath = filepath.Join(t.TempDir(), "service.db")
|
||||
if err := model.InitDB(); err != nil {
|
||||
t.Fatalf("failed to init db: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
nodeAgentTokenCache.reset()
|
||||
if err := model.CloseDB(); err != nil {
|
||||
t.Fatalf("failed to close db: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func generateCertificatePair(t *testing.T, dnsNames []string) (string, string) {
|
||||
t.Helper()
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateKey failed: %v", err)
|
||||
}
|
||||
template := &x509.Certificate{
|
||||
Subject: pkix.Name{
|
||||
CommonName: dnsNames[0],
|
||||
},
|
||||
DNSNames: dnsNames,
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
IsCA: false,
|
||||
SerialNumber: big.NewInt(time.Now().UnixNano()),
|
||||
}
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateCertificate failed: %v", err)
|
||||
}
|
||||
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
|
||||
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
|
||||
return string(certPEM), string(keyPEM)
|
||||
}
|
||||
@@ -0,0 +1,227 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
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 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 isUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return domain, nil
|
||||
}
|
||||
|
||||
func DeleteManagedDomain(id uint) error {
|
||||
domain, err := model.GetManagedDomainByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return domain.Delete()
|
||||
}
|
||||
|
||||
func MatchManagedDomainCertificate(rawDomain string) (*ManagedDomainMatchResult, error) {
|
||||
domain := normalizeManagedDomain(rawDomain)
|
||||
if err := validateManagedDomainPattern(domain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
managedDomains, err := model.ListEnabledManagedDomainsWithCertificate()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
candidates := make([]ManagedDomainMatchCandidate, 0)
|
||||
for _, item := range managedDomains {
|
||||
if item.CertID == nil || *item.CertID == 0 {
|
||||
continue
|
||||
}
|
||||
matchType := detectManagedDomainMatchType(item.Domain, domain)
|
||||
if matchType == "" {
|
||||
continue
|
||||
}
|
||||
certificate, err := model.GetTLSCertificateByID(*item.CertID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("托管域名 %s 关联证书不存在", item.Domain)
|
||||
}
|
||||
candidates = append(candidates, ManagedDomainMatchCandidate{
|
||||
ManagedDomainID: item.ID,
|
||||
Domain: item.Domain,
|
||||
MatchType: matchType,
|
||||
CertificateID: certificate.ID,
|
||||
CertificateName: certificate.Name,
|
||||
})
|
||||
}
|
||||
sortManagedDomainCandidates(candidates)
|
||||
result := &ManagedDomainMatchResult{
|
||||
Domain: domain,
|
||||
Matched: len(candidates) > 0,
|
||||
Candidates: candidates,
|
||||
}
|
||||
if len(candidates) > 0 {
|
||||
candidate := candidates[0]
|
||||
result.Candidate = &candidate
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func buildManagedDomain(existing *model.ManagedDomain, input ManagedDomainInput) (*model.ManagedDomain, error) {
|
||||
domain := normalizeManagedDomain(input.Domain)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
if err := validateManagedDomainPattern(domain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if input.CertID != nil && *input.CertID != 0 {
|
||||
if _, err := model.GetTLSCertificateByID(*input.CertID); err != nil {
|
||||
return nil, errors.New("所选证书不存在")
|
||||
}
|
||||
} else {
|
||||
input.CertID = nil
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.ManagedDomain{}
|
||||
}
|
||||
existing.Domain = domain
|
||||
existing.CertID = input.CertID
|
||||
existing.Enabled = input.Enabled
|
||||
existing.Remark = remark
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
func normalizeManagedDomain(domain string) string {
|
||||
return strings.ToLower(strings.TrimSpace(domain))
|
||||
}
|
||||
|
||||
func validateManagedDomainPattern(domain string) error {
|
||||
if domain == "" {
|
||||
return errors.New("域名不能为空")
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
if strings.Contains(domain, "*") {
|
||||
if !strings.HasPrefix(domain, "*.") || strings.Count(domain, "*") != 1 {
|
||||
return errors.New("通配符域名仅支持 *.example.com 格式")
|
||||
}
|
||||
return validateHostname(strings.TrimPrefix(domain, "*."))
|
||||
}
|
||||
return validateHostname(domain)
|
||||
}
|
||||
|
||||
func validateHostname(domain string) error {
|
||||
if domain == "" {
|
||||
return errors.New("域名不能为空")
|
||||
}
|
||||
if len(domain) > 253 {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
labels := strings.Split(domain, ".")
|
||||
if len(labels) < 2 {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
for _, label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
if label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
for _, r := range label {
|
||||
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
|
||||
continue
|
||||
}
|
||||
return errors.New("域名格式不合法")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func detectManagedDomainMatchType(pattern string, domain string) string {
|
||||
if pattern == domain {
|
||||
return ManagedDomainMatchTypeExact
|
||||
}
|
||||
if !strings.HasPrefix(pattern, "*.") {
|
||||
return ""
|
||||
}
|
||||
suffix := strings.TrimPrefix(pattern, "*.")
|
||||
if !strings.HasSuffix(domain, "."+suffix) {
|
||||
return ""
|
||||
}
|
||||
prefix := strings.TrimSuffix(domain, "."+suffix)
|
||||
if prefix == "" || strings.Contains(prefix, ".") {
|
||||
return ""
|
||||
}
|
||||
return ManagedDomainMatchTypeWildcard
|
||||
}
|
||||
|
||||
func sortManagedDomainCandidates(candidates []ManagedDomainMatchCandidate) {
|
||||
sort.Slice(candidates, func(i int, j int) bool {
|
||||
left := candidates[i]
|
||||
right := candidates[j]
|
||||
if left.MatchType != right.MatchType {
|
||||
return left.MatchType == ManagedDomainMatchTypeExact
|
||||
}
|
||||
if len(left.Domain) != len(right.Domain) {
|
||||
return len(left.Domain) > len(right.Domain)
|
||||
}
|
||||
return left.ManagedDomainID < right.ManagedDomainID
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package service
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMatchManagedDomainCertificatePrefersExactMatch(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
wildcardCertPEM, wildcardKeyPEM := generateCertificatePair(t, []string{"*.example.com"})
|
||||
wildcardCert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "wildcard-cert",
|
||||
CertPEM: wildcardCertPEM,
|
||||
KeyPEM: wildcardKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create wildcard certificate: %v", err)
|
||||
}
|
||||
exactCertPEM, exactKeyPEM := generateCertificatePair(t, []string{"api.example.com"})
|
||||
exactCert, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "exact-cert",
|
||||
CertPEM: exactCertPEM,
|
||||
KeyPEM: exactKeyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create exact certificate: %v", err)
|
||||
}
|
||||
if _, err = CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "*.example.com",
|
||||
CertID: &wildcardCert.ID,
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to create wildcard managed domain: %v", err)
|
||||
}
|
||||
if _, err = CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "api.example.com",
|
||||
CertID: &exactCert.ID,
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to create exact managed domain: %v", err)
|
||||
}
|
||||
|
||||
result, err := MatchManagedDomainCertificate("api.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("MatchManagedDomainCertificate failed: %v", err)
|
||||
}
|
||||
if !result.Matched || result.Candidate == nil {
|
||||
t.Fatal("expected exact domain to be matched")
|
||||
}
|
||||
if result.Candidate.MatchType != ManagedDomainMatchTypeExact {
|
||||
t.Fatalf("expected exact match first, got %s", result.Candidate.MatchType)
|
||||
}
|
||||
if result.Candidate.CertificateID != exactCert.ID {
|
||||
t.Fatalf("expected exact certificate %d, got %d", exactCert.ID, result.Candidate.CertificateID)
|
||||
}
|
||||
if len(result.Candidates) != 2 {
|
||||
t.Fatalf("expected 2 match candidates, got %d", len(result.Candidates))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatchManagedDomainCertificateSupportsWildcard(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
certPEM, keyPEM := generateCertificatePair(t, []string{"*.example.com"})
|
||||
certificate, err := CreateTLSCertificate(TLSCertificateInput{
|
||||
Name: "wildcard-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create certificate: %v", err)
|
||||
}
|
||||
if _, err = CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "*.example.com",
|
||||
CertID: &certificate.ID,
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to create managed domain: %v", err)
|
||||
}
|
||||
|
||||
result, err := MatchManagedDomainCertificate("edge.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("MatchManagedDomainCertificate failed: %v", err)
|
||||
}
|
||||
if !result.Matched || result.Candidate == nil {
|
||||
t.Fatal("expected wildcard domain to be matched")
|
||||
}
|
||||
if result.Candidate.MatchType != ManagedDomainMatchTypeWildcard {
|
||||
t.Fatalf("expected wildcard match, got %s", result.Candidate.MatchType)
|
||||
}
|
||||
|
||||
deepResult, err := MatchManagedDomainCertificate("deep.edge.example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("MatchManagedDomainCertificate failed: %v", err)
|
||||
}
|
||||
if deepResult.Matched {
|
||||
t.Fatal("expected single-level wildcard not to match deep subdomain")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateManagedDomainRejectsInvalidWildcard(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
_, err := CreateManagedDomain(ManagedDomainInput{
|
||||
Domain: "*.*.example.com",
|
||||
Enabled: true,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid wildcard domain to fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,506 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"openflare/common"
|
||||
"openflare/model"
|
||||
"openflare/utils/geoip"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type NodeInput struct {
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
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"`
|
||||
}
|
||||
|
||||
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"`
|
||||
AgentToken string `json:"agent_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("节点名不能为空")
|
||||
}
|
||||
node := &model.Node{
|
||||
Name: name,
|
||||
IP: ip,
|
||||
GeoName: geoName,
|
||||
GeoLatitude: geoLatitude,
|
||||
GeoLongitude: geoLongitude,
|
||||
GeoManualOverride: geoManualOverride,
|
||||
AgentVersion: "",
|
||||
NginxVersion: "",
|
||||
Status: NodeStatusPending,
|
||||
AutoUpdateEnabled: input.AutoUpdateEnabled,
|
||||
}
|
||||
node.NodeID, err = newServerNodeID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.AgentToken, err = newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
if err := node.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("节点标识生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
refreshAgentTokenCache(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
|
||||
}
|
||||
node.Name = name
|
||||
node.IP = ip
|
||||
node.GeoName = geoName
|
||||
node.GeoLatitude = geoLatitude
|
||||
node.GeoLongitude = geoLongitude
|
||||
node.GeoManualOverride = geoManualOverride
|
||||
node.AutoUpdateEnabled = input.AutoUpdateEnabled
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, strings.TrimSpace(node.IP))
|
||||
}
|
||||
if err = node.Update(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refreshAgentTokenCache(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
|
||||
}
|
||||
invalidateAgentTokenCache(node.AgentToken)
|
||||
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
|
||||
}
|
||||
refreshAgentTokenCache(node)
|
||||
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
|
||||
}
|
||||
refreshAgentTokenCache(node)
|
||||
slog.Info("openresty restart requested", "node_id", node.NodeID, "name", node.Name)
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
func AuthenticateAgentToken(token string) (*model.Node, error) {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return nil, errors.New("缺少 Agent Token")
|
||||
}
|
||||
return authenticateAgentTokenWithCache(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.OptionMap["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,
|
||||
GeoName: strings.TrimSpace(node.GeoName),
|
||||
GeoLatitude: node.GeoLatitude,
|
||||
GeoLongitude: node.GeoLongitude,
|
||||
GeoManualOverride: node.GeoManualOverride,
|
||||
AgentToken: node.AgentToken,
|
||||
UpdateChannel: strings.TrimSpace(node.UpdateChannel),
|
||||
UpdateTag: strings.TrimSpace(node.UpdateTag),
|
||||
RestartOpenrestyRequested: node.RestartOpenrestyRequested,
|
||||
AgentVersion: node.AgentVersion,
|
||||
NginxVersion: node.NginxVersion,
|
||||
OpenrestyStatus: normalizeOpenrestyStatus(node.OpenrestyStatus),
|
||||
OpenrestyMessage: strings.TrimSpace(node.OpenrestyMessage),
|
||||
Status: status,
|
||||
CurrentVersion: node.CurrentVersion,
|
||||
LastSeenAt: node.LastSeenAt,
|
||||
LastError: node.LastError,
|
||||
CreatedAt: node.CreatedAt,
|
||||
UpdatedAt: node.UpdatedAt,
|
||||
AutoUpdateEnabled: node.AutoUpdateEnabled,
|
||||
UpdateRequested: node.UpdateRequested,
|
||||
}
|
||||
if view.UpdateChannel == "" {
|
||||
view.UpdateChannel = ReleaseChannelStable.String()
|
||||
}
|
||||
return view
|
||||
}
|
||||
|
||||
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 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 cloneCoordinate(value *float64) *float64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func buildNodeAgentReleaseView(node *model.Node, release *githubReleaseResponse, channel ReleaseChannel) *NodeAgentReleaseInfo {
|
||||
currentVersion := strings.TrimSpace(node.AgentVersion)
|
||||
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 RegisterNodeWithAgentToken(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
|
||||
}
|
||||
refreshAgentTokenCache(node)
|
||||
slog.Info("agent register succeeded on reserved node", "node_id", node.NodeID, "name", node.Name)
|
||||
return &AgentRegistrationResponse{
|
||||
NodeID: node.NodeID,
|
||||
AgentToken: node.AgentToken,
|
||||
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,
|
||||
AgentToken: agentToken,
|
||||
}
|
||||
applyNodeRuntime(node, payload, false)
|
||||
if err = node.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("节点标识生成冲突,请重试")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
refreshAgentTokenCache(node)
|
||||
slog.Info("agent discovery register succeeded", "node_id", node.NodeID, "name", node.Name)
|
||||
return &AgentRegistrationResponse{
|
||||
NodeID: node.NodeID,
|
||||
AgentToken: node.AgentToken,
|
||||
Name: node.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeAgentNodePayload(payload AgentNodePayload) AgentNodePayload {
|
||||
payload.Name = strings.TrimSpace(payload.Name)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
payload.AgentVersion = strings.TrimSpace(payload.AgentVersion)
|
||||
payload.NginxVersion = strings.TrimSpace(payload.NginxVersion)
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.LastError = strings.TrimSpace(payload.LastError)
|
||||
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
payload.OpenrestyMessage = strings.TrimSpace(payload.OpenrestyMessage)
|
||||
return payload
|
||||
}
|
||||
|
||||
func validateAgentNodePayload(payload AgentNodePayload) error {
|
||||
if payload.IP == "" {
|
||||
return errors.New("ip 不能为空")
|
||||
}
|
||||
if payload.AgentVersion == "" {
|
||||
return errors.New("agent_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)
|
||||
}
|
||||
}
|
||||
node.IP = strings.TrimSpace(payload.IP)
|
||||
node.AgentVersion = strings.TrimSpace(payload.AgentVersion)
|
||||
node.NginxVersion = strings.TrimSpace(payload.NginxVersion)
|
||||
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
node.OpenrestyMessage = strings.TrimSpace(payload.OpenrestyMessage)
|
||||
node.Status = NodeStatusOnline
|
||||
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
node.LastSeenAt = time.Now()
|
||||
node.LastError = strings.TrimSpace(payload.LastError)
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"openflare/model"
|
||||
"time"
|
||||
|
||||
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 cachedMissingAgentToken struct {
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type agentTokenAuthCache struct {
|
||||
positive *ristretto.Cache[string, cachedAgentNode]
|
||||
negative *ristretto.Cache[string, cachedMissingAgentToken]
|
||||
now func() time.Time
|
||||
loadNodeByToken func(string) (*model.Node, error)
|
||||
}
|
||||
|
||||
var nodeAgentTokenCache = newAgentTokenAuthCache()
|
||||
|
||||
func newAgentTokenAuthCache() *agentTokenAuthCache {
|
||||
return &agentTokenAuthCache{
|
||||
positive: mustNewAgentTokenPositiveCache(),
|
||||
negative: mustNewAgentTokenNegativeCache(),
|
||||
now: time.Now,
|
||||
loadNodeByToken: func(token string) (*model.Node, error) {
|
||||
return model.GetNodeByAgentToken(token)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func mustNewAgentTokenPositiveCache() *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 mustNewAgentTokenNegativeCache() *ristretto.Cache[string, cachedMissingAgentToken] {
|
||||
cache, err := ristretto.NewCache(&ristretto.Config[string, cachedMissingAgentToken]{
|
||||
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, cachedMissingAgentToken{
|
||||
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 authenticateAgentTokenWithCache(token string) (*model.Node, error) {
|
||||
return nodeAgentTokenCache.authenticate(token)
|
||||
}
|
||||
|
||||
func refreshAgentTokenCache(node *model.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
nodeAgentTokenCache.storeNode(
|
||||
node.AgentToken,
|
||||
node,
|
||||
nodeAgentTokenCache.now().Add(agentTokenPositiveCacheTTL),
|
||||
)
|
||||
}
|
||||
|
||||
func invalidateAgentTokenCache(token string) {
|
||||
nodeAgentTokenCache.invalidate(token)
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"openflare/model"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestAgentTokenAuthCacheUsesPositiveCacheUntilLogicalExpiry(t *testing.T) {
|
||||
cache := newAgentTokenAuthCache()
|
||||
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",
|
||||
AgentToken: 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 TestAgentTokenAuthCacheRefreshesAfterMissingEntryExpires(t *testing.T) {
|
||||
cache := newAgentTokenAuthCache()
|
||||
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",
|
||||
AgentToken: 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,141 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"openflare/model"
|
||||
"time"
|
||||
|
||||
"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"`
|
||||
}
|
||||
|
||||
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"`
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
return &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),
|
||||
DiskIO24h: buildDiskIOTrendPoints(now, trendSnapshots),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func latestMetricSnapshot(snapshots []*model.NodeMetricSnapshot) *model.NodeMetricSnapshot {
|
||||
for _, snapshot := range snapshots {
|
||||
if snapshot != nil {
|
||||
return snapshot
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func latestTrafficReport(reports []*model.NodeRequestReport) *model.NodeRequestReport {
|
||||
for _, report := range reports {
|
||||
if report != nil {
|
||||
return report
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeObservabilityLimit(limit int) int {
|
||||
if limit <= 0 {
|
||||
return defaultObservabilityLimit
|
||||
}
|
||||
if limit > maxObservabilityLimit {
|
||||
return maxObservabilityLimit
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func normalizeObservabilityWindow(hours int) time.Duration {
|
||||
if hours <= 0 {
|
||||
return defaultObservabilityWindow
|
||||
}
|
||||
return time.Duration(hours) * time.Hour
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,349 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"openflare/model"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
NodeHealthEventStatusActive = "active"
|
||||
NodeHealthEventStatusResolved = "resolved"
|
||||
NodeHealthSeverityInfo = "info"
|
||||
NodeHealthSeverityWarning = "warning"
|
||||
NodeHealthSeverityCritical = "critical"
|
||||
nodeAccessLogRetentionWindow = 24 * time.Hour
|
||||
)
|
||||
|
||||
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"`
|
||||
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"`
|
||||
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 := 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 := 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),
|
||||
RawJSON: marshalJSON(profile),
|
||||
}
|
||||
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,
|
||||
OpenrestyRxBytes: snapshot.OpenrestyRxBytes,
|
||||
OpenrestyTxBytes: snapshot.OpenrestyTxBytes,
|
||||
OpenrestyConnections: snapshot.OpenrestyConnections,
|
||||
RawJSON: marshalJSON(snapshot),
|
||||
}
|
||||
return tx.Where("node_id = ? AND captured_at = ?", nodeID, record.CapturedAt).Assign(record).FirstOrCreate(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),
|
||||
RawJSON: marshalJSON(report),
|
||||
}
|
||||
return tx.Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, record.WindowStartedAt, record.WindowEndedAt).Assign(record).FirstOrCreate(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: strings.TrimSpace(item.Path),
|
||||
StatusCode: item.StatusCode,
|
||||
RawJSON: marshalJSON(item),
|
||||
}
|
||||
if resolver != nil {
|
||||
record.Region = resolver.Resolve(record.RemoteAddr)
|
||||
}
|
||||
if err := tx.Where(
|
||||
"node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?",
|
||||
nodeID,
|
||||
record.LoggedAt,
|
||||
record.RemoteAddr,
|
||||
record.Host,
|
||||
record.Path,
|
||||
record.StatusCode,
|
||||
).Assign(record).FirstOrCreate(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Where("node_id = ? AND logged_at < ?", nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow)).Delete(&model.NodeAccessLog{}).Error
|
||||
}
|
||||
|
||||
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentNodeHealthEvent, reportedAt time.Time) error {
|
||||
activeTypes := make(map[string]AgentNodeHealthEvent, len(events))
|
||||
for _, event := range events {
|
||||
eventType := normalizeHealthEventType(event.EventType)
|
||||
if eventType == "" {
|
||||
continue
|
||||
}
|
||||
event.EventType = eventType
|
||||
event.Severity = normalizeHealthSeverity(event.Severity)
|
||||
if event.TriggeredAtUnix <= 0 {
|
||||
event.TriggeredAtUnix = reportedAt.Unix()
|
||||
}
|
||||
activeTypes[eventType] = event
|
||||
}
|
||||
|
||||
var activeEvents []*model.NodeHealthEvent
|
||||
if err := tx.Where("node_id = ? AND status = ?", nodeID, NodeHealthEventStatusActive).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 = strings.TrimSpace(event.Message)
|
||||
existing.LastTriggeredAt = triggeredAt
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.RawJSON = marshalJSON(event)
|
||||
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: strings.TrimSpace(event.Message),
|
||||
FirstTriggeredAt: triggeredAt,
|
||||
LastTriggeredAt: triggeredAt,
|
||||
ReportedAt: reportedAt,
|
||||
RawJSON: marshalJSON(event),
|
||||
}
|
||||
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 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,172 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
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,225 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"time"
|
||||
)
|
||||
|
||||
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) []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
|
||||
points[index].OpenrestyRxBytes += snapshot.OpenrestyRxBytes
|
||||
points[index].OpenrestyTxBytes += snapshot.OpenrestyTxBytes
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.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,32 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"openflare/model"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestBuildDiskIOTrendPointsUsesCounterDelta(t *testing.T) {
|
||||
now := time.Date(2026, 3, 14, 18, 30, 0, 0, time.UTC)
|
||||
start := trendWindowStart(now)
|
||||
|
||||
points := buildDiskIOTrendPoints(now, []*model.NodeMetricSnapshot{
|
||||
{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: start.Add(22 * time.Hour),
|
||||
DiskReadBytes: 100,
|
||||
DiskWriteBytes: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: start.Add(23 * time.Hour),
|
||||
DiskReadBytes: 250,
|
||||
DiskWriteBytes: 260,
|
||||
},
|
||||
})
|
||||
|
||||
last := points[len(points)-1]
|
||||
if last.DiskReadBytes != 150 || last.DiskWriteBytes != 60 {
|
||||
t.Fatalf("expected disk io trend to use counter delta, got %+v", last)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package service
|
||||
|
||||
import "fmt"
|
||||
|
||||
const (
|
||||
openRestyObservabilityInitLuaPath = "init.lua"
|
||||
openRestyObservabilityLogLuaPath = "log.lua"
|
||||
openRestyObservabilityReadLuaPath = "read.lua"
|
||||
)
|
||||
|
||||
func renderOpenRestyObservabilityTemplateBlock() string {
|
||||
return stringsJoinLines(
|
||||
" lua_shared_dict openflare_observability 10m;",
|
||||
fmt.Sprintf(" init_worker_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityInitLuaPath),
|
||||
fmt.Sprintf(" log_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityLogLuaPath),
|
||||
"",
|
||||
fmt.Sprintf(" server {"),
|
||||
fmt.Sprintf(" listen %s;", nginxObservabilityListenPlaceholder),
|
||||
" server_name openflare-observability;",
|
||||
" access_log off;",
|
||||
"",
|
||||
" location = /openflare/observability {",
|
||||
" default_type application/json;",
|
||||
fmt.Sprintf(" content_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityReadLuaPath),
|
||||
" }",
|
||||
"",
|
||||
" location = /openflare/stub_status {",
|
||||
" stub_status;",
|
||||
" }",
|
||||
" }",
|
||||
"",
|
||||
)
|
||||
}
|
||||
|
||||
func stringsJoinLines(lines ...string) string {
|
||||
if len(lines) == 0 {
|
||||
return ""
|
||||
}
|
||||
result := ""
|
||||
for index, line := range lines {
|
||||
if index > 0 {
|
||||
result += "\n"
|
||||
}
|
||||
result += line
|
||||
}
|
||||
return result + "\n"
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/url"
|
||||
"openflare/model"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var proxyHeaderKeyPattern = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
|
||||
|
||||
type ProxyRouteCustomHeaderInput struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
|
||||
type ProxyRouteInput struct {
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
CustomHeaders []ProxyRouteCustomHeaderInput `json:"custom_headers"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
func ListProxyRoutes() ([]*model.ProxyRoute, error) {
|
||||
return model.ListProxyRoutes()
|
||||
}
|
||||
|
||||
func CreateProxyRoute(input ProxyRouteInput) (*model.ProxyRoute, error) {
|
||||
route, err := buildProxyRoute(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = route.Insert(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func UpdateProxyRoute(id uint, input ProxyRouteInput) (*model.ProxyRoute, error) {
|
||||
route, err := model.GetProxyRouteByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
route, err = buildProxyRoute(route, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = route.Update(); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New("域名已存在")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func DeleteProxyRoute(id uint) error {
|
||||
route, err := model.GetProxyRouteByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return route.Delete()
|
||||
}
|
||||
|
||||
func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.ProxyRoute, error) {
|
||||
domain := strings.ToLower(strings.TrimSpace(input.Domain))
|
||||
originURL := strings.TrimSpace(input.OriginURL)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeadersJSON, err := json.Marshal(customHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if domain == "" {
|
||||
return nil, errors.New("域名不能为空")
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, errors.New("域名格式不合法")
|
||||
}
|
||||
if err := validateOriginURL(originURL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !input.EnableHTTPS {
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
}
|
||||
if input.EnableHTTPS {
|
||||
if input.CertID == nil || *input.CertID == 0 {
|
||||
return nil, errors.New("启用 HTTPS 时必须选择证书")
|
||||
}
|
||||
if _, err := model.GetTLSCertificateByID(*input.CertID); err != nil {
|
||||
return nil, errors.New("所选证书不存在")
|
||||
}
|
||||
}
|
||||
if input.RedirectHTTP && !input.EnableHTTPS {
|
||||
return nil, errors.New("仅启用 HTTPS 后才能开启 HTTP 重定向")
|
||||
}
|
||||
if route == nil {
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.Domain = domain
|
||||
route.OriginURL = originURL
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.CustomHeaders = string(customHeadersJSON)
|
||||
route.Remark = remark
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func normalizeCustomHeaders(headers []ProxyRouteCustomHeaderInput) ([]ProxyRouteCustomHeaderInput, error) {
|
||||
if len(headers) == 0 {
|
||||
return []ProxyRouteCustomHeaderInput{}, nil
|
||||
}
|
||||
normalized := make([]ProxyRouteCustomHeaderInput, 0, len(headers))
|
||||
for _, header := range headers {
|
||||
key := strings.TrimSpace(header.Key)
|
||||
value := strings.TrimSpace(header.Value)
|
||||
if key == "" && value == "" {
|
||||
continue
|
||||
}
|
||||
if key == "" {
|
||||
return nil, errors.New("自定义请求头名称不能为空")
|
||||
}
|
||||
if !proxyHeaderKeyPattern.MatchString(key) {
|
||||
return nil, errors.New("自定义请求头名称格式不合法")
|
||||
}
|
||||
if strings.ContainsAny(key, "\r\n") || strings.ContainsAny(value, "\r\n") {
|
||||
return nil, errors.New("自定义请求头不能包含换行")
|
||||
}
|
||||
normalized = append(normalized, ProxyRouteCustomHeaderInput{
|
||||
Key: key,
|
||||
Value: value,
|
||||
})
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeStoredCustomHeaders(raw string) ([]ProxyRouteCustomHeaderInput, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []ProxyRouteCustomHeaderInput{}, nil
|
||||
}
|
||||
var headers []ProxyRouteCustomHeaderInput
|
||||
if err := json.Unmarshal([]byte(text), &headers); err != nil {
|
||||
return nil, errors.New("自定义请求头配置格式不合法")
|
||||
}
|
||||
return normalizeCustomHeaders(headers)
|
||||
}
|
||||
|
||||
func validateOriginURL(raw string) error {
|
||||
if raw == "" {
|
||||
return errors.New("源站地址不能为空")
|
||||
}
|
||||
parsed, err := url.ParseRequestURI(raw)
|
||||
if err != nil {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
||||
return errors.New("源站地址必须以 http:// 或 https:// 开头")
|
||||
}
|
||||
if parsed.Host == "" {
|
||||
return errors.New("源站地址格式不合法")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
return err != nil && strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"openflare/model"
|
||||
"strings"
|
||||
)
|
||||
|
||||
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"`
|
||||
}
|
||||
|
||||
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,
|
||||
}, 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 isUniqueConstraintError(err) {
|
||||
return nil, errors.New("证书名称已存在")
|
||||
}
|
||||
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("证书文件和私钥文件不能为空")
|
||||
}
|
||||
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 isUniqueConstraintError(err) {
|
||||
return nil, errors.New("certificate name already exists")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return certificate, nil
|
||||
}
|
||||
|
||||
func DeleteTLSCertificate(id uint) error {
|
||||
var routeCount int64
|
||||
if err := model.DB.Model(&model.ProxyRoute{}).Where("cert_id = ?", id).Count(&routeCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if routeCount > 0 {
|
||||
return errors.New("证书仍被反代规则引用,无法删除")
|
||||
}
|
||||
certificate, err := model.GetTLSCertificateByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return certificate.Delete()
|
||||
}
|
||||
|
||||
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("证书名称不能为空")
|
||||
}
|
||||
if certPEM == "" || keyPEM == "" {
|
||||
return nil, errors.New("证书内容和私钥内容不能为空")
|
||||
}
|
||||
parsed, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("证书或私钥格式不合法: %w", err)
|
||||
}
|
||||
if len(parsed.Certificate) == 0 {
|
||||
return nil, errors.New("证书内容不合法")
|
||||
}
|
||||
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
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,73 @@
|
||||
//go:build !windows
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
var unixRename = os.Rename
|
||||
|
||||
func replaceAndRestartServer(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
_ = os.Remove(backupPath)
|
||||
if err := unixRename(execPath, backupPath); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("备份当前服务端二进制失败: %w", err)
|
||||
}
|
||||
if err := replaceFileUnix(tmpPath, execPath); err != nil {
|
||||
_ = unixRename(backupPath, execPath)
|
||||
return fmt.Errorf("替换服务端二进制失败: %w", err)
|
||||
}
|
||||
_ = os.Remove(backupPath)
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil {
|
||||
return fmt.Errorf("重启服务失败: %w", err)
|
||||
}
|
||||
return fmt.Errorf("unreachable after exec")
|
||||
}
|
||||
|
||||
func replaceFileUnix(srcPath string, dstPath string) error {
|
||||
if err := unixRename(srcPath, dstPath); err == nil {
|
||||
return nil
|
||||
} else if linkErr, ok := err.(*os.LinkError); !ok || linkErr.Err != syscall.EXDEV {
|
||||
return err
|
||||
}
|
||||
|
||||
sourceFile, err := os.Open(srcPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer sourceFile.Close()
|
||||
|
||||
info, err := sourceFile.Stat()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
destinationFile, err := os.OpenFile(dstPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
copyErr := func() error {
|
||||
defer destinationFile.Close()
|
||||
if _, err = io.Copy(destinationFile, sourceFile); err != nil {
|
||||
return err
|
||||
}
|
||||
if err = destinationFile.Sync(); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}()
|
||||
if copyErr != nil {
|
||||
return copyErr
|
||||
}
|
||||
|
||||
if err = os.Chmod(dstPath, info.Mode().Perm()); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Remove(srcPath)
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
//go:build !windows
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestReplaceFileUnixFallsBackOnCrossDeviceRename(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
srcPath := filepath.Join(tempDir, "source.bin")
|
||||
dstPath := filepath.Join(tempDir, "target.bin")
|
||||
|
||||
if err := os.WriteFile(srcPath, []byte("new-binary"), 0o755); err != nil {
|
||||
t.Fatalf("failed to write source file: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(dstPath, []byte("old-binary"), 0o755); err != nil {
|
||||
t.Fatalf("failed to write target file: %v", err)
|
||||
}
|
||||
|
||||
originalRename := unixRename
|
||||
unixRename = func(oldPath string, newPath string) error {
|
||||
if oldPath == srcPath && newPath == dstPath {
|
||||
return &os.LinkError{Op: "rename", Old: oldPath, New: newPath, Err: syscall.EXDEV}
|
||||
}
|
||||
return os.Rename(oldPath, newPath)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
unixRename = originalRename
|
||||
})
|
||||
|
||||
if err := replaceFileUnix(srcPath, dstPath); err != nil {
|
||||
t.Fatalf("expected cross-device fallback to succeed: %v", err)
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(dstPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read target file: %v", err)
|
||||
}
|
||||
if string(content) != "new-binary" {
|
||||
t.Fatalf("unexpected target content: %s", string(content))
|
||||
}
|
||||
if _, err = os.Stat(srcPath); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("expected source file to be removed, got err=%v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
//go:build windows
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func replaceAndRestartServer(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
scriptPath := execPath + ".update.cmd"
|
||||
script := fmt.Sprintf(`@echo off
|
||||
setlocal
|
||||
:waitloop
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 (
|
||||
ping 127.0.0.1 -n 2 >nul
|
||||
goto waitloop
|
||||
)
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 exit /b 1
|
||||
start "" %s
|
||||
del /Q "%s" >nul 2>nul
|
||||
del /Q "%%~f0" >nul 2>nul
|
||||
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
|
||||
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("写入升级重启脚本失败: %w", err)
|
||||
}
|
||||
|
||||
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
|
||||
if err := cmd.Start(); err != nil {
|
||||
_ = os.Remove(scriptPath)
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("调度升级重启失败: %w", err)
|
||||
}
|
||||
|
||||
os.Exit(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildWindowsCommandLine(execPath string, args []string) string {
|
||||
parts := []string{quoteWindowsArg(execPath)}
|
||||
for _, arg := range args {
|
||||
parts = append(parts, quoteWindowsArg(arg))
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func quoteWindowsArg(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -0,0 +1,414 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"openflare/common"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user