[优化] 改名

This commit is contained in:
ryan
2026-03-15 16:04:54 +08:00
parent d68773c554
commit 32d90ba641
304 changed files with 629 additions and 969 deletions
+111
View File
@@ -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 ""
}
+100
View File
@@ -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))
}
}
+376
View File
@@ -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
}
+814
View File
@@ -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
}
+248
View File
@@ -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
}
+50
View File
@@ -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)
}
+227
View File
@@ -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")
}
}
+506
View File
@@ -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
+349
View File
@@ -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"
}
+183
View File
@@ -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")
}
+150
View File
@@ -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, `"`, `""`) + `"`
}
+414
View File
@@ -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")
}
}