This commit is contained in:
ryan
2026-06-18 16:51:06 +08:00
parent 61cfacba55
commit e3bfd9ca6d
40 changed files with 2819 additions and 116 deletions
@@ -0,0 +1,95 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"log/slog"
"net"
"strings"
pkggeoip "github.com/rain-kl/openflare/pkg/geoip"
)
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
return pkggeoip.NewMaxMindGeoIPService()
}
type accessLogRegionResolver struct {
provider pkggeoip.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 ""
}
@@ -11,6 +11,7 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/db"
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
"gorm.io/gorm"
)
@@ -89,8 +90,12 @@ func sourceSupportFiles(files []SupportFile) []SupportFile {
}
func isRuntimeGeneratedSupportFile(path string) bool {
path = strings.TrimSpace(path)
return strings.HasPrefix(path, "runtime/")
switch strings.TrimSpace(path) {
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName:
return true
default:
return false
}
}
func isActiveConfigNotFound(err error) bool {
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"testing"
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
)
func TestIsRuntimeGeneratedSupportFile(t *testing.T) {
tests := []struct {
path string
want bool
}{
{path: "pow_config.json", want: true},
{path: "waf_config.json", want: true},
{path: openrestyrender.SourceConfigFileName, want: true},
{path: "runtime/custom.json", want: false},
{path: "certs/example.pem", want: false},
}
for _, tc := range tests {
if got := isRuntimeGeneratedSupportFile(tc.path); got != tc.want {
t.Fatalf("isRuntimeGeneratedSupportFile(%q) = %v, want %v", tc.path, got, tc.want)
}
}
}
func TestSourceSupportFilesFiltersRuntimeGeneratedFiles(t *testing.T) {
files := []SupportFile{
{Path: "certs/example.pem", Content: "pem"},
{Path: "pow_config.json", Content: "{}"},
{Path: "waf_config.json", Content: "{}"},
{Path: openrestyrender.SourceConfigFileName, Content: "{}"},
{Path: "routes/extra.json", Content: "{}"},
}
filtered := sourceSupportFiles(files)
if len(filtered) != 2 {
t.Fatalf("expected 2 support files, got %d: %+v", len(filtered), filtered)
}
if filtered[0].Path != "certs/example.pem" || filtered[1].Path != "routes/extra.json" {
t.Fatalf("unexpected filtered files: %+v", filtered)
}
}
@@ -17,5 +17,4 @@ const (
errIPInvalid = "ip 格式无效"
errAgentVersionRequired = "version 不能为空"
errNodeIDConflict = "节点标识生成冲突,请重试"
errPagesPackageNotFound = "Pages 部署包尚未实现"
)
@@ -11,6 +11,7 @@ import (
"strings"
"time"
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
"github.com/Rain-kl/Wavelet/internal/model"
)
@@ -97,6 +98,41 @@ func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveNa
now := time.Now()
node.LastSeenAt = &now
node.LastError = truncateForDatabase(payload.LastError, 16000)
if !node.GeoManualOverride {
applyGeoInfoFromIP(node, node.IP)
}
}
func applyGeoInfoFromIP(node *model.OpenFlareNode, 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 := ofgeoip.GeoInfoFromIP(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 cloneCoordinate(value *float64) *float64 {
if value == nil {
return nil
}
cloned := *value
return &cloned
}
func truncateForDatabase(value string, max int) string {
@@ -199,6 +235,7 @@ func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenF
}
appendIfChanged("name", previous.Name, current.Name)
appendIfChanged("ip", previous.IP, current.IP)
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
appendIfChanged("version", previous.Version, current.Version)
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
@@ -210,12 +247,25 @@ func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenF
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 !lastSeenAtEqual(previous.LastSeenAt, 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
}
func lastSeenAtEqual(before *time.Time, after *time.Time) bool {
if before == nil || after == nil {
return before == after
@@ -241,7 +291,8 @@ func isUniqueConstraintError(err error) bool {
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
func refreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
if node == nil {
return
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
pkggeoip "github.com/rain-kl/openflare/pkg/geoip"
)
type fakeGeoIPProvider struct {
info *pkggeoip.GeoInfo
}
func (f *fakeGeoIPProvider) Name() string { return "fake-geoip" }
func (f *fakeGeoIPProvider) GetGeoInfo(ip net.IP) (*pkggeoip.GeoInfo, error) {
return f.info, nil
}
func (f *fakeGeoIPProvider) UpdateDatabase() error { return nil }
func (f *fakeGeoIPProvider) Close() error { return nil }
func withFakeGeoIPProvider(t *testing.T, info *pkggeoip.GeoInfo) {
t.Helper()
previous := pkggeoip.CurrentProvider
pkggeoip.CurrentProvider = &fakeGeoIPProvider{info: info}
t.Cleanup(func() {
pkggeoip.CurrentProvider = previous
})
}
func geoipFloat(value float64) *float64 {
return &value
}
func TestApplyGeoInfoFromIP(t *testing.T) {
latitude := 31.2304
longitude := 121.4737
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
Name: "Shanghai",
Latitude: geoipFloat(latitude),
Longitude: geoipFloat(longitude),
})
node := &model.OpenFlareNode{IP: "203.0.113.10"}
applyGeoInfoFromIP(node, node.IP)
if node.GeoName != "Shanghai" {
t.Fatalf("expected geo_name Shanghai, got %q", node.GeoName)
}
if node.GeoLatitude == nil || *node.GeoLatitude != latitude {
t.Fatalf("unexpected geo_latitude: %+v", node.GeoLatitude)
}
if node.GeoLongitude == nil || *node.GeoLongitude != longitude {
t.Fatalf("unexpected geo_longitude: %+v", node.GeoLongitude)
}
}
func TestApplyGeoInfoFromIPSkipsInvalidIP(t *testing.T) {
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{Name: "Should Not Apply"})
node := &model.OpenFlareNode{
IP: "203.0.113.10",
GeoName: "Existing",
GeoLatitude: geoipFloat(1),
GeoLongitude: geoipFloat(2),
}
applyGeoInfoFromIP(node, "not-an-ip")
if node.GeoName != "" || node.GeoLatitude != nil || node.GeoLongitude != nil {
t.Fatalf("expected geo fields to be cleared on invalid IP, got %+v", node)
}
}
func TestApplyNodeRuntimeRespectsGeoManualOverride(t *testing.T) {
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
Name: "Shanghai",
Latitude: geoipFloat(31.2304),
Longitude: geoipFloat(121.4737),
})
node := &model.OpenFlareNode{
GeoManualOverride: true,
GeoName: "Manual",
GeoLatitude: geoipFloat(10),
GeoLongitude: geoipFloat(20),
}
applyNodeRuntime(node, NodePayload{
IP: "203.0.113.10",
Version: "1.0.0",
}, true)
if node.GeoName != "Manual" {
t.Fatalf("expected manual geo_name to be preserved, got %q", node.GeoName)
}
if node.GeoLatitude == nil || *node.GeoLatitude != 10 {
t.Fatalf("expected manual geo_latitude to be preserved, got %+v", node.GeoLatitude)
}
}
func TestCollectHeartbeatChangesTracksGeoFields(t *testing.T) {
before := &model.OpenFlareNode{
IP: "10.0.0.1",
GeoName: "Old Region",
}
after := &model.OpenFlareNode{
IP: "203.0.113.10",
GeoName: "New Region",
GeoLatitude: geoipFloat(31.2304),
GeoLongitude: geoipFloat(121.4737),
}
changes := collectHeartbeatChanges(before, after)
if changes["ip"] != after.IP {
t.Fatalf("expected ip change, got %+v", changes)
}
if changes["geo_name"] != after.GeoName {
t.Fatalf("expected geo_name change, got %+v", changes)
}
if changes["geo_latitude"] != after.GeoLatitude {
t.Fatalf("expected geo_latitude change, got %+v", changes)
}
if changes["geo_longitude"] != after.GeoLongitude {
t.Fatalf("expected geo_longitude change, got %+v", changes)
}
}
@@ -28,7 +28,7 @@ func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode,
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil {
return nil, err
}
refreshAccessTokenCache(ctx, authNode)
RefreshAccessTokenCache(ctx, authNode)
return &RegistrationResponse{
NodeID: authNode.NodeID,
AccessToken: authNode.AccessToken,
@@ -74,7 +74,7 @@ func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*Registrat
}
return nil, err
}
refreshAccessTokenCache(ctx, record)
RefreshAccessTokenCache(ctx, record)
return &RegistrationResponse{
NodeID: record.NodeID,
AccessToken: record.AccessToken,
@@ -116,24 +116,29 @@ func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload N
}
}
refreshAccessTokenCache(ctx, authNode)
RefreshAccessTokenCache(ctx, authNode)
reportedAt := time.Now()
if authNode.LastSeenAt != nil {
reportedAt = *authNode.LastSeenAt
}
persistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
PersistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
activeConfig, err := getActiveConfigMeta(ctx)
if err != nil && !isActiveConfigNotFound(err) {
return nil, err
}
wafIPGroups, err := ChangedWAFIPGroupsForAgent(ctx, nil, payload.WAFIPGroupChecksums)
if err != nil {
return nil, err
}
return &HeartbeatResponse{
Node: authNode,
AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow),
ActiveConfig: activeConfig,
WAFIPGroups: nil,
WAFIPGroups: wafIPGroups,
}, nil
}
@@ -149,9 +154,13 @@ func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
return config, nil
}
// SyncWAFIPGroups is a stub until full WAF agent sync is migrated.
func SyncWAFIPGroups(_ context.Context, _ WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
return &WAFIPGroupSyncResult{Groups: []WAFIPGroup{}}, nil
// SyncWAFIPGroups returns WAF IP groups whose checksums differ from the agent state.
func SyncWAFIPGroups(ctx context.Context, input WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
groups, err := ChangedWAFIPGroupsForAgent(ctx, input.IDs, input.Checksums)
if err != nil {
return nil, err
}
return &WAFIPGroupSyncResult{Groups: groups}, nil
}
// ReportApplyLog records an agent apply result.
@@ -7,6 +7,7 @@ import (
"context"
"encoding/json"
"errors"
"log/slog"
"strings"
"time"
@@ -18,12 +19,14 @@ import (
)
const (
healthEventStatusActive = "active"
healthEventStatusResolved = "resolved"
healthSeverityInfo = "info"
healthSeverityWarning = "warning"
healthSeverityCritical = "critical"
accessLogPathMaxLength = 100
healthEventStatusActive = "active"
healthEventStatusResolved = "resolved"
healthSeverityInfo = "info"
healthSeverityWarning = "warning"
healthSeverityCritical = "critical"
nodeAccessLogRetentionDays = 90
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
accessLogPathMaxLength = 100
)
// NodeSystemProfile is the agent-reported system profile.
@@ -102,7 +105,8 @@ type NodeHealthEvent struct {
Metadata map[string]string `json:"metadata"`
}
func persistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
if strings.TrimSpace(nodeID) == "" {
return
}
@@ -275,15 +279,29 @@ func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficRep
}
func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, 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.OpenFlareAccessLog{
NodeID: nodeID,
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
Region: "",
Host: strings.TrimSpace(item.Host),
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
StatusCode: item.StatusCode,
}
if resolver != nil {
record.Region = resolver.Resolve(record.RemoteAddr)
}
exists, err := accessLogExists(tx, record)
if err != nil {
return err
@@ -295,16 +313,27 @@ func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, rep
return err
}
}
return nil
_, err = deleteAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
return err
}
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error {
return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
}
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
activeTypes := make(map[string]NodeHealthEvent, len(events))
for _, event := range events {
eventType := normalizeHealthEventType(event.EventType)
if eventType == "" {
continue
}
if len(managedEventTypes) > 0 {
if _, ok := managedEventTypes[eventType]; !ok {
continue
}
}
event.EventType = eventType
event.Severity = normalizeHealthSeverity(event.Severity)
if event.TriggeredAtUnix <= 0 {
@@ -314,7 +343,21 @@ func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEv
}
var activeEvents []*model.OpenFlareHealthEvent
if err := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive).Find(&activeEvents).Error; err != nil {
query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive)
if len(managedEventTypes) > 0 {
scopedTypes := make([]string, 0, len(managedEventTypes))
for eventType := range managedEventTypes {
eventType = normalizeHealthEventType(eventType)
if eventType != "" {
scopedTypes = append(scopedTypes, eventType)
}
}
if len(scopedTypes) == 0 {
return nil
}
query = query.Where("event_type IN ?", scopedTypes)
}
if err := query.Find(&activeEvents).Error; err != nil {
return err
}
@@ -391,6 +434,11 @@ func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEnde
return count > 0, nil
}
func deleteAccessLogsByNodeBefore(tx *gorm.DB, nodeID string, before time.Time) (int64, error) {
result := tx.Where("node_id = ? AND logged_at < ?", nodeID, before).Delete(&model.OpenFlareAccessLog{})
return result.RowsAffected, result.Error
}
func accessLogExists(tx *gorm.DB, record *model.OpenFlareAccessLog) (bool, error) {
var count int64
if err := tx.Model(&model.OpenFlareAccessLog{}).
@@ -438,6 +486,11 @@ func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
return time.Unix(unixSeconds, 0).UTC()
}
// MarshalJSON serializes a value for database JSON columns.
func MarshalJSON(value any) string {
return marshalJSON(value)
}
func marshalJSON(value any) string {
if value == nil {
return ""
@@ -5,8 +5,10 @@ package agent
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/compat"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/gin-gonic/gin"
)
@@ -98,7 +100,7 @@ func GetActiveConfigHandler(c *gin.Context) {
compat.OK(c, config)
}
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent (stub).
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
func SyncWAFIPGroupsHandler(c *gin.Context) {
var input WAFIPGroupSyncInput
if !compat.BindJSON(c, &input) {
@@ -129,13 +131,33 @@ func ReportApplyLogHandler(c *gin.Context) {
compat.OK(c, log)
}
// DownloadPagesPackageHandler is a stub until Pages agent packaging is migrated.
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
func DownloadPagesPackageHandler(c *gin.Context) {
c.JSON(http.StatusNotFound, compat.Envelope{
Success: false,
Message: errPagesPackageNotFound,
Data: nil,
})
deploymentID, ok := pagesDeploymentIDParam(c)
if !ok {
return
}
filePath, fileName, err := pages.GetDeploymentPackagePath(c.Request.Context(), deploymentID)
if err != nil {
compat.Fail(c, err.Error())
return
}
c.Header("Content-Disposition", "attachment; filename="+fileName)
c.File(filePath)
}
func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("deployment_id")
if raw == "" {
compat.Fail(c, "无效的 ID")
return 0, false
}
id64, err := strconv.ParseUint(raw, 10, 64)
if err != nil || id64 == 0 {
compat.Fail(c, "无效的 ID")
return 0, false
}
return uint(id64), true
}
// AgentWebSocketHandler upgrades an authenticated agent websocket connection.
@@ -0,0 +1,200 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"sort"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
)
type snapshotWAFRuleGroupRef struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
}
type snapshotWAFSection struct {
RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
}
type activeConfigSnapshot struct {
WAF snapshotWAFSection `json:"waf"`
}
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
if len(targetIDs) == 0 {
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
if err != nil {
return nil, err
}
changed := make([]WAFIPGroup, 0, len(groups))
for _, group := range groups {
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
continue
}
changed = append(changed, group)
}
return changed, nil
}
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
return []WAFIPGroup{}, nil
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
if err != nil {
return nil, err
}
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
for _, group := range groups {
groupByID[group.ID] = group
}
result := make([]WAFIPGroup, 0, len(ids))
for _, id := range ids {
group := groupByID[id]
if group == nil {
continue
}
agentGroup, err := buildAgentWAFIPGroup(group)
if err != nil {
return nil, err
}
result = append(result, agentGroup)
}
return result, nil
}
func buildAgentWAFIPGroup(group *model.OpenFlareWAFIPGroup) (WAFIPGroup, error) {
if group == nil {
return WAFIPGroup{}, errors.New("IP 组不存在")
}
ips, err := decodeWAFIPGroupStringList(group.IPList)
if err != nil {
return WAFIPGroup{}, err
}
if !group.Enabled {
ips = []string{}
}
agentGroup := WAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ips,
}
agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup)
return agentGroup, nil
}
func checksumAgentWAFIPGroup(group WAFIPGroup) string {
payload := struct {
ID uint `json:"id"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
}{
ID: group.ID,
Enabled: group.Enabled,
IPList: append([]string{}, group.IPList...),
}
sort.Strings(payload.IPList)
data, _ := json.Marshal(payload)
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
if isActiveConfigNotFound(err) {
return []uint{}, nil
}
return nil, err
}
snapshot, err := parseActiveConfigSnapshot(version.SnapshotJSON)
if err != nil {
return nil, err
}
idSet := make(map[uint]struct{})
for _, group := range snapshot.WAF.RuleGroups {
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids, nil
}
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
return &activeConfigSnapshot{}, nil
}
var snapshot activeConfigSnapshot
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
return nil, err
}
if snapshot.WAF.RuleGroups == nil {
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
}
return &snapshot, nil
}
func decodeWAFIPGroupStringList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var items []string
if err := json.Unmarshal([]byte(text), &items); err != nil {
return nil, err
}
return items, nil
}
func uniqueUintIDs(ids []uint) []uint {
normalized := make([]uint, 0, len(ids))
seen := make(map[uint]struct{}, len(ids))
for _, id := range ids {
if id == 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
return normalized
}
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strconv"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupWAFIPGroupTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareWAFIPGroup{},
&configVersionRecord{},
))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
t.Helper()
snapshot := map[string]any{
"routes": []any{},
"waf": map[string]any{
"rule_groups": []map[string]any{
{
"id": 1,
"name": "agent refs",
"enabled": true,
"ip_blacklist_group_ids": []uint{ipGroupID},
},
},
"bindings": []any{},
},
}
snapshotJSON, err := json.Marshal(snapshot)
require.NoError(t, err)
require.NoError(t, db.DB(ctx).Create(&configVersionRecord{
Version: "20260618-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "test-checksum",
IsActive: true,
}).Error)
}
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "agent runtime group",
Type: "manual",
Enabled: true,
IPList: `["203.0.113.44"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.Equal(t, ipGroup.ID, groups[0].ID)
assert.Equal(t, "203.0.113.44", groups[0].IPList[0])
assert.NotEmpty(t, groups[0].Checksum)
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
same, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err)
assert.Empty(t, same)
ipGroup.IPList = `["203.0.113.45"]`
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err)
require.Len(t, delta, 1)
assert.Equal(t, ipGroup.ID, delta[0].ID)
assert.Equal(t, "203.0.113.45", delta[0].IPList[0])
assert.NotEqual(t, groups[0].Checksum, delta[0].Checksum)
}
func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "sync group",
Type: "manual",
Enabled: true,
IPList: `["198.51.100.10"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
IDs: []uint{ipGroup.ID},
Checksums: map[string]string{},
})
require.NoError(t, err)
require.Len(t, result.Groups, 1)
assert.Equal(t, ipGroup.ID, result.Groups[0].ID)
assert.Equal(t, "198.51.100.10", result.Groups[0].IPList[0])
result, err = SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
IDs: []uint{ipGroup.ID},
Checksums: map[string]string{
strconv.FormatUint(uint64(ipGroup.ID), 10): result.Groups[0].Checksum,
},
})
require.NoError(t, err)
assert.Empty(t, result.Groups)
}
func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "disabled group",
Type: "manual",
Enabled: true,
IPList: `["203.0.113.10"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
ipGroup.Enabled = false
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.False(t, groups[0].Enabled)
assert.Empty(t, groups[0].IPList)
assert.NotEmpty(t, groups[0].Checksum)
}
@@ -5,6 +5,7 @@ package auth
const (
errInvalidParams = "无效的参数"
errUnauthorized = "无权进行此操作,未登录或 token 无效"
errPasswordLoginDisabled = "管理员关闭了密码登录"
errUsernameOrPasswordWrong = "用户名或密码错误"
errBannedAccount = "用户已被封禁"
@@ -0,0 +1,353 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
const (
errGitHubOAuthDisabled = "管理员未开启通过 GitHub 登录以及注册"
errWeChatOAuthDisabled = "管理员未开启通过微信登录以及注册"
errRegistrationClosed = "管理员关闭了新用户注册"
errGitHubAlreadyBound = "该 GitHub 账户已被绑定"
errWeChatAlreadyBound = "该微信账号已被绑定"
)
type githubOAuthResponse struct {
AccessToken string `json:"access_token"`
}
type githubUser struct {
Login string `json:"login"`
Name string `json:"name"`
Email string `json:"email"`
}
type wechatLoginResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data string `json:"data"`
}
// GitHubOAuth handles the legacy GET /oauth/github shortcut.
func GitHubOAuth(ctx context.Context, c *gin.Context, code string) (LegacyUser, error) {
if current := currentUserFromLegacyToken(ctx, c); current != nil {
if err := GitHubBind(ctx, c, current, code); err != nil {
return LegacyUser{}, err
}
return LegacyUser{}, nil
}
if !model.GitHubOAuthEnabled {
return LegacyUser{}, errors.New(errGitHubOAuthDisabled)
}
githubUser, err := getGitHubUserInfoByCode(code)
if err != nil {
return LegacyUser{}, err
}
user, err := findUserByShortcutBinding(ctx, githubUser.Login, "github", "GitHub")
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return LegacyUser{}, errors.New(errRegistrationClosed)
}
return LegacyUser{}, err
}
if !user.IsActive {
return LegacyUser{}, errors.New(errBannedAccount)
}
return finishLegacyLogin(ctx, c, user)
}
// GitHubBind binds a GitHub account to the current user.
func GitHubBind(ctx context.Context, c *gin.Context, current *model.User, code string) error {
if current == nil {
return errors.New(errUnauthorized)
}
if !model.GitHubOAuthEnabled {
return errors.New(errGitHubOAuthDisabled)
}
githubUser, err := getGitHubUserInfoByCode(code)
if err != nil {
return err
}
if _, err := findUserByShortcutBinding(ctx, githubUser.Login, "github", "GitHub"); err == nil {
return errors.New(errGitHubAlreadyBound)
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
return bindShortcutExternalAccount(ctx, current.ID, githubUser.Login, githubUser.Login, githubUser.Email, "github", "GitHub")
}
// WeChatOAuth handles the legacy GET /oauth/wechat shortcut.
func WeChatOAuth(ctx context.Context, c *gin.Context, code string) (LegacyUser, error) {
if !model.WeChatAuthEnabled {
return LegacyUser{}, errors.New(errWeChatOAuthDisabled)
}
wechatID, err := getWeChatIDByCode(code)
if err != nil {
return LegacyUser{}, err
}
user, err := findUserByShortcutBinding(ctx, wechatID, "wechat", "WeChat")
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return LegacyUser{}, errors.New(errRegistrationClosed)
}
return LegacyUser{}, err
}
if !user.IsActive {
return LegacyUser{}, errors.New(errBannedAccount)
}
return finishLegacyLogin(ctx, c, user)
}
// WeChatBind binds a WeChat account to the current user.
func WeChatBind(ctx context.Context, userID uint64, code string) error {
if userID == 0 {
return errors.New(errUnauthorized)
}
if !model.WeChatAuthEnabled {
return errors.New(errWeChatOAuthDisabled)
}
wechatID, err := getWeChatIDByCode(code)
if err != nil {
return err
}
if _, err := findUserByShortcutBinding(ctx, wechatID, "wechat", "WeChat"); err == nil {
return errors.New(errWeChatAlreadyBound)
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
return bindShortcutExternalAccount(ctx, userID, wechatID, wechatID, "", "wechat", "WeChat")
}
// EmailBind binds a verified email address to the current user.
func EmailBind(ctx context.Context, userID uint64, email, code string) error {
email = strings.TrimSpace(email)
code = strings.TrimSpace(code)
if userID == 0 {
return errors.New(errUnauthorized)
}
if email == "" || code == "" {
return errors.New(errInvalidParams)
}
if !verifyEmailCode(ctx, email, "register", code) {
return errors.New(errEmailCodeInvalid)
}
var user model.User
if err := db.DB(ctx).Where("id = ?", userID).First(&user).Error; err != nil {
return errors.New(errUserNotFound)
}
user.Email = email
return db.DB(ctx).Model(&user).Update("email", email).Error
}
func finishLegacyLogin(ctx context.Context, c *gin.Context, user *model.User) (LegacyUser, error) {
if user == nil {
return LegacyUser{}, errors.New(errUserNotFound)
}
user.LastLoginAt = time.Now()
if err := db.DB(ctx).Model(user).Update("last_login_at", user.LastLoginAt).Error; err != nil {
return LegacyUser{}, err
}
if err := setLoginSession(ctx, c, user); err != nil {
return LegacyUser{}, errors.New(errSaveSessionFailed)
}
token, err := issueLegacyAccessToken(ctx, user)
if err != nil {
return LegacyUser{}, err
}
logger.InfoF(ctx, "[LoginAudit] successful legacy shortcut login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, user, c.ClientIP())
return ToLegacyUser(user, token), nil
}
func getGitHubUserInfoByCode(code string) (*githubUser, error) {
code = strings.TrimSpace(code)
if code == "" {
return nil, errors.New(errInvalidParams)
}
values := map[string]string{
"client_id": model.GitHubClientId,
"client_secret": model.GitHubClientSecret,
"code": code,
}
jsonData, err := json.Marshal(values)
if err != nil {
return nil, err
}
client := http.Client{Timeout: 5 * time.Second}
req, err := http.NewRequest(http.MethodPost, "https://github.com/login/oauth/access_token", bytes.NewBuffer(jsonData))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
res, err := client.Do(req)
if err != nil {
slog.Error("github oauth access token request failed", "error", err)
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
}
defer res.Body.Close()
var oauthResponse githubOAuthResponse
if err := json.NewDecoder(res.Body).Decode(&oauthResponse); err != nil {
return nil, err
}
if strings.TrimSpace(oauthResponse.AccessToken) == "" {
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
}
req, err = http.NewRequest(http.MethodGet, "https://api.github.com/user", nil)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oauthResponse.AccessToken))
res2, err := client.Do(req)
if err != nil {
slog.Error("github user info request failed", "error", err)
return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!")
}
defer res2.Body.Close()
var ghUser githubUser
if err := json.NewDecoder(res2.Body).Decode(&ghUser); err != nil {
return nil, err
}
if strings.TrimSpace(ghUser.Login) == "" {
return nil, errors.New("返回值非法,用户字段为空,请稍后重试!")
}
return &ghUser, nil
}
func getWeChatIDByCode(code string) (string, error) {
code = strings.TrimSpace(code)
if code == "" {
return "", errors.New(errInvalidParams)
}
serverAddress := strings.TrimRight(strings.TrimSpace(model.WeChatServerAddress), "/")
if serverAddress == "" {
return "", errors.New(errWeChatOAuthDisabled)
}
req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("%s/api/wechat/user?code=%s", serverAddress, code), nil)
if err != nil {
return "", err
}
req.Header.Set("Authorization", model.WeChatServerToken)
client := http.Client{Timeout: 5 * time.Second}
httpResponse, err := client.Do(req)
if err != nil {
return "", err
}
defer func(body io.ReadCloser) {
if closeErr := body.Close(); closeErr != nil {
slog.Error("failed to close wechat response body", "error", closeErr)
}
}(httpResponse.Body)
var res wechatLoginResponse
if err := json.NewDecoder(httpResponse.Body).Decode(&res); err != nil {
return "", err
}
if !res.Success {
if strings.TrimSpace(res.Message) == "" {
return "", errors.New(errInvalidParams)
}
return "", errors.New(res.Message)
}
if strings.TrimSpace(res.Data) == "" {
return "", errors.New(errEmailCodeInvalid)
}
return strings.TrimSpace(res.Data), nil
}
func findUserByShortcutBinding(ctx context.Context, externalID string, sourceNames ...string) (*model.User, error) {
externalID = strings.TrimSpace(externalID)
if externalID == "" {
return nil, gorm.ErrRecordNotFound
}
query := db.DB(ctx).
Table("w_external_accounts AS ea").
Select("u.*").
Joins("JOIN w_users u ON u.id = ea.user_id").
Where("ea.external_id = ?", externalID)
if len(sourceNames) > 0 {
lowered := make([]string, 0, len(sourceNames))
for _, name := range sourceNames {
trimmed := strings.ToLower(strings.TrimSpace(name))
if trimmed != "" {
lowered = append(lowered, trimmed)
}
}
if len(lowered) > 0 {
query = query.
Joins("JOIN w_auth_sources s ON s.id = ea.auth_source_id").
Where("LOWER(s.name) IN ?", lowered)
}
}
var user model.User
if err := query.First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
func bindShortcutExternalAccount(ctx context.Context, userID uint64, externalID, externalUsername, email string, sourceNames ...string) error {
source, err := resolveShortcutAuthSource(ctx, sourceNames...)
if err != nil {
return err
}
return model.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID,
UserID: userID,
ExternalID: strings.TrimSpace(externalID),
ExternalUsername: strings.TrimSpace(externalUsername),
Email: strings.TrimSpace(email),
})
}
func resolveShortcutAuthSource(ctx context.Context, sourceNames ...string) (*model.AuthSource, error) {
for _, name := range sourceNames {
source, err := model.GetAuthSourceByName(ctx, name)
if err == nil {
return source, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
}
return nil, errors.New("认证源不存在")
}
@@ -10,6 +10,7 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/relay"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -143,6 +144,8 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil {
return nil, fmt.Errorf("update flared heartbeat: %w", err)
}
agent.RefreshAccessTokenCache(ctx, node)
persistFlaredObservability(ctx, node.NodeID, payload, now)
activeConfig, err := getActiveConfigMeta(ctx)
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
@@ -0,0 +1,50 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"fmt"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
"github.com/Rain-kl/Wavelet/internal/db"
"go.uber.org/zap"
)
const flaredRuntimeUnhealthyEventType = "flared_runtime_unhealthy"
func persistFlaredObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) {
connected := make([]string, 0, len(payload.ConnectedRelays))
for _, relay := range payload.ConnectedRelays {
connected = append(connected, fmt.Sprintf("%s:%s", relay.RelayNodeID, relay.Status))
}
managedTypes := map[string]struct{}{
flaredRuntimeUnhealthyEventType: {},
}
var events []agent.NodeHealthEvent
if payload.TunnelStatus == "unhealthy" {
events = append(events, agent.NodeHealthEvent{
EventType: flaredRuntimeUnhealthyEventType,
Severity: "critical",
Message: "openflared runtime is not healthy",
TriggeredAtUnix: reportedAt.Unix(),
Metadata: map[string]string{
"tunnel_status": payload.TunnelStatus,
"client_version": payload.ClientVersion,
"current_version": payload.CurrentVersion,
"current_checksum": payload.CurrentChecksum,
"connected_relays": strings.Join(connected, ","),
},
})
}
conn := db.DB(ctx)
if conn == nil {
return
}
if err := agent.ReconcileScopedNodeHealthEvents(conn, nodeID, events, reportedAt, managedTypes); err != nil {
zap.L().Error("persist flared health events failed", zap.String("node_id", nodeID), zap.Error(err))
}
}
@@ -0,0 +1,71 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupFlaredObservabilityTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareNode{},
&model.OpenFlareHealthEvent{},
))
db.SetDB(sqliteDB)
option.ResetInitializationForTest()
agent.ResetAuthCacheForTest()
return func() {
db.SetDB(nil)
option.ResetInitializationForTest()
agent.ResetAuthCacheForTest()
}
}
func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) {
cleanup := setupFlaredObservabilityTestDB(t)
defer cleanup()
ctx := context.Background()
node := &model.OpenFlareNode{
NodeID: "node-flared-unhealthy",
Name: "flared-unhealthy",
AccessToken: "tunnel-token-unhealthy",
Status: "pending",
NodeType: "tunnel_client",
}
require.NoError(t, db.DB(ctx).Create(node).Error)
_, err := Heartbeat(ctx, node, HeartbeatPayload{
ClientVersion: "v0.2.0",
FrpVersion: "0.61.0",
TunnelStatus: "unhealthy",
CurrentVersion: "v1",
CurrentChecksum: "checksum-1",
})
require.NoError(t, err)
events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 20)
require.NoError(t, err)
require.Len(t, events, 1)
assert.Equal(t, flaredRuntimeUnhealthyEventType, events[0].EventType)
assert.Equal(t, "active", events[0].Status)
}
@@ -34,6 +34,11 @@ func IsValidProvider(provider string) bool {
return pkggeoip.IsValidProvider(provider)
}
// GeoInfoFromIP resolves geographic information using the configured default provider.
func GeoInfoFromIP(ip net.IP) (*pkggeoip.GeoInfo, error) {
return pkggeoip.GetGeoInfo(ip)
}
// Lookup resolves geographic information for rawIP using the given provider.
func Lookup(provider, rawIP string) (*LookupView, error) {
trimmedProvider := strings.TrimSpace(provider)
@@ -46,6 +46,11 @@ func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) {
&model.OpenFlareNode{},
&model.OpenFlareOption{},
&model.OpenFlareApplyLog{},
&model.OpenFlareNodeSystemProfile{},
&model.OpenFlareMetricSnapshot{},
&model.OpenFlareHealthEvent{},
&model.OpenFlareNodeObservationFrps{},
&model.OpenFlareNodeObservationFrpc{},
&configVersionRecord{},
))
@@ -184,6 +184,21 @@ func TestGETOptionRequiresRootAuth(t *testing.T) {
})
}
func TestGETNodesWithOpenFlareToken(t *testing.T) {
dbConn, r := setupAuthOptionIntegration(t)
require.NoError(t, dbConn.AutoMigrate(&model.OpenFlareNode{}))
seedUser(t, dbConn, "admin", "password123", true)
rootToken := loginAndGetToken(t, r, "admin", "password123")
w := performJSONRequest(t, r, http.MethodGet, "/api/nodes/", nil, map[string]string{
compat.OpenFlareTokenHeader(): rootToken,
})
assert.Equal(t, http.StatusOK, w.Code)
env := decodeEnvelope(t, w)
assert.True(t, env.Success, "message=%s", env.Message)
}
func TestOptionHotReloadAfterUpdate(t *testing.T) {
dbConn, r := setupAuthOptionIntegration(t)
seedUser(t, dbConn, "admin", "password123", true)
@@ -0,0 +1,52 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package legacy
import (
ofauth "github.com/Rain-kl/Wavelet/internal/apps/openflare/auth"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/compat"
"github.com/gin-gonic/gin"
)
// GitHubOAuth handles GET /oauth/github for legacy GitHub login or bind.
func GitHubOAuth(c *gin.Context) {
user, err := ofauth.GitHubOAuth(c.Request.Context(), c, c.Query("code"))
if err != nil {
compat.Fail(c, err.Error())
return
}
if user.ID == 0 {
compat.OKMessage(c, "bind")
return
}
compat.OK(c, user)
}
// WeChatOAuth handles GET /oauth/wechat for legacy WeChat login.
func WeChatOAuth(c *gin.Context) {
user, err := ofauth.WeChatOAuth(c.Request.Context(), c, c.Query("code"))
if err != nil {
compat.Fail(c, err.Error())
return
}
compat.OK(c, user)
}
// WeChatBind handles GET /oauth/wechat/bind for legacy WeChat account binding.
func WeChatBind(c *gin.Context) {
if err := ofauth.WeChatBind(c.Request.Context(), callerUserID(c), c.Query("code")); err != nil {
compat.Fail(c, err.Error())
return
}
compat.OKMessage(c, "")
}
// EmailBind handles GET /oauth/email/bind for legacy email binding.
func EmailBind(c *gin.Context) {
if err := ofauth.EmailBind(c.Request.Context(), callerUserID(c), c.Query("email"), c.Query("code")); err != nil {
compat.Fail(c, err.Error())
return
}
compat.OKMessage(c, "")
}
@@ -9,11 +9,6 @@ import (
"github.com/gin-gonic/gin"
)
// bridgeOpenFlareToken maps OpenFlare-Token to X-Access-Token for compat auth middleware.
func bridgeOpenFlareToken() gin.HandlerFunc {
return compat.BridgeOpenFlareToken()
}
// legacyCapAuth verifies PoW CAPTCHA for legacy login using OpenFlare response format.
func legacyCapAuth(scope string) gin.HandlerFunc {
mgr := cap.GetDefaultManager()
@@ -4,10 +4,14 @@
// Package legacy registers OpenFlare /api/* compatibility routes for the old frontend.
package legacy
import "github.com/gin-gonic/gin"
import (
"github.com/Rain-kl/Wavelet/internal/apps/openflare/compat"
"github.com/gin-gonic/gin"
)
// RegisterRoutes mounts all OpenFlare legacy API routes under the /api group.
func RegisterRoutes(apiGroup *gin.RouterGroup) {
apiGroup.Use(compat.BridgeOpenFlareToken())
registerAuthRoutes(apiGroup)
registerOptionRoutes(apiGroup)
registerOriginRoutes(apiGroup)
@@ -16,12 +16,16 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) {
oauthGroup := apiGroup.Group("/oauth")
{
oauthGroup.GET("/github", GitHubOAuth)
oauthGroup.GET("/wechat", WeChatOAuth)
oauthGroup.GET("/wechat/bind", compat.BridgeOpenFlareToken(), compat.UserAuth(), WeChatBind)
oauthGroup.GET("/email/bind", compat.BridgeOpenFlareToken(), compat.UserAuth(), EmailBind)
oauthGroup.GET("/:source/authorize", OAuthAuthorize)
oauthGroup.GET("/:source/callback", OAuthCallback)
oauthGroup.POST("/link-existing", LinkExistingOAuthAccount)
externalAccounts := oauthGroup.Group("/external-accounts")
externalAccounts.Use(bridgeOpenFlareToken(), compat.UserAuth())
externalAccounts.Use(compat.UserAuth())
{
externalAccounts.GET("/", ListExternalAccounts)
externalAccounts.POST("/:id/delete", DeleteExternalAccount)
@@ -41,7 +45,7 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) {
userGroup.GET("/logout", Logout)
selfGroup := userGroup.Group("/")
selfGroup.Use(bridgeOpenFlareToken(), compat.UserAuth())
selfGroup.Use(compat.UserAuth())
{
selfGroup.GET("/self", GetSelf)
selfGroup.POST("/self/update", UpdateSelf)
@@ -50,7 +54,7 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) {
}
adminGroup := userGroup.Group("/")
adminGroup.Use(bridgeOpenFlareToken(), compat.AdminAuth())
adminGroup.Use(compat.AdminAuth())
{
adminGroup.GET("/", GetAllUsers)
adminGroup.GET("/search", SearchUsers)
@@ -63,7 +67,7 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) {
}
authSourceGroup := apiGroup.Group("/auth-sources")
authSourceGroup.Use(bridgeOpenFlareToken(), compat.RootAuth())
authSourceGroup.Use(compat.RootAuth())
{
authSourceGroup.GET("/", ListAuthSources)
authSourceGroup.POST("/", CreateAuthSource)
@@ -20,7 +20,7 @@ func RegisterRoutes(apiGroup *gin.RouterGroup) {
apiGroup.GET("/about", getAboutHandler)
optionRoute := apiGroup.Group("/option")
optionRoute.Use(compat.BridgeOpenFlareToken(), compat.RootAuth())
optionRoute.Use(compat.RootAuth())
{
optionRoute.GET("/", listOptionsHandler)
optionRoute.POST("/update", updateOptionHandler)
@@ -30,7 +30,7 @@ func RegisterRoutes(apiGroup *gin.RouterGroup) {
}
uptimeKumaRoute := apiGroup.Group("/uptimekuma")
uptimeKumaRoute.Use(compat.BridgeOpenFlareToken(), compat.RootAuth())
uptimeKumaRoute.Use(compat.RootAuth())
{
uptimeKumaRoute.POST("/sync", syncUptimeKumaHandler)
}
+19 -16
View File
@@ -4,20 +4,23 @@
package pages
const (
errPagesProjectNotFound = "Pages 项目不存在"
errPagesSlugExists = "Pages 项目标识已存在"
errPagesNameRequired = "Pages 项目名称不能为空"
errPagesSlugInvalid = "Pages 项目标识只能包含小写字母、数字和连字符"
errPagesDeleteReferenced = "Pages 项目已被规则引用,不能删除"
errPagesDeploymentNotFound = "Pages 部署不存在"
errPagesDeploymentMismatch = "Pages 部署不属于该项目"
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
errPagesPackageMissing = "缺少 Pages 部署包"
errPagesPackageNotZip = "Pages 部署包必须是 .zip 文件"
errPagesPackageInvalidZip = "Pages 部署包不是有效 zip 文件"
errPagesPackageEmpty = "Pages 部署包不能为空"
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空"
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL"
errPagesProjectNotFound = "Pages 项目不存在"
errPagesSlugExists = "Pages 项目标识已存在"
errPagesNameRequired = "Pages 项目名称不能为空"
errPagesSlugInvalid = "Pages 项目标识只能包含小写字母、数字和连字符"
errPagesDeleteReferenced = "Pages 项目已被规则引用,不能删除"
errPagesDeploymentNotFound = "Pages 部署不存在"
errPagesDeploymentMismatch = "Pages 部署不属于该项目"
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
errPagesPackageMissing = "缺少 Pages 部署包"
errPagesPackageNotZip = "Pages 部署包必须是 .zip 文件"
errPagesPackageInvalidZip = "Pages 部署包不是有效 zip 文件"
errPagesPackageEmpty = "Pages 部署包不能为空"
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空"
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL"
errPagesPackagePathEmpty = "Pages 部署包路径为空"
errPagesPackageNotInActiveConfig = "Pages 部署尚未进入激活配置"
errPagesInvalidSnapshotFormat = "配置快照格式无效"
)
@@ -5,6 +5,7 @@ package pages
import (
"context"
"encoding/json"
"errors"
"fmt"
"mime/multipart"
@@ -341,6 +342,77 @@ func ActivateDeployment(ctx context.Context, projectID uint, deploymentID uint)
return GetProject(ctx, project.ID)
}
// GetDeploymentPackagePath returns the on-disk artifact path and download filename for an agent package request.
func GetDeploymentPackagePath(ctx context.Context, deploymentID uint) (string, string, error) {
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID)
if err != nil {
return "", "", err
}
if err = ensureDeploymentInActiveSnapshot(ctx, deployment.ID); err != nil {
return "", "", err
}
if strings.TrimSpace(deployment.ArtifactPath) == "" {
return "", "", errors.New(errPagesPackagePathEmpty)
}
if _, err = os.Stat(deployment.ArtifactPath); err != nil {
return "", "", fmt.Errorf("Pages 部署包不存在: %w", err)
}
return deployment.ArtifactPath, fmt.Sprintf("pages-deployment-%d.zip", deployment.ID), nil
}
func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) error {
version, err := model.GetActiveConfigVersion(ctx)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errPagesPackageNotInActiveConfig)
}
return err
}
routes, err := parseSnapshotRoutes(version.SnapshotJSON)
if err != nil {
return err
}
for _, route := range routes {
if route.UpstreamType != "pages" || route.PagesDeployment == nil {
continue
}
if route.PagesDeployment.DeploymentID == deploymentID {
return nil
}
}
return errors.New(errPagesPackageNotInActiveConfig)
}
type snapshotPagesDeployment struct {
DeploymentID uint `json:"deployment_id"`
}
type snapshotRouteRef struct {
UpstreamType string `json:"upstream_type"`
PagesDeployment *snapshotPagesDeployment `json:"pages_deployment"`
}
func parseSnapshotRoutes(snapshotJSON string) ([]snapshotRouteRef, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
return []snapshotRouteRef{}, nil
}
if strings.HasPrefix(text, "[") {
var routes []snapshotRouteRef
if err := json.Unmarshal([]byte(text), &routes); err != nil {
return nil, errors.New(errPagesInvalidSnapshotFormat)
}
return routes, nil
}
var snapshot struct {
Routes []snapshotRouteRef `json:"routes"`
}
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
return nil, errors.New(errPagesInvalidSnapshotFormat)
}
return snapshot.Routes, nil
}
// DeleteDeployment 删除 Pages 部署。
func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error {
project, err := model.GetPagesProjectByID(ctx, projectID)
@@ -4,7 +4,13 @@
package pages
import (
"archive/zip"
"bytes"
"context"
"fmt"
"mime/multipart"
"net/http/httptest"
"strconv"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
@@ -26,6 +32,7 @@ func setupPagesTestDB(t *testing.T) func() {
&model.PagesProject{},
&model.PagesDeployment{},
&model.PagesDeploymentFile{},
&model.ConfigVersion{},
))
db.SetDB(sqliteDB)
@@ -82,3 +89,80 @@ func TestCreateProjectRejectsUnsafeFallbackPath(t *testing.T) {
require.Error(t, err)
assert.Contains(t, err.Error(), "回退路径")
}
func TestGetDeploymentPackagePathRequiresActiveConfigSnapshot(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
project, err := CreateProject(ctx, Input{
Name: "Published Site",
Slug: "published-site",
Enabled: true,
})
require.NoError(t, err)
deployment, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "site.zip", testPagesZip(t, map[string]string{
"index.html": "ok",
})), "root")
require.NoError(t, err)
_, err = ActivateDeployment(ctx, project.ID, deployment.ID)
require.NoError(t, err)
_, _, err = GetDeploymentPackagePath(ctx, deployment.ID)
require.Error(t, err)
assert.Contains(t, err.Error(), "激活配置")
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "v2026-001",
SnapshotJSON: fmt.Sprintf(`{"routes":[{"upstream_type":"pages","pages_deployment":{"deployment_id":%d}}]}`, deployment.ID),
MainConfig: "",
RenderedConfig: "",
SupportFilesJSON: "[]",
Checksum: "test-checksum",
IsActive: true,
CreatedBy: "test",
}).Error)
filePath, fileName, err := GetDeploymentPackagePath(ctx, deployment.ID)
require.NoError(t, err)
assert.NotEmpty(t, filePath)
assert.Equal(t, "pages-deployment-"+strconv.FormatUint(uint64(deployment.ID), 10)+".zip", fileName)
}
func testPagesZip(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
writer := zip.NewWriter(&buffer)
for name, content := range files {
file, err := writer.Create(name)
require.NoError(t, err)
_, err = file.Write([]byte(content))
require.NoError(t, err)
}
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func testPagesMultipartFile(t *testing.T, fileName string, content []byte) *multipart.FileHeader {
t.Helper()
var body bytes.Buffer
writer := multipart.NewWriter(&body)
part, err := writer.CreateFormFile("package", fileName)
require.NoError(t, err)
_, err = part.Write(content)
require.NoError(t, err)
require.NoError(t, writer.Close())
req := httptest.NewRequest("POST", "/", &body)
req.Header.Set("Content-Type", writer.FormDataContentType())
require.NoError(t, req.ParseMultipartForm(int64(len(content))+1024))
file, header, err := req.FormFile("package")
require.NoError(t, err)
file.Close()
return header
}
@@ -9,19 +9,38 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
const nodeStatusOnline = "online"
// ProxyStat describes a single frps proxy reported by the relay.
type ProxyStat struct {
Name string `json:"name"`
Type string `json:"type"`
Status string `json:"status"`
ClientVersion string `json:"client_version"`
LastStartTime string `json:"last_start_time"`
LastCloseTime string `json:"last_close_time"`
ClientAddr string `json:"client_addr"`
}
// HeartbeatPayload is sent by OpenFlareRelay on each heartbeat.
type HeartbeatPayload struct {
Version string `json:"version"`
ExtVersion string `json:"frp_version"`
RelayStatus string `json:"relay_status"`
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
ExtVersion string `json:"frp_version"`
RelayStatus string `json:"relay_status"`
FrpsConnCount int `json:"frps_connections"`
FrpsProxyCount int `json:"frps_proxy_count"`
FrpsClientCount int `json:"frps_client_count"`
FrpsProxies []ProxyStat `json:"frps_proxies,omitempty"`
Name string `json:"name"`
IP string `json:"ip"`
Profile *agent.NodeSystemProfile `json:"profile,omitempty"`
Snapshot *agent.NodeMetricSnapshot `json:"snapshot,omitempty"`
HealthEvents []agent.NodeHealthEvent `json:"health_events,omitempty"`
}
// Config is the frps configuration sent to the relay.
@@ -109,6 +128,11 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil {
return nil, fmt.Errorf("update relay heartbeat: %w", err)
}
if err := reconcileRelayHealthEvents(ctx, node.NodeID, payload.RelayStatus, now); err != nil {
return nil, fmt.Errorf("reconcile relay health events: %w", err)
}
agent.RefreshAccessTokenCache(ctx, node)
persistRelayHeartbeatObservability(ctx, node.NodeID, payload, now)
return &HeartbeatResponse{
RelayConfig: buildRelayConfig(node),
@@ -0,0 +1,174 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"encoding/json"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupRelayTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareNode{},
&model.OpenFlareOption{},
&model.OpenFlareNodeSystemProfile{},
&model.OpenFlareMetricSnapshot{},
&model.OpenFlareHealthEvent{},
&model.OpenFlareNodeObservationFrps{},
))
db.SetDB(sqliteDB)
option.ResetInitializationForTest()
agent.ResetAuthCacheForTest()
return func() {
db.SetDB(nil)
option.ResetInitializationForTest()
agent.ResetAuthCacheForTest()
}
}
func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) {
cleanup := setupRelayTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC().Truncate(time.Second)
node := &model.OpenFlareNode{
NodeID: "node-relay-observe",
Name: "relay-1",
AccessToken: "relay-token",
Status: "pending",
NodeType: "tunnel_relay",
RelayStatus: "unknown",
}
require.NoError(t, db.DB(ctx).Create(node).Error)
proxies := []ProxyStat{
{
Name: "proxy-a",
Type: "http",
Status: "online",
ClientVersion: "0.61.0",
ClientAddr: "10.0.0.2:12345",
},
}
_, err := Heartbeat(ctx, node, HeartbeatPayload{
Version: "v0.1.0",
ExtVersion: "0.61.0",
RelayStatus: "healthy",
FrpsConnCount: 7,
FrpsProxyCount: 3,
FrpsClientCount: 2,
FrpsProxies: proxies,
Name: "relay-runtime",
IP: "203.0.113.9",
Profile: &agent.NodeSystemProfile{
Hostname: "relay-runtime",
OSName: "Ubuntu",
OSVersion: "24.04",
Architecture: "amd64",
CPUCores: 4,
ReportedAtUnix: now.Unix(),
},
Snapshot: &agent.NodeMetricSnapshot{
CapturedAtUnix: now.Unix(),
CPUUsagePercent: 12.5,
NetworkRxBytes: 1024,
NetworkTxBytes: 2048,
},
HealthEvents: []agent.NodeHealthEvent{},
})
require.NoError(t, err)
var stored model.OpenFlareNode
require.NoError(t, db.DB(ctx).Where("node_id = ?", node.NodeID).First(&stored).Error)
assert.Equal(t, "online", stored.Status)
assert.Equal(t, "healthy", stored.RelayStatus)
assert.Equal(t, "203.0.113.9", stored.IP)
assert.Equal(t, "v0.1.0", stored.Version)
assert.Equal(t, "0.61.0", stored.ExtVersion)
profile, err := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID)
require.NoError(t, err)
assert.Equal(t, "relay-runtime", profile.Hostname)
assert.Equal(t, "Ubuntu", profile.OSName)
snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10)
require.NoError(t, err)
require.Len(t, snapshots, 1)
assert.Equal(t, 12.5, snapshots[0].CPUUsagePercent)
frpsObs, err := model.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1)
require.NoError(t, err)
require.Len(t, frpsObs, 1)
assert.Equal(t, 7, frpsObs[0].FrpsConnections)
assert.Equal(t, 3, frpsObs[0].FrpsProxyCount)
assert.Equal(t, 2, frpsObs[0].FrpsClientCount)
var decoded []ProxyStat
require.NoError(t, json.Unmarshal([]byte(frpsObs[0].FrpsProxies), &decoded))
require.Len(t, decoded, 1)
assert.Equal(t, "proxy-a", decoded[0].Name)
assert.Equal(t, "online", decoded[0].Status)
}
func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) {
cleanup := setupRelayTestDB(t)
defer cleanup()
ctx := context.Background()
node := &model.OpenFlareNode{
NodeID: "node-relay-unhealthy",
Name: "relay-unhealthy",
AccessToken: "relay-token-unhealthy",
Status: "pending",
NodeType: "tunnel_relay",
RelayStatus: "healthy",
}
require.NoError(t, db.DB(ctx).Create(node).Error)
_, err := Heartbeat(ctx, node, HeartbeatPayload{
Version: "v0.1.0",
ExtVersion: "0.61.0",
RelayStatus: "unhealthy",
})
require.NoError(t, err)
events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, true, 10)
require.NoError(t, err)
require.Len(t, events, 1)
assert.Equal(t, relayFrpsUnhealthyEventType, events[0].EventType)
assert.Equal(t, "active", events[0].Status)
_, err = Heartbeat(ctx, node, HeartbeatPayload{
Version: "v0.1.0",
ExtVersion: "0.61.0",
RelayStatus: "healthy",
})
require.NoError(t, err)
events, err = model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 10)
require.NoError(t, err)
require.Len(t, events, 1)
assert.Equal(t, "resolved", events[0].Status)
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"go.uber.org/zap"
"gorm.io/gorm"
)
const relayFrpsUnhealthyEventType = "frps_unhealthy"
func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus string, reportedAt time.Time) error {
if relayStatus == "unknown" {
return nil
}
managedTypes := map[string]struct{}{
relayFrpsUnhealthyEventType: {},
}
events := []agent.NodeHealthEvent{}
if relayStatus == "unhealthy" {
events = append(events, agent.NodeHealthEvent{
EventType: relayFrpsUnhealthyEventType,
Severity: "critical",
Message: "frps runtime is not healthy",
TriggeredAtUnix: reportedAt.Unix(),
Metadata: map[string]string{
"relay_status": relayStatus,
},
})
}
conn := db.DB(ctx)
if conn == nil {
return nil
}
return conn.Transaction(func(tx *gorm.DB) error {
return agent.ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, managedTypes)
})
}
func persistRelayHeartbeatObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) {
agent.PersistHeartbeatObservability(ctx, nodeID, agent.NodePayload{
Profile: payload.Profile,
Snapshot: payload.Snapshot,
HealthEvents: payload.HealthEvents,
}, reportedAt)
conn := db.DB(ctx)
if conn == nil {
return
}
frpsObs := &model.OpenFlareNodeObservationFrps{
NodeID: nodeID,
CapturedAt: reportedAt,
FrpsConnections: payload.FrpsConnCount,
FrpsProxyCount: payload.FrpsProxyCount,
FrpsClientCount: payload.FrpsClientCount,
FrpsProxies: agent.MarshalJSON(payload.FrpsProxies),
}
if err := conn.Create(frpsObs).Error; err != nil {
zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err))
}
}