refactor(backend): rename OpenFlare directory to lowercase openflare

This commit is contained in:
ryan
2026-08-30 17:43:23 +08:00
parent 06d5fedbfc
commit c93ff6674f
543 changed files with 819 additions and 819 deletions
@@ -0,0 +1,293 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cloudflare manages Cloudflare DNS pointing for OpenFlare domains.
package cloudflare
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"Wavelet/pkg/httppool"
)
const (
defaultAPIBaseURL = "https://api.cloudflare.com/client/v4"
defaultHTTPTimeout = 20 * time.Second
maxRequestAttempts = 3
maxResponseBodyBytes = 1 << 20
defaultRetryDelay = 200 * time.Millisecond
maxRetryAfterSeconds = 2
)
// Zone is a Cloudflare DNS zone.
type Zone struct {
ID string `json:"id"`
Name string `json:"name"`
}
// DNSRecord is a Cloudflare DNS record.
type DNSRecord struct {
ID string `json:"id"`
Type string `json:"type"`
Name string `json:"name"`
Content string `json:"content"`
Proxied bool `json:"proxied"`
TTL int `json:"ttl"`
}
// RecordInput is the desired Cloudflare DNS record payload.
type RecordInput struct {
Type string `json:"type"`
Name string `json:"name"`
Content string `json:"content"`
Proxied bool `json:"proxied"`
TTL int `json:"ttl"`
}
// Client describes the Cloudflare operations used by pointing reconciliation.
type Client interface {
VerifyToken(context.Context) error
FindZone(context.Context, string) (*Zone, error)
GetRecord(context.Context, string, string) (*DNSRecord, error)
ListARecords(context.Context, string, string) ([]DNSRecord, error)
CreateARecord(context.Context, string, RecordInput) (*DNSRecord, error)
UpdateARecord(context.Context, string, string, RecordInput) (*DNSRecord, error)
DeleteRecord(context.Context, string, string) error
}
// HTTPClient implements Client with Cloudflare's v4 HTTP API.
type HTTPClient struct {
token string
baseURL string
httpClient *http.Client
}
// ClientOption configures HTTPClient.
type ClientOption func(*HTTPClient)
// WithBaseURL overrides the Cloudflare API base URL.
func WithBaseURL(baseURL string) ClientOption {
return func(client *HTTPClient) { client.baseURL = strings.TrimRight(baseURL, "/") }
}
// WithHTTPClient overrides the HTTP transport.
func WithHTTPClient(httpClient *http.Client) ClientOption {
return func(client *HTTPClient) { client.httpClient = httpClient }
}
// NewHTTPClient creates a Cloudflare HTTP client.
func NewHTTPClient(token string, options ...ClientOption) *HTTPClient {
client := &HTTPClient{
token: strings.TrimSpace(token),
baseURL: defaultAPIBaseURL,
httpClient: httppool.NewClient(defaultHTTPTimeout),
}
for _, option := range options {
option(client)
}
return client
}
type apiError struct {
Code int `json:"code"`
Message string `json:"message"`
}
type apiEnvelope[T any] struct {
Success bool `json:"success"`
Errors []apiError `json:"errors"`
Result T `json:"result"`
}
// VerifyToken verifies that the configured API token is active.
func (client *HTTPClient) VerifyToken(ctx context.Context) error {
var result struct {
Status string `json:"status"`
}
if err := client.do(ctx, http.MethodGet, "/user/tokens/verify", nil, nil, &result); err != nil {
return err
}
if result.Status != "active" {
return errors.New("cloudflare API Token 未激活")
}
return nil
}
// FindZone returns the exact Cloudflare zone name.
func (client *HTTPClient) FindZone(ctx context.Context, name string) (*Zone, error) {
query := url.Values{"name": {strings.TrimSpace(name)}, "status": {"active"}, "per_page": {"2"}}
var zones []Zone
if err := client.do(ctx, http.MethodGet, "/zones", query, nil, &zones); err != nil {
return nil, err
}
if len(zones) != 1 {
return nil, fmt.Errorf("cloudflare 中未找到唯一 Zone %s", name)
}
return &zones[0], nil
}
// GetRecord returns a DNS record by ID.
func (client *HTTPClient) GetRecord(ctx context.Context, zoneID, recordID string) (*DNSRecord, error) {
var record DNSRecord
path := "/zones/" + url.PathEscape(zoneID) + "/dns_records/" + url.PathEscape(recordID)
if err := client.do(ctx, http.MethodGet, path, nil, nil, &record); err != nil {
return nil, err
}
return &record, nil
}
// ListARecords lists exact-name A records.
func (client *HTTPClient) ListARecords(ctx context.Context, zoneID, name string) ([]DNSRecord, error) {
query := url.Values{"type": {"A"}, "name": {strings.TrimSpace(name)}, "per_page": {"100"}}
var records []DNSRecord
path := "/zones/" + url.PathEscape(zoneID) + "/dns_records"
if err := client.do(ctx, http.MethodGet, path, query, nil, &records); err != nil {
return nil, err
}
return records, nil
}
// CreateARecord creates an A record.
func (client *HTTPClient) CreateARecord(ctx context.Context, zoneID string, input RecordInput) (*DNSRecord, error) {
var record DNSRecord
path := "/zones/" + url.PathEscape(zoneID) + "/dns_records"
if err := client.do(ctx, http.MethodPost, path, nil, input, &record); err != nil {
return nil, err
}
return &record, nil
}
// UpdateARecord replaces an A record.
func (client *HTTPClient) UpdateARecord(ctx context.Context, zoneID, recordID string, input RecordInput) (*DNSRecord, error) {
var record DNSRecord
path := "/zones/" + url.PathEscape(zoneID) + "/dns_records/" + url.PathEscape(recordID)
if err := client.do(ctx, http.MethodPut, path, nil, input, &record); err != nil {
return nil, err
}
return &record, nil
}
// DeleteRecord deletes a DNS record.
func (client *HTTPClient) DeleteRecord(ctx context.Context, zoneID, recordID string) error {
var result struct {
ID string `json:"id"`
}
path := "/zones/" + url.PathEscape(zoneID) + "/dns_records/" + url.PathEscape(recordID)
return client.do(ctx, http.MethodDelete, path, nil, nil, &result)
}
func (client *HTTPClient) do(ctx context.Context, method, path string, query url.Values, body, result any) error {
encodedBody, err := encodeRequestBody(body)
if err != nil {
return err
}
requestURL := buildRequestURL(client.baseURL, path, query)
for attempt := range maxRequestAttempts {
statusCode, retryHeader, responseBody, requestErr := client.send(ctx, method, requestURL, encodedBody)
if requestErr != nil {
return requestErr
}
if statusCode == http.StatusTooManyRequests && attempt < maxRequestAttempts-1 {
if waitErr := waitForRetry(ctx, retryAfter(retryHeader)); waitErr != nil {
return waitErr
}
continue
}
return decodeAPIResponse(statusCode, responseBody, result)
}
return errors.New("cloudflare API 请求超过重试次数")
}
func encodeRequestBody(body any) ([]byte, error) {
if body == nil {
return nil, nil
}
encodedBody, err := json.Marshal(body)
if err != nil {
return nil, fmt.Errorf("encode Cloudflare request: %w", err)
}
return encodedBody, nil
}
func buildRequestURL(baseURL, path string, query url.Values) string {
requestURL := baseURL + path
if len(query) > 0 {
requestURL += "?" + query.Encode()
}
return requestURL
}
func (client *HTTPClient) send(ctx context.Context, method, requestURL string, body []byte) (int, string, []byte, error) {
request, err := http.NewRequestWithContext(ctx, method, requestURL, bytes.NewReader(body))
if err != nil {
return 0, "", nil, fmt.Errorf("create Cloudflare request: %w", err)
}
request.Header.Set("Authorization", "Bearer "+client.token)
request.Header.Set("Content-Type", "application/json")
response, err := client.httpClient.Do(request)
if err != nil {
return 0, "", nil, fmt.Errorf("cloudflare API 请求失败: %w", err)
}
responseBody, readErr := io.ReadAll(io.LimitReader(response.Body, maxResponseBodyBytes))
closeErr := response.Body.Close()
if readErr != nil {
return 0, "", nil, fmt.Errorf("read Cloudflare response: %w", readErr)
}
if closeErr != nil {
return 0, "", nil, fmt.Errorf("close Cloudflare response: %w", closeErr)
}
return response.StatusCode, response.Header.Get("Retry-After"), responseBody, nil
}
func waitForRetry(ctx context.Context, delay time.Duration) error {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}
func decodeAPIResponse(statusCode int, responseBody []byte, result any) error {
var envelope apiEnvelope[json.RawMessage]
if err := json.Unmarshal(responseBody, &envelope); err != nil {
return fmt.Errorf("decode Cloudflare response: %w", err)
}
if statusCode < http.StatusOK || statusCode >= http.StatusMultipleChoices || !envelope.Success {
message := "cloudflare API 请求失败"
if len(envelope.Errors) > 0 && strings.TrimSpace(envelope.Errors[0].Message) != "" {
message = envelope.Errors[0].Message
}
return errors.New(message)
}
if result == nil || len(envelope.Result) == 0 || string(envelope.Result) == "null" {
return nil
}
if err := json.Unmarshal(envelope.Result, result); err != nil {
return fmt.Errorf("decode Cloudflare result: %w", err)
}
return nil
}
func retryAfter(value string) time.Duration {
seconds, err := strconv.Atoi(strings.TrimSpace(value))
if err != nil || seconds <= 0 {
return defaultRetryDelay
}
if seconds > maxRetryAfterSeconds {
seconds = maxRetryAfterSeconds
}
return time.Duration(seconds) * time.Second
}
@@ -0,0 +1,90 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
)
func TestHTTPClientVerifyTokenAndManageARecord(t *testing.T) {
mux := http.NewServeMux()
mux.HandleFunc("/user/tokens/verify", func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("Authorization"); got != "Bearer test-token" {
t.Errorf("Authorization = %q, want Bearer test-token", got)
}
writeCFTestResponse(t, w, map[string]any{"status": "active"})
})
mux.HandleFunc("/zones", func(w http.ResponseWriter, r *http.Request) {
if got := r.URL.Query().Get("name"); got != "example.com" {
t.Errorf("zone name = %q, want example.com", got)
}
writeCFTestResponse(t, w, []map[string]any{{"id": "zone-1", "name": "example.com"}})
})
mux.HandleFunc("/zones/zone-1/dns_records", func(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
writeCFTestResponse(t, w, []map[string]any{})
case http.MethodPost:
var input RecordInput
if err := json.NewDecoder(r.Body).Decode(&input); err != nil {
t.Fatalf("Decode(create) error = %v", err)
}
if input.Type != "A" || input.Name != "api.example.com" || input.Content != "203.0.113.10" || !input.Proxied {
t.Errorf("create input = %+v", input)
}
writeCFTestResponse(t, w, map[string]any{"id": "record-1", "type": "A", "name": input.Name, "content": input.Content, "proxied": input.Proxied})
default:
w.WriteHeader(http.StatusMethodNotAllowed)
}
})
mux.HandleFunc("/zones/zone-1/dns_records/record-1", func(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodPut:
writeCFTestResponse(t, w, map[string]any{"id": "record-1", "type": "A", "name": "api.example.com", "content": "203.0.113.11", "proxied": false})
case http.MethodDelete:
writeCFTestResponse(t, w, map[string]any{"id": "record-1"})
default:
w.WriteHeader(http.StatusMethodNotAllowed)
}
})
server := httptest.NewServer(mux)
t.Cleanup(server.Close)
client := NewHTTPClient("test-token", WithBaseURL(server.URL), WithHTTPClient(server.Client()))
ctx := context.Background()
if err := client.VerifyToken(ctx); err != nil {
t.Fatalf("VerifyToken() error = %v", err)
}
zone, err := client.FindZone(ctx, "example.com")
if err != nil || zone.ID != "zone-1" {
t.Fatalf("FindZone() = %+v, %v", zone, err)
}
records, err := client.ListARecords(ctx, zone.ID, "api.example.com")
if err != nil || len(records) != 0 {
t.Fatalf("ListARecords() = %+v, %v", records, err)
}
record, err := client.CreateARecord(ctx, zone.ID, RecordInput{Type: "A", Name: "api.example.com", Content: "203.0.113.10", Proxied: true, TTL: 1})
if err != nil || record.ID != "record-1" {
t.Fatalf("CreateARecord() = %+v, %v", record, err)
}
if _, err := client.UpdateARecord(ctx, zone.ID, record.ID, RecordInput{Type: "A", Name: record.Name, Content: "203.0.113.11", TTL: 300}); err != nil {
t.Fatalf("UpdateARecord() error = %v", err)
}
if err := client.DeleteRecord(ctx, zone.ID, record.ID); err != nil {
t.Fatalf("DeleteRecord() error = %v", err)
}
}
func writeCFTestResponse(t *testing.T, w http.ResponseWriter, result any) {
t.Helper()
w.Header().Set("Content-Type", "application/json")
if err := json.NewEncoder(w).Encode(map[string]any{"success": true, "errors": []any{}, "result": result}); err != nil {
t.Fatalf("Encode(response) error = %v", err)
}
}
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
const (
errConnectionNotConfigured = "尚未配置 Cloudflare 连接"
errConnectionSourceInvalid = "无效的 Cloudflare 连接来源"
errStandaloneInputRequired = "请填写 Cloudflare API Token"
errStandaloneInputInvalid = "配置的 Cloudflare API Token 无效"
errDNSAccountInvalid = "请选择有效的 Cloudflare DNS 账号"
errGroupNameRequired = "分组名称不能为空"
errGroupNodeSame = "主节点和备用节点不能相同"
errNodeInvalid = "请选择有效的边缘节点"
errNodeIPv4Required = "生效节点必须配置合法 IPv4"
errGroupDisabled = "指向分组已停用"
errMemberExists = "该域名已加入其他指向分组"
errMultipleARecords = "检测到 Cloudflare 中存在多条同名 A 记录,请先手动清理"
errSyncFailed = "Cloudflare DNS 同步失败"
errDeleteRemoteFailed = "删除 Cloudflare DNS 记录失败"
errTaskDispatchFailed = "无法投递 Cloudflare 同步任务"
)
@@ -0,0 +1,474 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"context"
"encoding/json"
"errors"
"net"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/credential"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/logger"
"gorm.io/gorm"
)
var clientFactory = func(token string) Client { return NewHTTPClient(token) }
// SetClientFactoryForTest replaces Cloudflare client construction for tests.
func SetClientFactoryForTest(factory func(string) Client) func() {
previous := clientFactory
clientFactory = factory
return func() { clientFactory = previous }
}
// GetConnection returns the global connection state without its token.
func GetConnection(ctx context.Context) (*ConnectionView, error) {
item, err := repository.GetCFConnection(ctx)
if errors.Is(err, gorm.ErrRecordNotFound) {
return &ConnectionView{}, nil
}
if err != nil {
return nil, err
}
return connectionView(item), nil
}
// SaveConnection stores a DNS-account or standalone Cloudflare credential source.
func SaveConnection(ctx context.Context, input ConnectionInput) (*ConnectionView, error) {
source := strings.TrimSpace(input.Source)
item := &model.CFConnection{Source: source}
switch source {
case model.CFConnectionSourceDNSAccount:
account, err := repository.GetDNSAccountByID(ctx, input.DNSAccountID)
if err != nil || !strings.EqualFold(strings.TrimSpace(account.Type), "cloudflare") {
return nil, errors.New(errDNSAccountInvalid)
}
item.DNSAccountID = &account.ID
case model.CFConnectionSourceStandalone:
token := strings.TrimSpace(input.APIToken)
if token == "" {
return nil, errors.New(errStandaloneInputRequired)
}
payload, err := json.Marshal(map[string]string{"api_token": token})
if err != nil {
return nil, errors.New(errStandaloneInputInvalid)
}
sealed, err := credential.Seal(string(payload))
if err != nil {
return nil, errors.New(errStandaloneInputInvalid)
}
item.Authorization = sealed
default:
return nil, errors.New(errConnectionSourceInvalid)
}
if err := repository.UpsertCFConnection(ctx, item); err != nil {
return nil, err
}
return connectionView(item), nil
}
// ClearConnection removes the configured Cloudflare credential.
func ClearConnection(ctx context.Context) error {
return repository.DeleteCFConnection(ctx)
}
// VerifyConnection verifies and marks the configured token ready.
func VerifyConnection(ctx context.Context) (*ConnectionView, error) {
item, err := repository.GetCFConnection(ctx)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New(errConnectionNotConfigured)
}
return nil, err
}
token, err := resolveToken(ctx, item)
if err != nil {
return nil, err
}
if err = clientFactory(token).VerifyToken(ctx); err != nil {
item.Status = model.CFConnectionStatusError
item.VerifiedAt = nil
if persistErr := repository.UpsertCFConnection(ctx, item); persistErr != nil {
logger.ErrorF(ctx, "[Cloudflare] persist failed verification status failed: error=%v", persistErr)
}
return nil, errors.New(errStandaloneInputInvalid)
}
now := time.Now()
item.Status = model.CFConnectionStatusReady
item.VerifiedAt = &now
if err = repository.UpsertCFConnection(ctx, item); err != nil {
return nil, err
}
return connectionView(item), nil
}
func connectionView(item *model.CFConnection) *ConnectionView {
return &ConnectionView{
Configured: true,
Ready: item.Status == model.CFConnectionStatusReady,
Source: item.Source, DNSAccountID: item.DNSAccountID,
Status: item.Status, VerifiedAt: item.VerifiedAt,
}
}
func resolveToken(ctx context.Context, item *model.CFConnection) (string, error) {
if item == nil {
return "", errors.New(errConnectionNotConfigured)
}
stored := item.Authorization
if item.Source == model.CFConnectionSourceDNSAccount {
if item.DNSAccountID == nil {
return "", errors.New(errDNSAccountInvalid)
}
account, err := repository.GetDNSAccountByID(ctx, *item.DNSAccountID)
if err != nil || !strings.EqualFold(strings.TrimSpace(account.Type), "cloudflare") {
return "", errors.New(errDNSAccountInvalid)
}
stored = account.Authorization
} else if item.Source != model.CFConnectionSourceStandalone {
return "", errors.New(errConnectionSourceInvalid)
}
opened, err := credential.Open(stored)
if err != nil {
return "", errors.New(errStandaloneInputInvalid)
}
var authorization map[string]string
if err = json.Unmarshal([]byte(opened), &authorization); err != nil || strings.TrimSpace(authorization["api_token"]) == "" {
return "", errors.New(errStandaloneInputInvalid)
}
return strings.TrimSpace(authorization["api_token"]), nil
}
// ListNodeOptions lists edge nodes selectable by pointing groups.
func ListNodeOptions(ctx context.Context) ([]NodeOption, error) {
nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil {
return nil, err
}
items := make([]NodeOption, 0, len(nodes))
for _, node := range nodes {
if node.NodeType != "edge_node" {
continue
}
items = append(items, nodeOption(&node))
}
return items, nil
}
// ListGroups returns pointing group summaries.
func ListGroups(ctx context.Context) ([]GroupItem, error) {
groups, err := repository.ListCFPointingGroups(ctx)
if err != nil {
return nil, err
}
items := make([]GroupItem, 0, len(groups))
for i := range groups {
item, buildErr := buildGroupItem(ctx, &groups[i])
if buildErr != nil {
return nil, buildErr
}
items = append(items, *item)
}
return items, nil
}
// CreateGroup creates a pointing group with its primary node active.
func CreateGroup(ctx context.Context, input GroupInput) (*GroupItem, error) {
group, err := groupFromInput(ctx, nil, input)
if err != nil {
return nil, err
}
if err = repository.CreateCFPointingGroup(ctx, group); err != nil {
return nil, err
}
return buildGroupItem(ctx, group)
}
// UpdateGroup updates a pointing group and queues reconciliation when enabled.
func UpdateGroup(ctx context.Context, id uint, input GroupInput) (*GroupItem, error) {
existing, err := repository.GetCFPointingGroup(ctx, id)
if err != nil {
return nil, err
}
group, err := groupFromInput(ctx, existing, input)
if err != nil {
return nil, err
}
if err = repository.SaveCFPointingGroup(ctx, group); err != nil {
return nil, err
}
if err = repository.MarkCFPointingGroupMembersPending(ctx, id); err != nil {
return nil, err
}
if group.Enabled {
if _, err = DispatchGroupSync(ctx, id, "cloudflare_group_update"); err != nil {
return nil, errors.New(errTaskDispatchFailed)
}
}
return buildGroupItem(ctx, group)
}
func groupFromInput(ctx context.Context, existing *model.CFPointingGroup, input GroupInput) (*model.CFPointingGroup, error) {
name := strings.TrimSpace(input.Name)
if name == "" {
return nil, errors.New(errGroupNameRequired)
}
if input.BackupNodeID != nil && *input.BackupNodeID == input.PrimaryNodeID {
return nil, errors.New(errGroupNodeSame)
}
primary, err := validEdgeNode(ctx, input.PrimaryNodeID, true)
if err != nil {
return nil, err
}
if input.BackupNodeID != nil {
if _, err = validEdgeNode(ctx, *input.BackupNodeID, false); err != nil {
return nil, err
}
}
if existing == nil {
existing = &model.CFPointingGroup{}
}
existing.Name = name
existing.PrimaryNodeID = primary.ID
existing.ActiveNodeID = primary.ID
existing.BackupNodeID = input.BackupNodeID
existing.DefaultProxied = input.DefaultProxied
existing.Enabled = input.Enabled
return existing, nil
}
func validEdgeNode(ctx context.Context, id uint, requireIPv4 bool) (*model.OpenFlareNode, error) {
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil || node.NodeType != "edge_node" {
return nil, errors.New(errNodeInvalid)
}
if requireIPv4 && net.ParseIP(strings.TrimSpace(node.IP)).To4() == nil {
return nil, errors.New(errNodeIPv4Required)
}
return node, nil
}
func buildGroupItem(ctx context.Context, group *model.CFPointingGroup) (*GroupItem, error) {
primary, err := repository.GetOpenFlareNodeByID(ctx, group.PrimaryNodeID)
if err != nil {
return nil, err
}
active, err := repository.GetOpenFlareNodeByID(ctx, group.ActiveNodeID)
if err != nil {
return nil, err
}
count, err := repository.CountCFPointingMembersByGroupID(ctx, group.ID)
if err != nil {
return nil, err
}
item := &GroupItem{ID: group.ID, Name: group.Name, PrimaryNode: nodeOption(primary), ActiveNode: nodeOption(active), DefaultProxied: group.DefaultProxied, Enabled: group.Enabled, MemberCount: count, CreatedAt: group.CreatedAt, UpdatedAt: group.UpdatedAt}
if group.BackupNodeID != nil {
backup, backupErr := repository.GetOpenFlareNodeByID(ctx, *group.BackupNodeID)
if backupErr != nil {
return nil, backupErr
}
option := nodeOption(backup)
item.BackupNode = &option
}
return item, nil
}
func nodeOption(node *model.OpenFlareNode) NodeOption {
return NodeOption{ID: node.ID, Name: node.Name, IP: node.IP}
}
// GetGroup returns a group and its members.
func GetGroup(ctx context.Context, id uint) (*GroupDetail, error) {
group, err := repository.GetCFPointingGroup(ctx, id)
if err != nil {
return nil, err
}
item, err := buildGroupItem(ctx, group)
if err != nil {
return nil, err
}
members, err := listMemberItems(ctx, id)
if err != nil {
return nil, err
}
return &GroupDetail{Group: *item, Members: members}, nil
}
// CreateMember adds a ZoneDomain and queues its first synchronization.
func CreateMember(ctx context.Context, groupID uint, input MemberCreateInput) (*MemberItem, error) {
group, err := repository.GetCFPointingGroup(ctx, groupID)
if err != nil {
return nil, err
}
domain, err := repository.GetZoneDomainByID(ctx, input.ZoneDomainID)
if err != nil {
return nil, err
}
if _, err = repository.GetCFPointingMemberByZoneDomainID(ctx, domain.ID); err == nil {
return nil, errors.New(errMemberExists)
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
proxied := group.DefaultProxied
if input.Proxied != nil {
proxied = *input.Proxied
}
member := &model.CFPointingMember{GroupID: groupID, ZoneDomainID: domain.ID, Proxied: proxied, SyncStatus: model.CFMemberSyncPending}
if err = repository.CreateCFPointingMember(ctx, member); err != nil {
return nil, err
}
if group.Enabled {
if _, err = DispatchMemberSync(ctx, member.ID, "cloudflare_member_create"); err != nil {
return nil, errors.New(errTaskDispatchFailed)
}
}
return memberItem(member, domain), nil
}
// UpdateMember updates orange-cloud state and queues reconciliation.
func UpdateMember(ctx context.Context, groupID, memberID uint, input MemberUpdateInput) (*MemberItem, error) {
member, err := repository.GetCFPointingMember(ctx, groupID, memberID)
if err != nil {
return nil, err
}
member.Proxied = input.Proxied
member.SyncStatus = model.CFMemberSyncPending
member.LastError = ""
if err = repository.SaveCFPointingMember(ctx, member); err != nil {
return nil, err
}
group, err := repository.GetCFPointingGroup(ctx, groupID)
if err != nil {
return nil, err
}
if group.Enabled {
if _, err = DispatchMemberSync(ctx, member.ID, "cloudflare_member_update"); err != nil {
return nil, errors.New(errTaskDispatchFailed)
}
}
domain, err := repository.GetZoneDomainByID(ctx, member.ZoneDomainID)
if err != nil {
return nil, err
}
return memberItem(member, domain), nil
}
// RemoveMember deletes the managed remote A record before removing local state.
func RemoveMember(ctx context.Context, groupID, memberID uint) error {
member, err := repository.GetCFPointingMember(ctx, groupID, memberID)
if err != nil {
return err
}
if err = DeleteManagedRecord(ctx, member.ID); err != nil {
return errors.New(errDeleteRemoteFailed)
}
return repository.DeleteCFPointingMember(ctx, member)
}
// DeleteGroup removes every managed remote A record and then local state.
func DeleteGroup(ctx context.Context, groupID uint) error {
if _, err := repository.GetCFPointingGroup(ctx, groupID); err != nil {
return err
}
members, err := repository.ListCFPointingMembersByGroupID(ctx, groupID)
if err != nil {
return err
}
for _, member := range members {
if err = DeleteManagedRecord(ctx, member.ID); err != nil {
return errors.New(errDeleteRemoteFailed)
}
}
return repository.DeleteCFPointingGroupAndMembers(ctx, groupID)
}
// ListAvailableDomains returns ZoneDomains not yet assigned to a group.
func ListAvailableDomains(ctx context.Context) ([]AvailableDomain, error) {
domains, err := repository.ListAvailableCFZoneDomains(ctx)
if err != nil {
return nil, err
}
zones, err := repository.ListZones(ctx)
if err != nil {
return nil, err
}
zoneRoots := make(map[uint]string, len(zones))
for i := range zones {
zoneRoots[zones[i].ID] = zones[i].Domain
}
items := make([]AvailableDomain, 0, len(domains))
for _, domain := range domains {
items = append(items, AvailableDomain{
ID: domain.ID,
ZoneID: domain.ZoneID,
Domain: domain.Domain,
ZoneDomain: zoneRoots[domain.ZoneID],
})
}
return items, nil
}
func listMemberItems(ctx context.Context, groupID uint) ([]MemberItem, error) {
members, err := repository.ListCFPointingMembersByGroupID(ctx, groupID)
if err != nil {
return nil, err
}
items := make([]MemberItem, 0, len(members))
for i := range members {
domain, domainErr := repository.GetZoneDomainByID(ctx, members[i].ZoneDomainID)
if domainErr != nil {
if errors.Is(domainErr, gorm.ErrRecordNotFound) {
logger.WarnF(ctx, "[Cloudflare] cleaning up orphaned pointing member: member_id=%d zone_domain_id=%d", members[i].ID, members[i].ZoneDomainID)
if delErr := repository.DeleteCFPointingMember(ctx, &members[i]); delErr != nil {
logger.ErrorF(ctx, "[Cloudflare] delete orphaned member failed: member_id=%d error=%v", members[i].ID, delErr)
}
continue
}
return nil, domainErr
}
items = append(items, *memberItem(&members[i], domain))
}
return items, nil
}
func memberItem(member *model.CFPointingMember, domain *model.ZoneDomain) *MemberItem {
return &MemberItem{ID: member.ID, GroupID: member.GroupID, ZoneDomainID: member.ZoneDomainID, Domain: domain.Domain, ZoneID: domain.ZoneID, Proxied: member.Proxied, DesiredIP: member.DesiredIP, SyncStatus: member.SyncStatus, LastError: member.LastError, SyncedAt: member.SyncedAt}
}
// GetOverview returns readiness and aggregate sync counts.
func GetOverview(ctx context.Context) (*Overview, error) {
connection, err := GetConnection(ctx)
if err != nil {
return nil, err
}
groups, err := repository.ListCFPointingGroups(ctx)
if err != nil {
return nil, err
}
overview := &Overview{Connection: *connection, GroupCount: len(groups)}
for _, group := range groups {
members, listErr := repository.ListCFPointingMembersByGroupID(ctx, group.ID)
if listErr != nil {
return nil, listErr
}
for _, member := range members {
overview.MemberCount++
switch member.SyncStatus {
case model.CFMemberSyncOK:
overview.OKCount++
case model.CFMemberSyncError:
overview.ErrorCount++
default:
overview.PendingCount++
}
}
}
return overview, nil
}
@@ -0,0 +1,159 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"context"
"errors"
"net"
"strings"
"sync"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/logger"
"gorm.io/gorm"
)
const (
memberLockStripeCount = 64
memberLastErrorColumn = "last_error"
memberSyncStatusColumn = "sync_status"
)
var memberLocks [memberLockStripeCount]sync.Mutex
// ReconcileMember makes one Cloudflare A record match the local desired state.
func ReconcileMember(ctx context.Context, memberID uint) error {
lock := &memberLocks[memberID%memberLockStripeCount]
lock.Lock()
defer lock.Unlock()
if err := repository.UpdateCFPointingMemberColumns(ctx, memberID, map[string]any{memberSyncStatusColumn: model.CFMemberSyncing, memberLastErrorColumn: ""}); err != nil {
return err
}
if err := reconcileMember(ctx, memberID); err != nil {
if updateErr := repository.UpdateCFPointingMemberColumns(ctx, memberID, map[string]any{memberSyncStatusColumn: model.CFMemberSyncError, memberLastErrorColumn: err.Error()}); updateErr != nil {
logger.ErrorF(ctx, "[Cloudflare] persist member sync error failed: member_id=%d error=%v", memberID, updateErr)
}
return err
}
return nil
}
func reconcileMember(ctx context.Context, memberID uint) error {
state, err := repository.GetCFPointingMemberContext(ctx, memberID)
if err != nil {
return err
}
if !state.Group.Enabled {
return errors.New(errGroupDisabled)
}
ip := strings.TrimSpace(state.Node.IP)
if net.ParseIP(ip).To4() == nil {
return errors.New(errNodeIPv4Required)
}
connection, err := repository.GetCFConnection(ctx)
if err != nil || connection.Status != model.CFConnectionStatusReady {
return errors.New(errConnectionNotConfigured)
}
token, err := resolveToken(ctx, connection)
if err != nil {
return err
}
client := clientFactory(token)
zoneID := state.Member.CFZoneID
if zoneID == "" {
zone, findErr := client.FindZone(ctx, state.Zone.Domain)
if findErr != nil {
return findErr
}
zoneID = zone.ID
}
input := RecordInput{Type: "A", Name: state.Domain.Domain, Content: ip, Proxied: state.Member.Proxied, TTL: 300}
if input.Proxied {
input.TTL = 1
}
recordID := state.Member.CFRecordID
if recordID != "" {
if _, getErr := client.GetRecord(ctx, zoneID, recordID); getErr == nil {
record, updateErr := client.UpdateARecord(ctx, zoneID, recordID, input)
if updateErr != nil {
return updateErr
}
return markMemberSynced(ctx, memberID, zoneID, record.ID, ip)
}
}
records, err := client.ListARecords(ctx, zoneID, state.Domain.Domain)
if err != nil {
return err
}
var record *DNSRecord
switch len(records) {
case 0:
record, err = client.CreateARecord(ctx, zoneID, input)
case 1:
record, err = client.UpdateARecord(ctx, zoneID, records[0].ID, input)
default:
return errors.New(errMultipleARecords)
}
if err != nil {
return err
}
return markMemberSynced(ctx, memberID, zoneID, record.ID, ip)
}
func markMemberSynced(ctx context.Context, memberID uint, zoneID, recordID, ip string) error {
now := time.Now()
return repository.UpdateCFPointingMemberColumns(ctx, memberID, map[string]any{
"cf_zone_id": zoneID, "cf_record_id": recordID, "desired_ip": ip,
memberSyncStatusColumn: model.CFMemberSyncOK, memberLastErrorColumn: "", "synced_at": &now,
})
}
// DeleteManagedRecord deletes the cached or uniquely discoverable A record.
func DeleteManagedRecord(ctx context.Context, memberID uint) error {
state, err := repository.GetCFPointingMemberContext(ctx, memberID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
}
connection, err := repository.GetCFConnection(ctx)
if err != nil {
return err
}
token, err := resolveToken(ctx, connection)
if err != nil {
return err
}
client := clientFactory(token)
zoneID := state.Member.CFZoneID
if zoneID == "" {
zone, findErr := client.FindZone(ctx, state.Zone.Domain)
if findErr != nil {
return findErr
}
zoneID = zone.ID
}
if state.Member.CFRecordID != "" {
if deleteErr := client.DeleteRecord(ctx, zoneID, state.Member.CFRecordID); deleteErr == nil {
return nil
}
}
records, err := client.ListARecords(ctx, zoneID, state.Domain.Domain)
if err != nil {
return err
}
if len(records) == 0 {
return nil
}
if len(records) > 1 {
return errors.New(errMultipleARecords)
}
return client.DeleteRecord(ctx, zoneID, records[0].ID)
}
@@ -0,0 +1,219 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"context"
"errors"
"testing"
"Wavelet/openflare/plugins/server/kernel/credential"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
type fakeClient struct {
records []DNSRecord
created *RecordInput
updated *RecordInput
deleted []string
deleteErrors map[string]error
}
func (client *fakeClient) VerifyToken(context.Context) error { return nil }
func (client *fakeClient) FindZone(context.Context, string) (*Zone, error) {
return &Zone{ID: "zone-1", Name: "example.com"}, nil
}
func (client *fakeClient) GetRecord(context.Context, string, string) (*DNSRecord, error) {
return nil, errors.New("not found")
}
func (client *fakeClient) ListARecords(context.Context, string, string) ([]DNSRecord, error) {
return client.records, nil
}
func (client *fakeClient) CreateARecord(_ context.Context, _ string, input RecordInput) (*DNSRecord, error) {
client.created = &input
return &DNSRecord{ID: "record-created", Name: input.Name, Content: input.Content, Proxied: input.Proxied}, nil
}
func (client *fakeClient) UpdateARecord(_ context.Context, _, id string, input RecordInput) (*DNSRecord, error) {
client.updated = &input
return &DNSRecord{ID: id, Name: input.Name, Content: input.Content, Proxied: input.Proxied}, nil
}
func (client *fakeClient) DeleteRecord(_ context.Context, _, recordID string) error {
client.deleted = append(client.deleted, recordID)
return client.deleteErrors[recordID]
}
func setupCloudflareLogicDB(t *testing.T) (context.Context, uint) {
t.Helper()
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
if err != nil {
t.Fatalf("gorm.Open() error = %v", err)
}
if err := conn.AutoMigrate(
&model.CFConnection{}, &model.CFPointingGroup{}, &model.CFPointingMember{},
&model.Zone{}, &model.ZoneDomain{}, &model.OpenFlareNode{}, &model.DNSAccount{},
); err != nil {
t.Fatalf("AutoMigrate() error = %v", err)
}
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
sealed, err := credential.Seal(`{"api_token":"test-token"}`)
if err != nil {
t.Fatalf("credential.Seal() error = %v", err)
}
if err := repository.UpsertCFConnection(ctx, &model.CFConnection{Source: model.CFConnectionSourceStandalone, Authorization: sealed, Status: model.CFConnectionStatusReady}); err != nil {
t.Fatalf("UpsertCFConnection() error = %v", err)
}
zone := model.Zone{Domain: "example.com"}
node := model.OpenFlareNode{Name: "edge", NodeID: "node-1", NodeType: "edge_node", IP: "203.0.113.10"}
if err := conn.Create(&zone).Error; err != nil {
t.Fatalf("Create(zone) error = %v", err)
}
if err := conn.Create(&node).Error; err != nil {
t.Fatalf("Create(node) error = %v", err)
}
domain := model.ZoneDomain{ZoneID: zone.ID, Domain: "api.example.com"}
if err := conn.Create(&domain).Error; err != nil {
t.Fatalf("Create(domain) error = %v", err)
}
group := model.CFPointingGroup{Name: "primary", PrimaryNodeID: node.ID, ActiveNodeID: node.ID, DefaultProxied: true, Enabled: true}
if err := conn.Create(&group).Error; err != nil {
t.Fatalf("Create(group) error = %v", err)
}
member := model.CFPointingMember{GroupID: group.ID, ZoneDomainID: domain.ID, Proxied: true, SyncStatus: model.CFMemberSyncPending}
if err := conn.Create(&member).Error; err != nil {
t.Fatalf("Create(member) error = %v", err)
}
return ctx, member.ID
}
func TestReconcileMemberCreatesMissingARecord(t *testing.T) {
ctx, memberID := setupCloudflareLogicDB(t)
fake := &fakeClient{}
restore := SetClientFactoryForTest(func(string) Client { return fake })
t.Cleanup(restore)
if err := ReconcileMember(ctx, memberID); err != nil {
t.Fatalf("ReconcileMember() error = %v", err)
}
if fake.created == nil || fake.created.Content != "203.0.113.10" || !fake.created.Proxied || fake.created.TTL != 1 {
t.Errorf("CreateARecord input = %+v", fake.created)
}
member, err := repository.GetCFPointingMemberByID(ctx, memberID)
if err != nil {
t.Fatalf("GetCFPointingMemberByID() error = %v", err)
}
if member.SyncStatus != model.CFMemberSyncOK || member.CFRecordID != "record-created" || member.DesiredIP != "203.0.113.10" {
t.Errorf("reconciled member = %+v", member)
}
}
func TestReconcileMemberRejectsMultipleSameNameARecords(t *testing.T) {
ctx, memberID := setupCloudflareLogicDB(t)
fake := &fakeClient{records: []DNSRecord{{ID: "one"}, {ID: "two"}}}
restore := SetClientFactoryForTest(func(string) Client { return fake })
t.Cleanup(restore)
if err := ReconcileMember(ctx, memberID); err == nil {
t.Fatal("ReconcileMember() error = nil, want duplicate record error")
}
member, err := repository.GetCFPointingMemberByID(ctx, memberID)
if err != nil {
t.Fatalf("GetCFPointingMemberByID() error = %v", err)
}
if member.SyncStatus != model.CFMemberSyncError || member.LastError == "" {
t.Errorf("failed member = %+v", member)
}
}
func TestCreateMemberCopiesGroupDefaultProxied(t *testing.T) {
ctx, existingMemberID := setupCloudflareLogicDB(t)
existing, err := repository.GetCFPointingMemberByID(ctx, existingMemberID)
if err != nil {
t.Fatalf("GetCFPointingMemberByID() error = %v", err)
}
group, err := repository.GetCFPointingGroup(ctx, existing.GroupID)
if err != nil {
t.Fatalf("GetCFPointingGroup() error = %v", err)
}
zone := model.Zone{Domain: "example.net"}
if err := db.DB(ctx).Create(&zone).Error; err != nil {
t.Fatalf("Create(zone) error = %v", err)
}
domain := model.ZoneDomain{ZoneID: zone.ID, Domain: "www.example.net"}
if err := db.DB(ctx).Create(&domain).Error; err != nil {
t.Fatalf("Create(domain) error = %v", err)
}
restore := SetDispatchTaskForTest(func(context.Context, string, []byte, string) (string, error) { return "task-1", nil })
t.Cleanup(restore)
member, err := CreateMember(ctx, group.ID, MemberCreateInput{ZoneDomainID: domain.ID})
if err != nil {
t.Fatalf("CreateMember() error = %v", err)
}
if !member.Proxied {
t.Error("CreateMember() proxied = false, want group default true")
}
}
func TestDeleteManagedRecordFallsBackWhenCachedRecordIsStale(t *testing.T) {
ctx, memberID := setupCloudflareLogicDB(t)
if err := repository.UpdateCFPointingMemberColumns(ctx, memberID, map[string]any{
"cf_zone_id": "zone-1",
"cf_record_id": "stale-record",
}); err != nil {
t.Fatalf("UpdateCFPointingMemberColumns() error = %v", err)
}
fake := &fakeClient{
records: []DNSRecord{{ID: "actual-record"}},
deleteErrors: map[string]error{"stale-record": errors.New("not found")},
}
restore := SetClientFactoryForTest(func(string) Client { return fake })
t.Cleanup(restore)
if err := DeleteManagedRecord(ctx, memberID); err != nil {
t.Fatalf("DeleteManagedRecord() error = %v", err)
}
if len(fake.deleted) != 2 || fake.deleted[0] != "stale-record" || fake.deleted[1] != "actual-record" {
t.Errorf("deleted record IDs = %v, want [stale-record actual-record]", fake.deleted)
}
}
func TestUpdateMemberDoesNotDispatchWhenGroupIsDisabled(t *testing.T) {
ctx, memberID := setupCloudflareLogicDB(t)
member, err := repository.GetCFPointingMemberByID(ctx, memberID)
if err != nil {
t.Fatalf("GetCFPointingMemberByID() error = %v", err)
}
group, err := repository.GetCFPointingGroup(ctx, member.GroupID)
if err != nil {
t.Fatalf("GetCFPointingGroup() error = %v", err)
}
group.Enabled = false
if err = repository.SaveCFPointingGroup(ctx, group); err != nil {
t.Fatalf("SaveCFPointingGroup() error = %v", err)
}
dispatchCount := 0
restore := SetDispatchTaskForTest(func(context.Context, string, []byte, string) (string, error) {
dispatchCount++
return "task-1", nil
})
t.Cleanup(restore)
updated, err := UpdateMember(ctx, group.ID, memberID, MemberUpdateInput{Proxied: false})
if err != nil {
t.Fatalf("UpdateMember() error = %v", err)
}
if updated.SyncStatus != model.CFMemberSyncPending {
t.Errorf("UpdateMember() sync status = %q, want %q", updated.SyncStatus, model.CFMemberSyncPending)
}
if dispatchCount != 0 {
t.Errorf("dispatch count = %d, want 0", dispatchCount)
}
}
@@ -0,0 +1,384 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"errors"
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
func abortLogic(c *gin.Context, err error) bool {
if err == nil {
return false
}
switch {
case errors.Is(err, gorm.ErrRecordNotFound):
response.AbortNotFound(c, "Cloudflare 资源不存在")
case err.Error() == errMemberExists:
response.AbortConflict(c, err.Error())
case err.Error() == errTaskDispatchFailed:
response.AbortInternal(c, err.Error())
default:
response.AbortBadRequest(c, err.Error())
}
return true
}
// GetConnectionHandler returns Cloudflare connection readiness.
// @Summary 获取 Cloudflare 连接
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=cloudflare.ConnectionView}
// @Router /api/v1/d/cloudflare/connection [get]
func GetConnectionHandler(c *gin.Context) {
item, err := GetConnection(c.Request.Context())
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// SaveConnectionHandler saves a Cloudflare credential source.
// @Summary 保存 Cloudflare 连接
// @Tags openflare-cloudflare
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body cloudflare.ConnectionInput true "连接参数"
// @Success 200 {object} response.Any{data=cloudflare.ConnectionView}
// @Failure 400 {object} response.Any
// @Router /api/v1/d/cloudflare/connection [put]
func SaveConnectionHandler(c *gin.Context) {
var input ConnectionInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := SaveConnection(c.Request.Context(), input)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// VerifyConnectionHandler verifies the configured Cloudflare token.
// @Summary 测试 Cloudflare 连接
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=cloudflare.ConnectionView}
// @Failure 400 {object} response.Any
// @Router /api/v1/d/cloudflare/connection/verify [post]
func VerifyConnectionHandler(c *gin.Context) {
item, err := VerifyConnection(c.Request.Context())
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// ClearConnectionHandler clears the Cloudflare credential source.
// @Summary 清除 Cloudflare 连接
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any
// @Router /api/v1/d/cloudflare/connection/clear [post]
func ClearConnectionHandler(c *gin.Context) {
if abortLogic(c, ClearConnection(c.Request.Context())) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// OverviewHandler returns Cloudflare pointing health.
// @Summary 获取 Cloudflare 指向总览
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=cloudflare.Overview}
// @Router /api/v1/d/cloudflare/overview [get]
func OverviewHandler(c *gin.Context) {
item, err := GetOverview(c.Request.Context())
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// ListGroupsHandler lists pointing groups.
// @Summary 获取 Cloudflare 指向分组
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]cloudflare.GroupItem}
// @Router /api/v1/d/cloudflare/groups [get]
func ListGroupsHandler(c *gin.Context) {
items, err := ListGroups(c.Request.Context())
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(items))
}
// CreateGroupHandler creates a pointing group.
// @Summary 创建 Cloudflare 指向分组
// @Tags openflare-cloudflare
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body cloudflare.GroupInput true "分组参数"
// @Success 200 {object} response.Any{data=cloudflare.GroupItem}
// @Failure 400 {object} response.Any
// @Router /api/v1/d/cloudflare/groups [post]
func CreateGroupHandler(c *gin.Context) {
var input GroupInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := CreateGroup(c.Request.Context(), input)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// GetGroupHandler returns one pointing group and its members.
// @Summary 获取 Cloudflare 指向分组详情
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Success 200 {object} response.Any{data=cloudflare.GroupDetail}
// @Failure 404 {object} response.Any
// @Router /api/v1/d/cloudflare/groups/{id} [get]
func GetGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
item, err := GetGroup(c.Request.Context(), id)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// UpdateGroupHandler updates a pointing group.
// @Summary 更新 Cloudflare 指向分组
// @Tags openflare-cloudflare
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Param body body cloudflare.GroupInput true "分组参数"
// @Success 200 {object} response.Any{data=cloudflare.GroupItem}
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/d/cloudflare/groups/{id}/update [post]
func UpdateGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input GroupInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := UpdateGroup(c.Request.Context(), id, input)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// DeleteGroupHandler deletes a pointing group and its managed remote A records.
// @Summary 删除 Cloudflare 指向分组
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Success 200 {object} response.Any
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/d/cloudflare/groups/{id}/delete [post]
func DeleteGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if abortLogic(c, DeleteGroup(c.Request.Context(), id)) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// SyncGroupHandler queues a full group synchronization.
// @Summary 同步 Cloudflare 指向分组
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Success 200 {object} response.Any{data=cloudflare.SyncReceipt}
// @Failure 500 {object} response.Any
// @Router /api/v1/d/cloudflare/groups/{id}/sync [post]
func SyncGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
taskID, err := DispatchGroupSync(c.Request.Context(), id, "cloudflare_manual_group_sync")
if err != nil {
response.AbortInternal(c, errTaskDispatchFailed)
return
}
c.JSON(http.StatusOK, response.OK(&SyncReceipt{TaskID: taskID}))
}
// ListMembersHandler lists group members.
// @Summary 获取 Cloudflare 指向分组成员
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Success 200 {object} response.Any{data=[]cloudflare.MemberItem}
// @Router /api/v1/d/cloudflare/groups/{id}/members [get]
func ListMembersHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
item, err := GetGroup(c.Request.Context(), id)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item.Members))
}
// CreateMemberHandler adds a ZoneDomain to a pointing group.
// @Summary 添加 Cloudflare 指向成员
// @Tags openflare-cloudflare
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Param body body cloudflare.MemberCreateInput true "成员参数"
// @Success 200 {object} response.Any{data=cloudflare.MemberItem}
// @Failure 400 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/d/cloudflare/groups/{id}/members [post]
func CreateMemberHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input MemberCreateInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := CreateMember(c.Request.Context(), id, input)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// UpdateMemberHandler updates a member's orange-cloud state.
// @Summary 更新 Cloudflare 指向成员
// @Tags openflare-cloudflare
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Param memberId path int true "成员 ID"
// @Param body body cloudflare.MemberUpdateInput true "成员参数"
// @Success 200 {object} response.Any{data=cloudflare.MemberItem}
// @Router /api/v1/d/cloudflare/groups/{id}/members/{memberId}/update [post]
func UpdateMemberHandler(c *gin.Context) {
groupID, memberID, ok := memberParams(c)
if !ok {
return
}
var input MemberUpdateInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := UpdateMember(c.Request.Context(), groupID, memberID, input)
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// RemoveMemberHandler removes a member and its managed remote A record.
// @Summary 移出 Cloudflare 指向成员
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Param memberId path int true "成员 ID"
// @Success 200 {object} response.Any
// @Router /api/v1/d/cloudflare/groups/{id}/members/{memberId}/remove [post]
func RemoveMemberHandler(c *gin.Context) {
groupID, memberID, ok := memberParams(c)
if !ok {
return
}
if abortLogic(c, RemoveMember(c.Request.Context(), groupID, memberID)) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// SyncMemberHandler queues one member synchronization.
// @Summary 同步 Cloudflare 指向成员
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Param id path int true "分组 ID"
// @Param memberId path int true "成员 ID"
// @Success 200 {object} response.Any{data=cloudflare.SyncReceipt}
// @Router /api/v1/d/cloudflare/groups/{id}/members/{memberId}/sync [post]
func SyncMemberHandler(c *gin.Context) {
_, memberID, ok := memberParams(c)
if !ok {
return
}
taskID, err := DispatchMemberSync(c.Request.Context(), memberID, "cloudflare_manual_member_sync")
if err != nil {
response.AbortInternal(c, errTaskDispatchFailed)
return
}
c.JSON(http.StatusOK, response.OK(&SyncReceipt{TaskID: taskID}))
}
// ListAvailableDomainsHandler lists ZoneDomains not assigned to another group.
// @Summary 获取可加入 Cloudflare 指向的域名
// @Tags openflare-cloudflare
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]cloudflare.AvailableDomain}
// @Router /api/v1/d/cloudflare/domains/available [get]
func ListAvailableDomainsHandler(c *gin.Context) {
items, err := ListAvailableDomains(c.Request.Context())
if abortLogic(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(items))
}
func memberParams(c *gin.Context) (uint, uint, bool) {
groupID, ok := apiutil.IDParam(c)
if !ok {
return 0, 0, false
}
memberID, ok := apiutil.NamedIDParam(c, "memberId")
return groupID, memberID, ok
}
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
)
func TestConnectionHandlersNeverReturnAPIToken(t *testing.T) {
setupCloudflareLogicDB(t)
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(response.ErrorHandlerMiddleware())
router.PUT("/connection", SaveConnectionHandler)
router.GET("/connection", GetConnectionHandler)
save := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodPut, "/connection", strings.NewReader(`{"source":"standalone","api_token":"top-secret-token"}`))
request.Header.Set("Content-Type", "application/json")
router.ServeHTTP(save, request)
if save.Code != http.StatusOK {
t.Fatalf("PUT /connection status = %d, body = %s", save.Code, save.Body.String())
}
if strings.Contains(save.Body.String(), "top-secret-token") || strings.Contains(save.Body.String(), "api_token") {
t.Fatalf("PUT /connection leaked token: %s", save.Body.String())
}
get := httptest.NewRecorder()
router.ServeHTTP(get, httptest.NewRequest(http.MethodGet, "/connection", nil))
if get.Code != http.StatusOK {
t.Fatalf("GET /connection status = %d, body = %s", get.Code, get.Body.String())
}
if strings.Contains(get.Body.String(), "top-secret-token") || strings.Contains(get.Body.String(), "api_token") {
t.Fatalf("GET /connection leaked token: %s", get.Body.String())
}
}
func TestGetGroupWithOrphanedMemberHealsAndSucceeds(t *testing.T) {
ctx, memberID := setupCloudflareLogicDB(t)
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(response.ErrorHandlerMiddleware())
router.GET("/groups/:id", GetGroupHandler)
member, err := repository.GetCFPointingMemberByID(ctx, memberID)
if err != nil {
t.Fatalf("GetCFPointingMemberByID() error = %v", err)
}
// Simulate orphaned member by deleting the ZoneDomain directly
if err := db.DB(ctx).Exec("DELETE FROM of_zone_domains WHERE id = ?", member.ZoneDomainID).Error; err != nil {
t.Fatalf("DELETE FROM of_zone_domains error = %v", err)
}
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "/groups/1", nil)
router.ServeHTTP(recorder, request)
if recorder.Code != http.StatusOK {
t.Fatalf("GET /groups/1 status = %d, body = %s", recorder.Code, recorder.Body.String())
}
// Verify the orphaned member has been removed
_, err = repository.GetCFPointingMemberByID(ctx, memberID)
if err == nil {
t.Errorf("GetCFPointingMemberByID() should return not found after healing")
}
}
@@ -0,0 +1,367 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/task"
"gorm.io/gorm"
)
const (
// SyncMemberTask is the Asynq task type for one Cloudflare member.
SyncMemberTask = "cloudflare:sync_member"
// SyncGroupTask is the Asynq task type for a Cloudflare group.
SyncGroupTask = "cloudflare:sync_group"
// SyncByNodeTask is the Asynq task type for members targeting one node.
SyncByNodeTask = "cloudflare:sync_by_node"
// TaskTypeSyncMember is the task metadata type for member synchronization.
TaskTypeSyncMember = "of_cloudflare_sync_member"
// TaskTypeSyncGroup is the task metadata type for group synchronization.
TaskTypeSyncGroup = "of_cloudflare_sync_group"
// TaskTypeSyncByNode is the task metadata type for node-triggered synchronization.
TaskTypeSyncByNode = "of_cloudflare_sync_by_node"
)
// SyncMemberMeta describes one-member reconciliation (admin-dispatchable).
var SyncMemberMeta = task.TaskMeta{
Type: TaskTypeSyncMember,
AsynqTask: SyncMemberTask,
Name: "Cloudflare 域名同步",
Description: "同步单个域名的 Cloudflare A 记录",
SupportsTime: false,
MaxRetry: 3,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.TaskParam{
{
Name: "member_id",
Label: "成员 ID",
Type: "number",
Required: true,
Placeholder: "请输入 Cloudflare 指向成员 ID",
Description: "of_cf_pointing_members 表中的成员主键 ID",
},
},
}
// SyncGroupMeta describes group reconciliation (admin-dispatchable).
var SyncGroupMeta = task.TaskMeta{
Type: TaskTypeSyncGroup,
AsynqTask: SyncGroupTask,
Name: "Cloudflare 分组同步",
Description: "同步指向分组内全部域名",
SupportsTime: false,
MaxRetry: 2,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.TaskParam{
{
Name: "group_id",
Label: "分组 ID",
Type: "number",
Required: true,
Placeholder: "请输入 Cloudflare 指向分组 ID",
Description: "of_cf_pointing_groups 表中的分组主键 ID",
},
},
}
// SyncByNodeMeta describes node-triggered reconciliation (internal only).
var SyncByNodeMeta = task.TaskMeta{
Type: TaskTypeSyncByNode,
AsynqTask: SyncByNodeTask,
Name: "Cloudflare 节点同步",
Description: "同步当前指向指定节点的全部域名",
SupportsTime: false,
MaxRetry: 2,
Queue: task.QueueDefault,
Retryable: true,
InternalOnly: true,
}
// SyncMemberPayload identifies one member.
type SyncMemberPayload struct {
MemberID uint `json:"member_id"`
}
// SyncGroupPayload identifies one group.
type SyncGroupPayload struct {
GroupID uint `json:"group_id"`
}
// SyncByNodePayload identifies one active node.
type SyncByNodePayload struct {
NodeID uint `json:"node_id"`
}
var dispatchTaskFn = task.DispatchTask
// SetDispatchTaskForTest replaces task dispatch for tests.
func SetDispatchTaskForTest(fn func(context.Context, string, []byte, string) (string, error)) func() {
previous := dispatchTaskFn
dispatchTaskFn = fn
return func() { dispatchTaskFn = previous }
}
// DispatchMemberSync queues one member reconciliation.
func DispatchMemberSync(ctx context.Context, memberID uint, triggeredBy string) (string, error) {
return dispatch(ctx, TaskTypeSyncMember, SyncMemberPayload{MemberID: memberID}, triggeredBy)
}
// DispatchGroupSync queues a group reconciliation.
func DispatchGroupSync(ctx context.Context, groupID uint, triggeredBy string) (string, error) {
return dispatch(ctx, TaskTypeSyncGroup, SyncGroupPayload{GroupID: groupID}, triggeredBy)
}
// DispatchNodeSync queues reconciliation for members targeting a node.
func DispatchNodeSync(ctx context.Context, nodeID uint, triggeredBy string) (string, error) {
return dispatch(ctx, TaskTypeSyncByNode, SyncByNodePayload{NodeID: nodeID}, triggeredBy)
}
func dispatch(ctx context.Context, taskType string, payload any, triggeredBy string) (string, error) {
encoded, err := json.Marshal(payload)
if err != nil {
return "", err
}
return dispatchTaskFn(ctx, taskType, encoded, triggeredBy)
}
// SyncMemberTaskHandler reconciles one member.
type SyncMemberTaskHandler struct{}
// ValidatePayload validates a one-member task payload.
func (handler *SyncMemberTaskHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New("任务参数不能为空")
}
var input SyncMemberPayload
if err := decodePayload(payload, &input); err != nil {
return nil, fmt.Errorf("无效的 Cloudflare 成员同步参数: %w", err)
}
if input.MemberID == 0 {
return nil, errors.New("成员 ID 不能为空或零")
}
return json.Marshal(input)
}
// Execute reconciles one member.
func (handler *SyncMemberTaskHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
normalized, err := handler.ValidatePayload(payload)
if err != nil {
return nil, task.PermanentError(err.Error())
}
var input SyncMemberPayload
_ = json.Unmarshal(normalized, &input)
state, loadErr := repository.GetCFPointingMemberContext(ctx, input.MemberID)
if loadErr != nil {
task.AppendLog(ctx, "加载成员上下文失败: member_id=%d error=%v", input.MemberID, loadErr)
} else {
task.AppendLog(ctx,
"开始域名同步: domain=%s zone=%s group=%s(#%d) node=%s(%s) proxied=%v member_id=%d",
state.Domain.Domain,
state.Zone.Domain,
state.Group.Name,
state.Group.ID,
state.Node.Name,
strings.TrimSpace(state.Node.IP),
state.Member.Proxied,
input.MemberID,
)
}
if err = ReconcileMember(ctx, input.MemberID); err != nil {
if state != nil {
task.AppendLog(ctx, "域名同步失败: domain=%s member_id=%d error=%v",
state.Domain.Domain, input.MemberID, err)
} else {
task.AppendLog(ctx, "域名同步失败: member_id=%d error=%v", input.MemberID, err)
}
return nil, fmt.Errorf("%s: %w", errSyncFailed, err)
}
message := "Cloudflare 域名同步成功"
if state != nil {
ip := strings.TrimSpace(state.Node.IP)
message = fmt.Sprintf("Cloudflare 域名同步成功: %s → %s (proxied=%v)",
state.Domain.Domain, ip, state.Member.Proxied)
task.AppendLog(ctx,
"域名同步成功: domain=%s desired_ip=%s proxied=%v group=%s node=%s",
state.Domain.Domain, ip, state.Member.Proxied, state.Group.Name, state.Node.Name,
)
} else {
task.AppendLog(ctx, "域名同步成功: member_id=%d", input.MemberID)
}
return &task.TaskResult{Message: message}, nil
}
// SyncGroupTaskHandler reconciles every member in a group.
type SyncGroupTaskHandler struct{}
// ValidatePayload validates a group task payload.
func (handler *SyncGroupTaskHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New("任务参数不能为空")
}
var input SyncGroupPayload
if err := decodePayload(payload, &input); err != nil {
return nil, fmt.Errorf("无效的 Cloudflare 分组同步参数: %w", err)
}
if input.GroupID == 0 {
return nil, errors.New("分组 ID 不能为空或零")
}
return json.Marshal(input)
}
// Execute reconciles every member in a group.
func (handler *SyncGroupTaskHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
normalized, err := handler.ValidatePayload(payload)
if err != nil {
return nil, task.PermanentError(err.Error())
}
var input SyncGroupPayload
if err = json.Unmarshal(normalized, &input); err != nil {
return nil, task.PermanentError(err.Error())
}
scopeName := fmt.Sprintf("#%d", input.GroupID)
activeNode := ""
if group, groupErr := repository.GetCFPointingGroup(ctx, input.GroupID); groupErr != nil {
task.AppendLog(ctx, "加载分组失败: group_id=%d error=%v", input.GroupID, groupErr)
} else {
scopeName = group.Name
if node, nodeErr := repository.GetOpenFlareNodeByID(ctx, group.ActiveNodeID); nodeErr != nil {
task.AppendLog(ctx, "加载生效节点失败: group=%s active_node_id=%d error=%v",
group.Name, group.ActiveNodeID, nodeErr)
} else {
activeNode = fmt.Sprintf("%s(%s)", node.Name, strings.TrimSpace(node.IP))
}
task.AppendLog(ctx,
"准备分组同步: group=%s id=%d enabled=%v active_node=%s default_proxied=%v",
group.Name, group.ID, group.Enabled, activeNode, group.DefaultProxied,
)
}
members, err := repository.ListCFPointingMembersByGroupID(ctx, input.GroupID)
return executeBatchSync(ctx, members, err, "分组", scopeName, input.GroupID, activeNode)
}
// SyncByNodeTaskHandler reconciles every member targeting a node.
type SyncByNodeTaskHandler struct{}
// ValidatePayload validates a node task payload.
func (handler *SyncByNodeTaskHandler) ValidatePayload(payload []byte) ([]byte, error) {
var input SyncByNodePayload
if err := decodePayload(payload, &input); err != nil || input.NodeID == 0 {
return nil, errors.New("无效的 Cloudflare 节点同步参数")
}
return json.Marshal(input)
}
// Execute reconciles every member targeting a node.
func (handler *SyncByNodeTaskHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
normalized, err := handler.ValidatePayload(payload)
if err != nil {
return nil, task.PermanentError(err.Error())
}
var input SyncByNodePayload
if err = json.Unmarshal(normalized, &input); err != nil {
return nil, task.PermanentError(err.Error())
}
scopeName := fmt.Sprintf("#%d", input.NodeID)
activeNode := ""
if node, nodeErr := repository.GetOpenFlareNodeByID(ctx, input.NodeID); nodeErr != nil {
task.AppendLog(ctx, "加载节点失败: node_id=%d error=%v", input.NodeID, nodeErr)
} else {
scopeName = node.Name
activeNode = fmt.Sprintf("%s(%s)", node.Name, strings.TrimSpace(node.IP))
task.AppendLog(ctx, "准备节点同步: node=%s id=%d ip=%s",
node.Name, node.ID, strings.TrimSpace(node.IP))
}
members, err := repository.ListCFPointingMembersByActiveNodeID(ctx, input.NodeID)
return executeBatchSync(ctx, members, err, "节点", scopeName, input.NodeID, activeNode)
}
func executeBatchSync(
ctx context.Context,
members []model.CFPointingMember,
listErr error,
scope, scopeName string,
scopeID uint,
activeNode string,
) (*task.TaskResult, error) {
if listErr != nil {
task.AppendLog(ctx, "列出%s成员失败: name=%s id=%d error=%v",
scope, scopeName, scopeID, listErr)
return nil, listErr
}
task.AppendLog(ctx, "开始%s同步: name=%s id=%d active_node=%s 域名数=%d",
scope, scopeName, scopeID, activeNode, len(members))
if len(members) == 0 {
message := fmt.Sprintf("Cloudflare %s同步完成: %s 无域名成员", scope, scopeName)
task.AppendLog(ctx, "%s", message)
return &task.TaskResult{Message: message}, nil
}
syncedCount := 0
for index, member := range members {
domainName := fmt.Sprintf("zone_domain_id=%d", member.ZoneDomainID)
if domain, domainErr := repository.GetZoneDomainByID(ctx, member.ZoneDomainID); domainErr == nil {
domainName = domain.Domain
} else if errors.Is(domainErr, gorm.ErrRecordNotFound) {
task.AppendLog(ctx, "[%d/%d] 域名记录已不存在,清理孤立成员: member_id=%d zone_domain_id=%d",
index+1, len(members), member.ID, member.ZoneDomainID)
if delErr := repository.DeleteCFPointingMember(ctx, &member); delErr != nil {
task.AppendLog(ctx, "[%d/%d] 清理孤立成员失败: member_id=%d error=%v",
index+1, len(members), member.ID, delErr)
}
continue
}
task.AppendLog(ctx, "[%d/%d] 同步域名 %s (member_id=%d proxied=%v)",
index+1, len(members), domainName, member.ID, member.Proxied)
if err := ReconcileMember(ctx, member.ID); err != nil {
task.AppendLog(ctx, "[%d/%d] 失败: domain=%s member_id=%d error=%v",
index+1, len(members), domainName, member.ID, err)
return nil, err
}
syncedCount++
task.AppendLog(ctx, "[%d/%d] 成功: domain=%s", index+1, len(members), domainName)
}
message := fmt.Sprintf("Cloudflare %s同步完成: %s 共 %d 个域名", scope, scopeName, syncedCount)
if activeNode != "" {
message = fmt.Sprintf("Cloudflare %s同步完成: %s → %s,共 %d 个域名",
scope, scopeName, activeNode, syncedCount)
}
task.AppendLog(ctx, "%s", message)
return &task.TaskResult{Message: message}, nil
}
func decodePayload(payload []byte, target any) error {
decoder := json.NewDecoder(bytes.NewReader(payload))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
return err
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return errors.New("unexpected trailing JSON value")
}
return nil
}
@@ -0,0 +1,19 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import "testing"
func TestSyncMemberTaskHandlerRejectsTrailingJSONContent(t *testing.T) {
handler := &SyncMemberTaskHandler{}
for _, payload := range [][]byte{
[]byte(`{"member_id":7} {}`),
[]byte(`{"member_id":7}}`),
[]byte(`{"member_id":7}]`),
} {
if normalized, err := handler.ValidatePayload(payload); err == nil {
t.Errorf("ValidatePayload(%s) = %s, nil; want non-nil error", payload, normalized)
}
}
}
@@ -0,0 +1,107 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cloudflare
import "time"
// ConnectionInput configures the global Cloudflare credential source.
type ConnectionInput struct {
Source string `json:"source"`
DNSAccountID uint `json:"dns_account_id"`
APIToken string `json:"api_token"`
}
// ConnectionView exposes connection state without credentials.
type ConnectionView struct {
Configured bool `json:"configured"`
Ready bool `json:"ready"`
Source string `json:"source"`
DNSAccountID *uint `json:"dns_account_id"`
Status string `json:"status"`
VerifiedAt *time.Time `json:"verified_at"`
}
// NodeOption is a selectable edge node.
type NodeOption struct {
ID uint `json:"id"`
Name string `json:"name"`
IP string `json:"ip"`
}
// GroupInput creates or updates a pointing group.
type GroupInput struct {
Name string `json:"name"`
PrimaryNodeID uint `json:"primary_node_id"`
BackupNodeID *uint `json:"backup_node_id"`
DefaultProxied bool `json:"default_proxied"`
Enabled bool `json:"enabled"`
}
// GroupItem is the admin-facing pointing group summary.
type GroupItem struct {
ID uint `json:"id"`
Name string `json:"name"`
PrimaryNode NodeOption `json:"primary_node"`
BackupNode *NodeOption `json:"backup_node"`
ActiveNode NodeOption `json:"active_node"`
DefaultProxied bool `json:"default_proxied"`
Enabled bool `json:"enabled"`
MemberCount int64 `json:"member_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// MemberCreateInput adds a ZoneDomain to a group. Nil Proxied copies the group default.
type MemberCreateInput struct {
ZoneDomainID uint `json:"zone_domain_id"`
Proxied *bool `json:"proxied"`
}
// MemberUpdateInput updates the effective orange-cloud state.
type MemberUpdateInput struct {
Proxied bool `json:"proxied"`
}
// MemberItem is the admin-facing member state.
type MemberItem struct {
ID uint `json:"id"`
GroupID uint `json:"group_id"`
ZoneDomainID uint `json:"zone_domain_id"`
Domain string `json:"domain"`
ZoneID uint `json:"zone_id"`
Proxied bool `json:"proxied"`
DesiredIP string `json:"desired_ip"`
SyncStatus string `json:"sync_status"`
LastError string `json:"last_error"`
SyncedAt *time.Time `json:"synced_at"`
}
// GroupDetail combines a group with its members.
type GroupDetail struct {
Group GroupItem `json:"group"`
Members []MemberItem `json:"members"`
}
// AvailableDomain is a ZoneDomain eligible for pointing.
type AvailableDomain struct {
ID uint `json:"id"`
ZoneID uint `json:"zone_id"`
Domain string `json:"domain"`
ZoneDomain string `json:"zone_domain"`
}
// Overview summarizes Cloudflare pointing readiness and sync health.
type Overview struct {
Connection ConnectionView `json:"connection"`
GroupCount int `json:"group_count"`
MemberCount int `json:"member_count"`
OKCount int `json:"ok_count"`
PendingCount int `json:"pending_count"`
ErrorCount int `json:"error_count"`
}
// SyncReceipt identifies a queued asynchronous synchronization.
type SyncReceipt struct {
TaskID string `json:"task_id"`
}
@@ -0,0 +1,33 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dashboard
import (
"sync"
"time"
)
const overviewCacheTTL = 30 * time.Second
var overviewCache struct {
mu sync.Mutex
payload *OverviewPayload
expiresAt time.Time
}
func getCachedOverview() (*OverviewPayload, bool) {
overviewCache.mu.Lock()
defer overviewCache.mu.Unlock()
if overviewCache.payload == nil || time.Now().After(overviewCache.expiresAt) {
return nil, false
}
return overviewCache.payload, true
}
func setCachedOverview(payload *OverviewPayload) {
overviewCache.mu.Lock()
defer overviewCache.mu.Unlock()
overviewCache.payload = payload
overviewCache.expiresAt = time.Now().Add(overviewCacheTTL)
}
@@ -0,0 +1,63 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dashboard provides helper utilities for dashboard API handlers.
package dashboard
import (
"strings"
"time"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/kernel/model"
)
const (
nodeStatusOnline = "online"
nodeStatusOffline = "offline"
nodeStatusPending = "pending"
dashboardDistributionLimit = 8
dashboardOverviewSnapshotLimit = 500
highCPUUsagePercentThreshold = 80
highMemoryUsagePercentThreshold = 85
highStorageUsagePercentThreshold = 85
)
func computeNodeStatus(node *model.OpenFlareNode) string {
if node == nil {
return nodeStatusOffline
}
if node.LastSeenAt == nil || node.LastSeenAt.IsZero() {
return nodeStatusPending
}
// 默认离线阈值 60 秒(与 node_offline_threshold 默认一致)
threshold := 60 * time.Second
if time.Since(*node.LastSeenAt) > threshold {
return nodeStatusOffline
}
return nodeStatusOnline
}
func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
if node == nil {
return time.Time{}
}
nodeType := strings.TrimSpace(node.NodeType)
if nodeType == "" {
nodeType = "edge_node"
}
if nodeType == "tunnel_relay" && ofws.IsRelayConnected(node.NodeID) {
return ofws.RelayWSConnectedLastSeenValue
}
if nodeType == "tunnel_client" && ofws.IsFlaredConnected(node.NodeID) {
return ofws.FlaredWSConnectedLastSeenValue
}
if ofws.IsAgentConnected(node.NodeID) {
return ofws.AgentWSConnectedLastSeenValue
}
if node.LastSeenAt == nil {
return time.Time{}
}
return *node.LastSeenAt
}
@@ -0,0 +1,411 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dashboard
import (
"context"
"sort"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/observability"
"Wavelet/openflare/plugins/server/kernel/model"
)
// Summary is the dashboard node summary section.
type Summary 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"`
}
// Traffic is the dashboard traffic section.
type Traffic 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"`
}
// Capacity is the dashboard capacity section.
type Capacity 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"`
}
// NodeHealth is a dashboard node health row.
type NodeHealth struct {
ID uint `json:"id"`
NodeID string `json:"node_id"`
Name string `json:"name"`
GeoName string `json:"geo_name"`
GeoLatitude *float64 `json:"geo_latitude"`
GeoLongitude *float64 `json:"geo_longitude"`
Status string `json:"status"`
OpenrestyStatus string `json:"openresty_status"`
CurrentVersion string `json:"current_version"`
LastSeenAt any `json:"last_seen_at"`
ActiveEventCount int `json:"active_event_count"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsagePercent float64 `json:"memory_usage_percent"`
StorageUsagePercent float64 `json:"storage_usage_percent"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
}
// OverviewView is the expanded dashboard overview payload.
type OverviewView struct {
GeneratedAt time.Time `json:"generated_at"`
Summary Summary `json:"summary"`
Traffic Traffic `json:"traffic"`
Capacity Capacity `json:"capacity"`
Distributions observability.TrafficDistributions `json:"distributions"`
Trends observability.NodeTrends `json:"trends"`
Nodes []NodeHealth `json:"nodes"`
}
// OverviewPayload is the compact legacy dashboard overview response.
type OverviewPayload struct {
GeneratedAt any `json:"generated_at"`
Summary Summary `json:"summary"`
Traffic Traffic `json:"traffic"`
Capacity Capacity `json:"capacity"`
Distributions distributionsPayload `json:"distributions"`
Trends trendsPayload `json:"trends"`
Nodes [][]any `json:"nodes"`
}
type distributionsPayload struct {
StatusCodes [][]any `json:"status_codes"`
TopDomains [][]any `json:"top_domains"`
SourceCountries [][]any `json:"source_countries"`
}
type trendsPayload struct {
Traffic24h [][]any `json:"traffic_24h"`
Capacity24h [][]any `json:"capacity_24h"`
Network24h [][]any `json:"network_24h"`
DiskIO24h [][]any `json:"disk_io_24h"`
}
// GetOverview aggregates dashboard overview data from nodes and observability tables.
func GetOverview(ctx context.Context) (*OverviewPayload, error) {
if payload, ok := getCachedOverview(); ok {
return payload, nil
}
view, err := buildOverviewView(ctx)
if err != nil {
return nil, err
}
payload := compressOverview(view)
setCachedOverview(payload)
return payload, nil
}
func buildOverviewView(ctx context.Context) (*OverviewView, error) {
now := time.Now()
since := now.Add(-24 * time.Hour)
nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil {
return nil, err
}
// Latest-per-node health: dedicated LIMIT 1 BY queries (not a global raw LIMIT).
latestSnapshotRows, err := repository.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", since)
if err != nil {
return nil, err
}
snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, "", since, dashboardOverviewSnapshotLimit)
if err != nil {
return nil, err
}
accessLogRegions, err := repository.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit)
if err != nil {
return nil, err
}
activeEvents, err := repository.ListOpenFlareActiveHealthEvents(ctx)
if err != nil {
return nil, err
}
// L1 business: trends + distributions + totals from access logs only.
view := &OverviewView{
GeneratedAt: now,
Nodes: make([]NodeHealth, 0, len(nodes)),
Distributions: observability.BuildTrafficDistributionsFromAccessLogs(
ctx, since, now, dashboardDistributionLimit, accessLogRegions,
),
Trends: observability.BuildNodeTrends(ctx, now, "", snapshots),
}
// Global traffic summary uses true window uniqExact for UV (not sum of hourly uniques).
if summary, sumErr := repository.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{
Since: since,
Until: now,
}); sumErr == nil {
view.Traffic.RequestCount = summary.RequestCount
view.Traffic.ErrorCount = summary.ErrorCount
view.Traffic.UniqueVisitors = summary.UniqueIPCount
view.Traffic.ReportedNodes = int(summary.NodeCount)
if summary.RequestCount > 0 {
// Average QPS over the 24h window.
view.Traffic.EstimatedQPS = float64(summary.RequestCount) / (24 * 3600)
}
} else {
// Fallback: sum hourly request/error buckets only (UV left from summary path).
applyTrafficTotalsFromTrend(&view.Traffic, view.Trends.Traffic24h)
}
nodeTraffic := map[string]model.OpenFlareAccessLogNodeAggregate{}
if aggregates, aggErr := repository.NodeAggregatesOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{
Since: since,
Until: now,
}); aggErr == nil {
for _, row := range aggregates {
nodeTraffic[row.NodeID] = row
}
}
var cpuNodeCount int
var memoryNodeCount int
latestSnapshots := observability.LatestMetricSnapshotsByNode(latestSnapshotRows)
activeEventsByNode := observability.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 == "unhealthy" {
view.Summary.UnhealthyNodes++
}
latestSnapshot := latestSnapshots[node.NodeID]
nodeActiveEvents := activeEventsByNode[node.NodeID]
nodeHealth := NodeHealth{
ID: node.ID,
NodeID: node.NodeID,
Name: node.Name,
GeoName: node.GeoName,
GeoLatitude: node.GeoLatitude,
GeoLongitude: node.GeoLongitude,
Status: computedStatus,
OpenrestyStatus: node.OpenrestyStatus,
CurrentVersion: node.CurrentVersion,
LastSeenAt: nodeViewLastSeenAt(&node),
ActiveEventCount: len(nodeActiveEvents),
}
cpuNodeCount, memoryNodeCount = applyNodeSnapshotMetrics(&nodeHealth, latestSnapshot, view, cpuNodeCount, memoryNodeCount)
if agg, ok := nodeTraffic[node.NodeID]; ok {
nodeHealth.RequestCount = agg.RequestCount
nodeHealth.ErrorCount = agg.ErrorCount
nodeHealth.UniqueVisitorCount = agg.UniqueIPCount
}
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 applyNodeSnapshotMetrics(nodeHealth *NodeHealth, snapshot *model.OpenFlareMetricSnapshot, view *OverviewView, cpuNodeCount, memoryNodeCount int) (int, int) {
if snapshot == nil {
return cpuNodeCount, memoryNodeCount
}
nodeHealth.CPUUsagePercent = snapshot.CPUUsagePercent
nodeHealth.MemoryUsagePercent = observability.Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes)
nodeHealth.StorageUsagePercent = observability.Percentage(snapshot.StorageUsedBytes, snapshot.StorageTotalBytes)
if snapshot.CPUUsagePercent > 0 {
view.Capacity.AverageCPUUsagePercent += snapshot.CPUUsagePercent
cpuNodeCount++
}
if nodeHealth.MemoryUsagePercent > 0 {
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
memoryNodeCount++
}
if snapshot.CPUUsagePercent >= highCPUUsagePercentThreshold {
view.Capacity.HighCPUNodes++
}
if nodeHealth.MemoryUsagePercent >= highMemoryUsagePercentThreshold {
view.Capacity.HighMemoryNodes++
}
if nodeHealth.StorageUsagePercent >= highStorageUsagePercentThreshold {
view.Capacity.HighStorageNodes++
}
return cpuNodeCount, memoryNodeCount
}
func applyTrafficTotalsFromTrend(traffic *Traffic, points []observability.TrafficTrendPoint) {
if traffic == nil {
return
}
traffic.RequestCount = 0
traffic.ErrorCount = 0
// Do not sum hourly unique visitors — that overcounts. UV must come from TrafficSummary.
for _, point := range points {
traffic.RequestCount += point.RequestCount
traffic.ErrorCount += point.ErrorCount
}
if traffic.RequestCount > 0 && traffic.ReportedNodes == 0 {
traffic.ReportedNodes = 1
}
}
func compressOverview(view *OverviewView) *OverviewPayload {
if view == nil {
return &OverviewPayload{
Distributions: distributionsPayload{
StatusCodes: [][]any{},
TopDomains: [][]any{},
SourceCountries: [][]any{},
},
Trends: trendsPayload{
Traffic24h: [][]any{},
Capacity24h: [][]any{},
Network24h: [][]any{},
DiskIO24h: [][]any{},
},
Nodes: [][]any{},
}
}
return &OverviewPayload{
GeneratedAt: view.GeneratedAt,
Summary: view.Summary,
Traffic: view.Traffic,
Capacity: view.Capacity,
Distributions: distributionsPayload{
StatusCodes: compressDistributionItems(view.Distributions.StatusCodes),
TopDomains: compressDistributionItems(view.Distributions.TopDomains),
SourceCountries: compressDistributionItems(view.Distributions.SourceCountries),
},
Trends: trendsPayload{
Traffic24h: compressTrafficTrendPoints(view.Trends.Traffic24h),
Capacity24h: compressCapacityTrendPoints(view.Trends.Capacity24h),
Network24h: compressNetworkTrendPoints(view.Trends.Network24h),
DiskIO24h: compressDiskIOTrendPoints(view.Trends.DiskIO24h),
},
Nodes: compressDashboardNodes(view.Nodes),
}
}
func compressDistributionItems(items []observability.DistributionItem) [][]any {
rows := make([][]any, 0, len(items))
for _, item := range items {
rows = append(rows, []any{item.Key, item.Value})
}
return rows
}
func compressTrafficTrendPoints(points []observability.TrafficTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
// Compact layout: [0] bucket, [1] request, [2] error, [3] uv, [4] 2xx, [5] 4xx, [6] 5xx
rows = append(rows, []any{
point.BucketStartedAt,
point.RequestCount,
point.ErrorCount,
point.UniqueVisitorCount,
point.Status2xxCount,
point.Status4xxCount,
point.Status5xxCount,
})
}
return rows
}
func compressCapacityTrendPoints(points []observability.CapacityTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
rows = append(rows, []any{
point.BucketStartedAt,
point.AverageCPUUsagePercent,
point.AverageMemoryUsagePercent,
point.ReportedNodes,
})
}
return rows
}
func compressNetworkTrendPoints(points []observability.NetworkTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
// Compact layout: [0] bucket, [1] bytes_received, [2] bytes_provided, [3] reported_nodes
rows = append(rows, []any{
point.BucketStartedAt,
point.BytesReceived,
point.BytesProvided,
point.ReportedNodes,
})
}
return rows
}
func compressDiskIOTrendPoints(points []observability.DiskIOTrendPoint) [][]any {
rows := make([][]any, 0, len(points))
for _, point := range points {
rows = append(rows, []any{
point.BucketStartedAt,
point.DiskReadBytes,
point.DiskWriteBytes,
point.ReportedNodes,
})
}
return rows
}
func compressDashboardNodes(nodes []NodeHealth) [][]any {
rows := make([][]any, 0, len(nodes))
for _, node := range nodes {
rows = append(rows, []any{
node.ID,
node.NodeID,
node.Name,
node.GeoName,
node.GeoLatitude,
node.GeoLongitude,
node.Status,
node.OpenrestyStatus,
node.CurrentVersion,
node.LastSeenAt,
node.ActiveEventCount,
node.CPUUsagePercent,
node.MemoryUsagePercent,
node.StorageUsagePercent,
node.RequestCount,
node.ErrorCount,
node.UniqueVisitorCount,
})
}
return rows
}
@@ -0,0 +1,190 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dashboard
import (
"context"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupDashboardTestDB(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{}))
db.SetDB(sqliteDB)
testhelper.SetupLogStoresForTest(t)
return func() {
db.SetDB(nil)
}
}
func TestGetOverviewStructure(t *testing.T) {
cleanup := setupDashboardTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
lastSeen := now.Add(-15 * time.Second) // within default 60s offline threshold
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-dashboard-1",
Name: "Edge 1",
IP: "10.0.0.1",
Status: "online",
OpenrestyStatus: "healthy",
CurrentVersion: "v1.0.0",
LastSeenAt: &lastSeen,
}).Error)
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-dashboard-2",
Name: "Edge 2",
IP: "10.0.0.2",
Status: "pending",
OpenrestyStatus: "unknown",
}).Error)
// Seed older + newer snapshots per node; health must use latest-per-node, not a global raw limit.
require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: "node-dashboard-1",
CapturedAt: now.Add(-2 * time.Hour),
CPUUsagePercent: 10,
MemoryUsedBytes: 1,
MemoryTotalBytes: 10,
}))
require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: "node-dashboard-1",
CapturedAt: now.Add(-time.Minute),
CPUUsagePercent: 55,
MemoryUsedBytes: 5,
MemoryTotalBytes: 10,
StorageUsedBytes: 2,
StorageTotalBytes: 10,
}))
// Business traffic from access logs (L1 authority): 12 requests, 1 server error, 4 unique IPs.
logs := make([]*model.OpenFlareAccessLog, 0, 12)
for i := 0; i < 11; i++ {
logs = append(logs, &model.OpenFlareAccessLog{
NodeID: "node-dashboard-1",
LoggedAt: now.Add(-time.Minute),
RemoteAddr: "10.0.0." + string(rune('1'+i%4)), // rough; fixed below
Host: "app.example.com",
Path: "/",
StatusCode: 200,
BytesSent: 100,
})
}
ips := []string{"10.0.0.10", "10.0.0.11", "10.0.0.12", "10.0.0.13"}
for i := 0; i < 11; i++ {
logs[i].RemoteAddr = ips[i%4]
}
logs = append(logs, &model.OpenFlareAccessLog{
NodeID: "node-dashboard-1",
LoggedAt: now.Add(-time.Minute),
RemoteAddr: ips[0],
Host: "app.example.com",
Path: "/err",
StatusCode: 502,
BytesSent: 10,
})
require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, logs))
overview, err := GetOverview(ctx)
require.NoError(t, err)
require.NotNil(t, overview)
assert.False(t, overview.GeneratedAt.(time.Time).IsZero())
assert.Equal(t, 2, overview.Summary.TotalNodes)
assert.Equal(t, 1, overview.Summary.OnlineNodes)
assert.Equal(t, 1, overview.Summary.PendingNodes)
assert.Equal(t, 0, overview.Summary.OfflineNodes)
assert.Equal(t, 0, overview.Summary.UnhealthyNodes)
assert.Equal(t, int64(12), overview.Traffic.RequestCount)
assert.Equal(t, int64(4), overview.Traffic.UniqueVisitors)
assert.Equal(t, int64(1), overview.Traffic.ErrorCount)
// QPS over 24h window
assert.InDelta(t, 12.0/(24*3600.0), overview.Traffic.EstimatedQPS, 0.0001)
assert.Equal(t, 1, overview.Traffic.ReportedNodes)
// Node-level traffic from access log aggregates
onlineNodeCheck := overview.Nodes
require.NotEmpty(t, onlineNodeCheck)
assert.InDelta(t, 55.0, overview.Capacity.AverageCPUUsagePercent, 1e-9)
assert.InDelta(t, 50.0, overview.Capacity.AverageMemoryUsagePercent, 1e-9)
assert.Equal(t, 0, overview.Capacity.HighCPUNodes)
assert.Equal(t, 0, overview.Capacity.HighMemoryNodes)
assert.Equal(t, 0, overview.Capacity.HighStorageNodes)
require.NotNil(t, overview.Distributions.StatusCodes)
require.NotNil(t, overview.Distributions.TopDomains)
require.NotNil(t, overview.Distributions.SourceCountries)
// Status/top domains come from access logs.
assert.NotEmpty(t, overview.Distributions.StatusCodes)
assert.NotEmpty(t, overview.Distributions.TopDomains)
assert.Empty(t, overview.Distributions.SourceCountries)
require.Len(t, overview.Trends.Traffic24h, 24)
require.Len(t, overview.Trends.Capacity24h, 24)
require.Len(t, overview.Trends.Network24h, 24)
require.Len(t, overview.Trends.DiskIO24h, 24)
for _, row := range overview.Trends.Traffic24h {
require.Len(t, row, 7)
}
for _, row := range overview.Trends.Capacity24h {
require.Len(t, row, 4)
}
for _, row := range overview.Trends.Network24h {
require.Len(t, row, 4)
}
for _, row := range overview.Trends.DiskIO24h {
require.Len(t, row, 4)
}
require.Len(t, overview.Nodes, 2)
for _, row := range overview.Nodes {
require.Len(t, row, 17)
}
nodeByID := make(map[string][]any, len(overview.Nodes))
for _, row := range overview.Nodes {
nodeByID[row[1].(string)] = row
}
onlineNode := nodeByID["node-dashboard-1"]
require.NotNil(t, onlineNode)
assert.Equal(t, "Edge 1", onlineNode[2])
assert.Equal(t, "online", onlineNode[6])
assert.Equal(t, "healthy", onlineNode[7])
// Latest-per-node health fields (indexes match compressDashboardNodes).
assert.InDelta(t, 55.0, onlineNode[11], 1e-9) // cpu_usage_percent from latest snapshot
assert.InDelta(t, 50.0, onlineNode[12], 1e-9) // memory_usage_percent
assert.Equal(t, int64(12), onlineNode[14]) // request_count from access logs
assert.Equal(t, int64(1), onlineNode[15]) // error_count
assert.Equal(t, int64(4), onlineNode[16]) // unique visitors
pendingNode := nodeByID["node-dashboard-2"]
require.NotNil(t, pendingNode)
assert.Equal(t, "Edge 2", pendingNode[2])
assert.Equal(t, "pending", pendingNode[6])
assert.Equal(t, "unknown", pendingNode[7])
assert.InDelta(t, 55.0, overview.Capacity.AverageCPUUsagePercent, 1e-9)
assert.Equal(t, 1, overview.Traffic.ReportedNodes)
}
@@ -0,0 +1,33 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package dashboard
import (
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// GetOverviewHandler 获取仪表盘概览数据。
// @Summary 获取仪表盘概览
// @Description 聚合节点与可观测性数据,返回 OpenFlare 控制台仪表盘概览,需要管理员权限
// @Tags openflare-dashboard
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=dashboard.OverviewPayload} "仪表盘概览"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/dashboard/overview [get]
func GetOverviewHandler(c *gin.Context) {
overview, err := GetOverview(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(overview))
}
@@ -0,0 +1,90 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package agent implements the OpenFlare agent protocol: node registration,
// heartbeat processing, access-log ingestion, and related middleware.
package agent
import (
"context"
"log/slog"
"net"
"strings"
"sync"
pkggeoip "Wavelet/openflare/share/geoip"
)
// 共享一个 GeoIP 服务实例:mmdb 打开(mmap + 解析元数据)成本不低,缺文件时还会
// 同步下载,绝不能每个上报批次重建。maxminddb.Reader 并发安全,无需额外加锁。
// 初始化失败不锁存:下一批上报会重试(与旧行为一致)。
var (
sharedAccessLogGeoMu sync.Mutex
sharedAccessLogGeoInstance pkggeoip.Service
)
func sharedAccessLogGeoService(ctx context.Context) pkggeoip.Service {
sharedAccessLogGeoMu.Lock()
defer sharedAccessLogGeoMu.Unlock()
if sharedAccessLogGeoInstance == nil {
service, err := pkggeoip.NewMaxMindGeoIPServiceWithContext(ctx, "", "")
if err != nil {
slog.WarnContext(ctx, "initialize access log geo service failed", "error", err)
return nil
}
sharedAccessLogGeoInstance = service
}
return sharedAccessLogGeoInstance
}
// resolveAccessLogRegion resolves the region name for an access-log remote address.
// mmdb Lookup 本身是内存映射 trie 查找(微秒级),无需再建应用层 IP 缓存。
// resolveAccessLogRegion resolves the region name for an access-log remote address.
// mmdb Lookup 本身是内存映射 trie 查找(微秒级),无需再建应用层 IP 缓存。
func resolveAccessLogRegion(ctx context.Context, rawIP string) string {
normalizedIP := normalizeAccessLogIP(rawIP)
if normalizedIP == "" {
return ""
}
service := sharedAccessLogGeoService(ctx)
if service == nil {
return ""
}
info, err := service.GetGeoInfo(net.ParseIP(normalizedIP))
if err != nil || info == nil {
return ""
}
region := strings.TrimSpace(info.Name)
if region == "" {
region = strings.TrimSpace(info.ISOCode)
}
return region
}
func normalizeAccessLogIP(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
if ip := net.ParseIP(trimmed); ip != nil {
return ip.String()
}
trimmed = strings.TrimPrefix(trimmed, "[")
trimmed = strings.TrimSuffix(trimmed, "]")
if ip := net.ParseIP(trimmed); ip != nil {
return ip.String()
}
host, _, err := net.SplitHostPort(strings.TrimSpace(raw))
if err != nil {
return ""
}
host = strings.TrimPrefix(host, "[")
host = strings.TrimSuffix(host, "]")
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
return ""
}
@@ -0,0 +1,161 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"sync"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"gorm.io/gorm"
)
const (
agentTokenPositiveCacheTTL = 2 * time.Minute
agentTokenNegativeCacheTTL = 10 * time.Minute
// ponytail: 上限仅防未授权口伪造 token 撑爆内存;打满后放弃缓存(回退 DB 查询),行为不变
maxAgentTokenNegativeCacheEntries = 10_000
)
type cachedAgentNode struct {
node *model.OpenFlareNode
expiresAt time.Time
}
type accessTokenAuthCache struct {
mu sync.RWMutex
positive map[string]cachedAgentNode
negative map[string]time.Time
now func() time.Time
loadNodeByToken func(context.Context, string) (*model.OpenFlareNode, error)
}
var tokenCache = newAccessTokenAuthCache()
func newAccessTokenAuthCache() *accessTokenAuthCache {
return &accessTokenAuthCache{
positive: make(map[string]cachedAgentNode),
negative: make(map[string]time.Time),
now: time.Now,
loadNodeByToken: repository.GetOpenFlareNodeByAccessToken,
}
}
func (c *accessTokenAuthCache) authenticate(ctx context.Context, token string) (*model.OpenFlareNode, 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(ctx, token)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
}
return nil, err
}
c.storeNode(token, node)
return cloneNode(node), nil
}
func (c *accessTokenAuthCache) getNode(token string, now time.Time) (*model.OpenFlareNode, bool) {
c.mu.RLock()
entry, ok := c.positive[token]
c.mu.RUnlock()
if !ok {
return nil, false
}
if now.After(entry.expiresAt) {
c.mu.Lock()
delete(c.positive, token)
c.mu.Unlock()
return nil, false
}
return cloneNode(entry.node), true
}
func (c *accessTokenAuthCache) isMissing(token string, now time.Time) bool {
c.mu.RLock()
expiresAt, ok := c.negative[token]
c.mu.RUnlock()
if !ok {
return false
}
if now.After(expiresAt) {
c.mu.Lock()
delete(c.negative, token)
c.mu.Unlock()
return false
}
return true
}
func (c *accessTokenAuthCache) storeNode(token string, node *model.OpenFlareNode) {
if token == "" || node == nil {
return
}
c.mu.Lock()
defer c.mu.Unlock()
delete(c.negative, token)
c.positive[token] = cachedAgentNode{
node: cloneNode(node),
expiresAt: c.now().Add(agentTokenPositiveCacheTTL),
}
}
func (c *accessTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
if token == "" {
return
}
c.mu.Lock()
defer c.mu.Unlock()
delete(c.positive, token)
if len(c.negative) >= maxAgentTokenNegativeCacheEntries {
c.evictExpiredMissingLocked(c.now())
if len(c.negative) >= maxAgentTokenNegativeCacheEntries {
return // 缓存满:放弃缓存该 token,认证仍走 DB,仅防内存无限增长
}
}
c.negative[token] = expiresAt
}
// evictExpiredMissingLocked 清理已过期的 negative 条目,须持写锁调用。
func (c *accessTokenAuthCache) evictExpiredMissingLocked(now time.Time) {
for token, expiresAt := range c.negative {
if now.After(expiresAt) {
delete(c.negative, token)
}
}
}
func (c *accessTokenAuthCache) reset() {
c.mu.Lock()
defer c.mu.Unlock()
c.positive = make(map[string]cachedAgentNode)
c.negative = make(map[string]time.Time)
}
// ResetAuthCacheForTest clears the in-memory access token cache for integration tests.
func ResetAuthCacheForTest() {
tokenCache.reset()
}
// AuthenticateAccessToken validates X-Agent-Token against of_nodes.access_token.
func AuthenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
token = strings.TrimSpace(token)
if token == "" {
return nil, errors.New(errMissingAgentToken)
}
return tokenCache.authenticate(ctx, token)
}
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"errors"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/pages"
openrestyrender "Wavelet/openflare/share/render/openresty"
"gorm.io/gorm"
)
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := repository.GetActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
return &ActiveConfigMeta{
Version: version.Version,
Checksum: version.Checksum,
}, nil
}
func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) {
version, err := repository.GetActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
var supportFiles []SupportFile
if strings.TrimSpace(version.SupportFilesJSON) != "" {
if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil {
return nil, err
}
}
// Main config version history is independent of Pages deployment history.
// Agents always receive pages routes bound to each project's current active
// deployment so config rollback never depends on pruned packages.
sourceJSON := version.SnapshotJSON
if rebound, rebindErr := pages.RebindSnapshotPagesToCurrentActive(ctx, version.SnapshotJSON); rebindErr != nil {
return nil, rebindErr
} else if strings.TrimSpace(rebound) != "" {
sourceJSON = rebound
}
return &ConfigResponse{
Version: version.Version,
Checksum: version.Checksum,
SourceConfigJSON: sourceJSON,
SupportFiles: sourceSupportFiles(supportFiles),
CreatedAt: version.CreatedAt,
}, nil
}
func sourceSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
result := make([]SupportFile, 0, len(files))
for _, file := range files {
if isRuntimeGeneratedSupportFile(file.Path) {
continue
}
result = append(result, file)
}
return result
}
func isRuntimeGeneratedSupportFile(path string) bool {
switch strings.TrimSpace(path) {
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName: // pow_config.json is legacy; waf_config.json is canonical
return true
default:
return false
}
}
func isActiveConfigNotFound(err error) bool {
return errors.Is(err, gorm.ErrRecordNotFound)
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"testing"
openrestyrender "Wavelet/openflare/share/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)
}
}
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
const (
errMissingAgentToken = "缺少 Agent Token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errInvalidAgentToken = "无权进行此操作,Agent Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errNodeMissingFromContext = "Node object missing from context"
errNoActiveConfig = "当前没有激活版本"
errNodeNotFound = "节点不存在"
errNodeIDRequired = "node_id 不能为空"
errVersionRequired = "version 不能为空"
errInvalidApplyResult = "result 仅支持 success、warning 或 failed"
errIPRequired = "ip 不能为空"
errIPInvalid = "ip 格式无效"
errAgentVersionRequired = "version 不能为空"
errNodeIDConflict = "节点标识生成冲突,请重试"
)
@@ -0,0 +1,308 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"crypto/rand"
"encoding/hex"
"net"
"strings"
"time"
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
)
const (
openrestyStatusHealthy = "healthy"
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
releaseChannelStable = "stable"
randomTokenBytes = 16
maxDatabaseTextLength = 16000
defaultAgentHeartbeatInterval = 3000 // 默认心跳间隔 3 秒(毫秒)
defaultAgentUpdateRepo = "Rain-kl/OpenFlare"
)
func newRandomToken() (string, error) {
buf := make([]byte, randomTokenBytes)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return hex.EncodeToString(buf), nil
}
func newServerNodeID() (string, error) {
token, err := newRandomToken()
if err != nil {
return "", err
}
return "node-" + token, nil
}
func normalizeOpenrestyStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case openrestyStatusHealthy:
return openrestyStatusHealthy
case openrestyStatusUnhealthy:
return openrestyStatusUnhealthy
default:
return openrestyStatusUnknown
}
}
func normalizeNodePayload(payload NodePayload) NodePayload {
payload.Name = strings.TrimSpace(payload.Name)
payload.IP = strings.TrimSpace(payload.IP)
payload.Version = strings.TrimSpace(payload.Version)
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
payload.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
// Align L2 edge_health with top-level status/message (PG is latest-state authority).
if payload.EdgeHealth != nil {
if s := strings.TrimSpace(payload.EdgeHealth.Status); s != "" {
if payload.OpenrestyStatus == "" || payload.OpenrestyStatus == openrestyStatusUnknown {
payload.OpenrestyStatus = normalizeOpenrestyStatus(s)
}
}
if m := strings.TrimSpace(payload.EdgeHealth.Message); m != "" && payload.OpenrestyMessage == "" {
payload.OpenrestyMessage = truncateForDatabase(m, maxDatabaseTextLength)
}
// CH series status must match the same authority as PG after normalize.
payload.EdgeHealth.Status = payload.OpenrestyStatus
payload.EdgeHealth.Message = payload.OpenrestyMessage
}
return payload
}
func validateNodePayload(payload NodePayload) error {
if payload.IP == "" {
return errPayload(errIPRequired)
}
if net.ParseIP(payload.IP) == nil {
return errPayload(errIPInvalid)
}
if payload.Version == "" {
return errPayload(errAgentVersionRequired)
}
return nil
}
type payloadError string
func (e payloadError) Error() string { return string(e) }
func errPayload(message string) error { return payloadError(message) }
func applyNodeRuntime(ctx context.Context, node *model.OpenFlareNode, payload NodePayload, preserveName bool) {
if !preserveName || strings.TrimSpace(node.Name) == "" {
if strings.TrimSpace(payload.Name) != "" {
node.Name = strings.TrimSpace(payload.Name)
}
}
if !node.IPManualOverride {
node.IP = strings.TrimSpace(payload.IP)
}
node.Version = strings.TrimSpace(payload.Version)
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
node.Status = nodeStatusOnline
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
now := time.Now()
node.LastSeenAt = &now
node.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
if !node.GeoManualOverride {
ofgeoip.ApplyNodeGeoFromIP(ctx, node, node.IP)
}
}
func truncateForDatabase(value string, maxVal int) string {
if maxVal <= 0 {
return ""
}
runes := []rune(strings.TrimSpace(value))
if len(runes) <= maxVal {
return string(runes)
}
return string(runes[:maxVal])
}
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
reported := normalizeIP(reportedIP)
remote := normalizeRemoteAddr(remoteAddr)
if reported == "" {
return remote
}
if isPublicNodeIP(reported) {
return reported
}
if isPublicNodeIP(remote) {
return remote
}
return reported
}
func normalizeIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
host := raw
if strings.Contains(raw, ":") {
if h, _, err := net.SplitHostPort(raw); err == nil {
host = h
}
}
host = strings.TrimPrefix(host, "[")
host = strings.TrimSuffix(host, "]")
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
return ""
}
func normalizeRemoteAddr(remoteAddr string) string {
remoteAddr = strings.TrimSpace(remoteAddr)
if remoteAddr == "" {
return ""
}
host, _, err := net.SplitHostPort(remoteAddr)
if err != nil {
return normalizeIP(remoteAddr)
}
return normalizeIP(host)
}
func isPublicNodeIP(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
if ip == nil {
return false
}
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func buildAgentSettings(ctx context.Context, node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings {
autoUpdate := false
if node != nil {
autoUpdate = node.AutoUpdateEnabled
}
if strings.TrimSpace(updateChannel) == "" {
updateChannel = releaseChannelStable
}
// 从 SystemConfig 读取配置,使用默认值作为降级
heartbeatInterval, _ := repository.GetIntByKey(ctx, model.ConfigKeyAgentHeartbeatInterval)
if heartbeatInterval <= 0 {
heartbeatInterval = defaultAgentHeartbeatInterval
}
wsUpgradeEnabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyAgentWebsocketUpgradeEnabled)
updateRepo, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentUpdateRepo)
if strings.TrimSpace(updateRepo.Value) == "" {
updateRepo.Value = defaultAgentUpdateRepo
}
return &Settings{
HeartbeatInterval: heartbeatInterval,
WebsocketUpgradeEnabled: wsUpgradeEnabled,
AutoUpdate: autoUpdate,
UpdateRepo: updateRepo.Value,
UpdateNow: updateNow,
UpdateChannel: updateChannel,
UpdateTag: strings.TrimSpace(updateTag),
RestartOpenrestyNow: restartOpenrestyNow,
}
}
func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenFlareNode) map[string]any {
if previous == nil || current == nil {
return map[string]any{}
}
changes := make(map[string]any)
appendIfChanged := func(key string, before any, after any) {
if before != after {
changes[key] = after
}
}
appendIfChanged("name", previous.Name, current.Name)
appendIfChanged("ip", previous.IP, current.IP)
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
appendIfChanged("version", previous.Version, current.Version)
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
appendIfChanged("status", previous.Status, current.Status)
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
appendIfChanged("last_error", previous.LastError, current.LastError)
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
changes["geo_latitude"] = current.GeoLatitude
}
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
changes["geo_longitude"] = current.GeoLongitude
}
if !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
}
return before.Equal(*after)
}
func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
payload.NodeID = strings.TrimSpace(payload.NodeID)
payload.Version = strings.TrimSpace(payload.Version)
payload.Result = strings.ToLower(strings.TrimSpace(payload.Result))
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), maxDatabaseTextLength)
payload.Checksum = strings.TrimSpace(payload.Checksum)
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
return payload
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
func RefreshAccessTokenCache(_ context.Context, node *model.OpenFlareNode) {
if node == nil {
return
}
tokenCache.storeNode(node.AccessToken, cloneNode(node))
}
func cloneNode(node *model.OpenFlareNode) *model.OpenFlareNode {
if node == nil {
return nil
}
cloned := *node
return &cloned
}
@@ -0,0 +1,133 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"net"
"testing"
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
pkggeoip "Wavelet/openflare/share/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"}
ofgeoip.ApplyNodeGeoFromIP(context.Background(), 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),
}
ofgeoip.ApplyNodeGeoFromIP(context.Background(), 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(context.Background(), 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)
}
}
@@ -0,0 +1,223 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"time"
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/node"
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/logger"
)
// RegisterWithAccessToken registers an agent on a reserved node token.
func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) {
_ = ofgeoip.EnsureRuntimeProvider(ctx)
payload = normalizeNodePayload(payload)
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
if err := validateNodePayload(payload); err != nil {
return nil, err
}
applyNodeRuntime(ctx, authNode, payload, true)
if err := repository.SaveOpenFlareNode(ctx, authNode); err != nil {
return nil, err
}
RefreshAccessTokenCache(ctx, authNode)
return &RegistrationResponse{
NodeID: authNode.NodeID,
AccessToken: authNode.AccessToken,
Name: authNode.Name,
}, nil
}
// RegisterWithDiscovery registers a new node using the global discovery token.
func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*RegistrationResponse, error) {
_ = ofgeoip.EnsureRuntimeProvider(ctx)
payload = normalizeNodePayload(payload)
if err := validateNodePayload(payload); err != nil {
return nil, err
}
nodeID, err := newServerNodeID()
if err != nil {
return nil, err
}
accessToken, err := newRandomToken()
if err != nil {
return nil, err
}
nodeName := payload.Name
if nodeName == "" {
nodeName = nodeID
}
record := &model.OpenFlareNode{
NodeID: nodeID,
Name: nodeName,
AccessToken: accessToken,
Status: nodeStatusOnline,
NodeType: "edge_node",
CapabilitiesJSON: "[]",
UpdateChannel: releaseChannelStable,
}
applyNodeRuntime(ctx, record, payload, false)
if err = repository.CreateOpenFlareNode(ctx, record); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errNodeIDConflict)
}
return nil, err
}
RefreshAccessTokenCache(ctx, record)
return &RegistrationResponse{
NodeID: record.NodeID,
AccessToken: record.AccessToken,
Name: record.Name,
}, nil
}
// HeartbeatNode updates runtime state and returns agent settings.
func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*HeartbeatResponse, error) {
_ = ofgeoip.EnsureRuntimeProvider(ctx)
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
payload.NodeID = authNode.NodeID
payload = normalizeNodePayload(payload)
if err := validateNodePayload(payload); err != nil {
return nil, err
}
previous := *authNode
updateNow := authNode.UpdateRequested
restartOpenrestyNow := authNode.RestartOpenrestyRequested
updateChannel := strings.TrimSpace(authNode.UpdateChannel)
updateTag := strings.TrimSpace(authNode.UpdateTag)
applyNodeRuntime(ctx, authNode, payload, true)
authNode.UpdateRequested = false
authNode.UpdateChannel = releaseChannelStable
authNode.UpdateTag = ""
authNode.RestartOpenrestyRequested = false
changes := collectHeartbeatChanges(&previous, authNode)
if len(changes) > 0 {
fields := make([]string, 0, len(changes))
for field := range changes {
fields = append(fields, field)
}
if err := repository.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil {
return nil, err
}
if previous.IP != authNode.IP {
if _, dispatchErr := cf.DispatchNodeSync(ctx, authNode.ID, "cloudflare_agent_ip_update"); dispatchErr != nil {
logger.ErrorF(ctx, "[Cloudflare] enqueue heartbeat node sync failed: node_id=%d error=%v", authNode.ID, dispatchErr)
}
}
}
RefreshAccessTokenCache(ctx, authNode)
reportedAt := time.Now()
if authNode.LastSeenAt != nil {
reportedAt = *authNode.LastSeenAt
}
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(ctx, authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow),
ActiveConfig: activeConfig,
WAFIPGroups: wafIPGroups,
}, nil
}
// GetActiveConfig returns the active configuration for an agent.
func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
config, err := getActiveConfigForAgent(ctx)
if err != nil {
if isActiveConfigNotFound(err) {
return nil, errors.New(errNoActiveConfig)
}
return nil, err
}
return config, 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.
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
now := time.Now()
payload = normalizeApplyLogPayload(payload)
if payload.NodeID == "" {
return nil, errors.New(errNodeIDRequired)
}
if payload.Version == "" {
return nil, errors.New(errVersionRequired)
}
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFailed {
return nil, errors.New(errInvalidApplyResult)
}
latest, err := repository.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID)
if err != nil {
return nil, err
}
if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) {
if err := repository.UpdateOpenFlareNodeFromApplyResult(ctx, payload.NodeID, payload.Result, payload.Version, payload.Message, now); err != nil {
return nil, err
}
return latest, nil
}
log := &model.OpenFlareApplyLog{
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,
}
if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); err != nil {
return nil, err
}
return log, nil
}
// ValidateDiscoveryToken delegates to the node package discovery token helper.
func ValidateDiscoveryToken(ctx context.Context, token string) error {
return node.ValidateDiscoveryToken(ctx, token)
}
@@ -0,0 +1,60 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
const (
agentTokenHeader = "X-Agent-Token" //nolint:gosec // HTTP header name, not a credential value
agentNodeContextKey = "agent_node"
)
// Auth validates X-Agent-Token against of_nodes.access_token.
func Auth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
node, err := AuthenticateAccessToken(c.Request.Context(), token)
if err != nil {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
c.Set(agentNodeContextKey, node)
c.Next()
}
}
// RegisterAuth accepts either a node access token or the global discovery token.
func RegisterAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
c.Set(agentNodeContextKey, node)
c.Next()
return
}
if err := ValidateDiscoveryToken(c.Request.Context(), token); err != nil {
response.AbortUnauthorized(c, errInvalidDiscoveryToken)
return
}
c.Set("discovery_enabled", true)
c.Next()
}
}
// NodeFromContext returns the authenticated agent node.
func NodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
value, ok := c.Get(agentNodeContextKey)
if !ok {
return nil, false
}
node, ok := value.(*model.OpenFlareNode)
return node, ok
}
@@ -0,0 +1,198 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupAgentAuthTestDB(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.SystemConfig{},
))
db.SetDB(sqliteDB)
tokenCache.reset()
return func() {
db.SetDB(nil)
tokenCache.reset()
}
}
func TestAuthenticateAccessToken(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-auth-1",
Name: "edge",
AccessToken: "valid-agent-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
t.Run("valid token", func(t *testing.T) {
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
require.NoError(t, err)
assert.Equal(t, "node-auth-1", node.NodeID)
})
t.Run("cached token", func(t *testing.T) {
originalLoader := tokenCache.loadNodeByToken
t.Cleanup(func() {
tokenCache.loadNodeByToken = originalLoader
})
tokenCache.loadNodeByToken = func(context.Context, string) (*model.OpenFlareNode, error) {
t.Fatal("db should not be queried for cached token")
return nil, nil
}
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
require.NoError(t, err)
assert.Equal(t, "node-auth-1", node.NodeID)
})
t.Run("missing token", func(t *testing.T) {
_, err := AuthenticateAccessToken(ctx, "")
require.Error(t, err)
assert.Contains(t, err.Error(), errMissingAgentToken)
})
t.Run("invalid token", func(t *testing.T) {
_, err := AuthenticateAccessToken(ctx, "invalid-token")
require.Error(t, err)
})
}
func TestAgentAuthMiddleware(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-mw-1",
Name: "edge",
AccessToken: "middleware-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
router := testhelper.NewTestGinEngine()
router.GET("/protected", Auth(), func(c *gin.Context) {
node, ok := NodeFromContext(c)
if !ok {
c.Status(http.StatusInternalServerError)
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"node_id": node.NodeID}))
})
t.Run("authorized request", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set(agentTokenHeader, "middleware-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
assert.Empty(t, apiResp.ErrorMsg)
})
t.Run("unauthorized request", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set(agentTokenHeader, "bad-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusUnauthorized, resp.Code)
})
}
func TestAgentRegisterAuthMiddleware(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-register-1",
Name: "edge",
AccessToken: "existing-node-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyAgentDiscoveryToken, "discovery-token"))
router := testhelper.NewTestGinEngine()
router.POST("/register", RegisterAuth(), func(c *gin.Context) {
if node, ok := NodeFromContext(c); ok {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
return
}
if _, ok := c.Get("discovery_enabled"); ok {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "discovery"}))
return
}
c.Status(http.StatusInternalServerError)
})
t.Run("existing node token", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/register", nil)
req.Header.Set(agentTokenHeader, "existing-node-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
data, ok := apiResp.Data.(map[string]any)
require.True(t, ok)
assert.Equal(t, "node", data["mode"])
})
t.Run("discovery token", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/register", nil)
req.Header.Set(agentTokenHeader, "discovery-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
data, ok := apiResp.Data.(map[string]any)
require.True(t, ok)
assert.Equal(t, "discovery", data["mode"])
})
}
@@ -0,0 +1,234 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"go.uber.org/zap"
)
const (
accessLogPathMaxLength = 100
accessLogUserAgentMaxLength = 512
accessLogCacheStatusMaxLength = 32
)
// PersistHeartbeatObservability stores profile, host metrics, edge health, and access logs.
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
if strings.TrimSpace(nodeID) == "" {
return
}
if payload.Profile == nil &&
payload.HostMetrics == nil &&
payload.EdgeHealth == nil &&
len(payload.AccessLogs) == 0 &&
len(payload.Buffered) == 0 &&
payload.HealthEvents == nil {
return
}
accessLogRecords, err := buildNodeAccessLogRecords(ctx, nodeID, payload.AccessLogs, payload.Buffered, reportedAt)
if err != nil {
zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
return
}
profile := buildNodeSystemProfileModel(nodeID, payload.Profile, reportedAt)
healthEvents := healthEventInputs(payload.HealthEvents)
if err := repository.PersistOpenFlareNodePGObservability(
ctx,
profile,
nodeID,
healthEvents,
payload.HealthEvents != nil,
reportedAt,
nil,
); err != nil {
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
return
}
if err := persistBufferedObservability(ctx, nodeID, payload.Buffered, reportedAt); err != nil {
zap.L().Error("persist buffered observability failed", zap.String("node_id", nodeID), zap.Error(err))
}
if err := persistNodeMetricSnapshot(ctx, nodeID, payload.HostMetrics, reportedAt); err != nil {
zap.L().Error("persist metric snapshot failed", zap.String("node_id", nodeID), zap.Error(err))
}
if err := persistNodeEdgeHealth(ctx, nodeID, payload.EdgeHealth, payload.OpenrestyStatus, reportedAt); err != nil {
zap.L().Error("persist edge health failed", zap.String("node_id", nodeID), zap.Error(err))
}
if err := persistNodeAccessLogs(ctx, nodeID, accessLogRecords, reportedAt); err != nil {
zap.L().Error("persist heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
}
}
func persistBufferedObservability(ctx context.Context, nodeID string, records []BufferedObservabilityRecord, reportedAt time.Time) error {
for _, record := range records {
if err := persistNodeMetricSnapshot(ctx, nodeID, record.HostMetrics, reportedAt); err != nil {
return err
}
if err := persistNodeEdgeHealth(ctx, nodeID, record.EdgeHealth, "", reportedAt); err != nil {
return err
}
}
return nil
}
func persistNodeEdgeHealth(ctx context.Context, nodeID string, health *NodeEdgeHealth, fallbackStatus string, reportedAt time.Time) error {
if health == nil {
return nil
}
status := strings.TrimSpace(health.Status)
if status == "" {
status = strings.TrimSpace(fallbackStatus)
}
if status == "" {
status = openrestyStatusUnknown
}
return repository.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{
NodeID: nodeID,
CapturedAt: timeFromUnix(health.CapturedAtUnix, reportedAt),
Status: status,
Connections: health.Connections,
})
}
func buildNodeSystemProfileModel(nodeID string, profile *NodeSystemProfile, reportedAt time.Time) *model.OpenFlareNodeSystemProfile {
if profile == nil {
return nil
}
return &model.OpenFlareNodeSystemProfile{
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),
}
}
func healthEventInputs(events []NodeHealthEvent) []repository.OpenFlareHealthEventInput {
if events == nil {
return nil
}
out := make([]repository.OpenFlareHealthEventInput, 0, len(events))
for _, event := range events {
out = append(out, repository.OpenFlareHealthEventInput{
EventType: event.EventType,
Severity: event.Severity,
Message: event.Message,
TriggeredAtUnix: event.TriggeredAtUnix,
Metadata: event.Metadata,
})
}
return out
}
func persistNodeMetricSnapshot(ctx context.Context, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error {
if snapshot == nil {
return nil
}
record := &model.OpenFlareMetricSnapshot{
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,
// NetworkRx/Tx no longer collected from agents; CH columns remain 0.
}
return repository.InsertOpenFlareMetricSnapshot(ctx, record)
}
func buildNodeAccessLogRecords(ctx context.Context, nodeID string, direct []NodeAccessLog, buffered []BufferedObservabilityRecord, reportedAt time.Time) ([]*model.OpenFlareAccessLog, error) {
total := len(direct)
for _, record := range buffered {
total += len(record.AccessLogs)
}
if total == 0 {
return nil, nil
}
records := make([]*model.OpenFlareAccessLog, 0, total)
appendLogs := func(logs []NodeAccessLog) {
for _, item := range logs {
bytesSent := max(item.BytesSent, 0)
requestLength := max(item.RequestLength, 0)
requestTimeMs := max(item.RequestTimeMs, 0)
record := &model.OpenFlareAccessLog{
NodeID: nodeID,
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
Region: resolveAccessLogRegion(ctx, item.RemoteAddr),
Host: strings.TrimSpace(item.Host),
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
UserAgent: truncateForDatabase(strings.TrimSpace(item.UserAgent), accessLogUserAgentMaxLength),
CacheStatus: truncateForDatabase(strings.TrimSpace(item.CacheStatus), accessLogCacheStatusMaxLength),
StatusCode: item.StatusCode,
BytesSent: bytesSent,
RequestLength: requestLength,
RequestTimeMs: requestTimeMs,
}
records = append(records, record)
}
}
appendLogs(direct)
for _, record := range buffered {
appendLogs(record.AccessLogs)
}
return records, nil
}
func persistNodeAccessLogs(ctx context.Context, _ string, records []*model.OpenFlareAccessLog, _ time.Time) error {
if len(records) == 0 {
return nil
}
return repository.InsertOpenFlareAccessLogsBatch(ctx, records)
}
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
func ReconcileScopedNodeHealthEvents(ctx context.Context, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
return repository.ReconcileOpenFlareHealthEvents(ctx, nodeID, healthEventInputs(events), reportedAt, managedEventTypes)
}
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
if unixSeconds <= 0 {
return fallback
}
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 ""
}
raw, err := json.Marshal(value)
if err != nil {
return ""
}
return string(raw)
}
@@ -0,0 +1,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"testing"
"time"
)
func TestBuildNodeAccessLogRecordsPreservesBytesSent(t *testing.T) {
reportedAt := time.Date(2026, 7, 12, 10, 0, 0, 0, time.UTC)
records, err := buildNodeAccessLogRecords(context.Background(), "node-a", []NodeAccessLog{
{
LoggedAtUnix: reportedAt.Unix(),
RemoteAddr: "203.0.113.10",
Host: "api.example.com",
Path: "/v1/ping",
StatusCode: 200,
BytesSent: 4096,
},
}, nil, reportedAt)
if err != nil {
t.Fatalf("buildNodeAccessLogRecords() error = %v", err)
}
if len(records) != 1 {
t.Fatalf("expected one access log record, got %d", len(records))
}
if records[0].BytesSent != 4096 {
t.Fatalf("BytesSent = %d, want 4096", records[0].BytesSent)
}
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import pkgprotocol "Wavelet/openflare/share/protocol"
// NodePayload is the data sent by an agent on registration or heartbeat.
type NodePayload = pkgprotocol.NodePayload
// NodeSystemProfile carries static host information reported by an agent.
type NodeSystemProfile = pkgprotocol.NodeSystemProfile
// NodeMetricSnapshot holds a point-in-time resource-usage sample from an agent.
type NodeMetricSnapshot = pkgprotocol.NodeMetricSnapshot
// NodeEdgeHealth is L2 OpenResty health + connections.
type NodeEdgeHealth = pkgprotocol.NodeEdgeHealth
// NodeAccessLog is a single access-log record forwarded by an agent.
type NodeAccessLog = pkgprotocol.NodeAccessLog
// BufferedObservabilityRecord bundles multiple observability payloads into one upload.
type BufferedObservabilityRecord = pkgprotocol.BufferedObservabilityRecord
// NodeHealthEvent represents a discrete health-state change on an agent node.
type NodeHealthEvent = pkgprotocol.NodeHealthEvent
// ApplyLogPayload carries the result of a configuration-apply attempt reported by an agent.
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
// Settings contains remote-control directives sent from the server to an agent.
type Settings = pkgprotocol.AgentSettings
// ActiveConfigMeta describes the currently active configuration version on the server.
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
// SupportFile represents a supplementary file bundled with an agent configuration package.
type SupportFile = pkgprotocol.SupportFile
// WAFIPGroup is a named IP-address group used in WAF allow/block rules.
type WAFIPGroup = pkgprotocol.WAFIPGroup
// WAFIPGroupSyncRequest is sent by an agent to request an incremental WAF IP-group sync.
type WAFIPGroupSyncRequest = pkgprotocol.WAFIPGroupSyncRequest
// WAFIPGroupSyncResponse carries the server's reply to a WAF IP-group sync request.
type WAFIPGroupSyncResponse = pkgprotocol.WAFIPGroupSyncResponse
// Backward-compatible names used by server routers and handlers.
// WAFIPGroupSyncInput is an alias for WAFIPGroupSyncRequest kept for backward compatibility.
type WAFIPGroupSyncInput = WAFIPGroupSyncRequest
// WAFIPGroupSyncResult is an alias for WAFIPGroupSyncResponse kept for backward compatibility.
type WAFIPGroupSyncResult = WAFIPGroupSyncResponse
@@ -0,0 +1,298 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net/http"
"strconv"
"Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/domain/pages"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/openflare/share/protocol"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// RegisterHandler registers or discovers an agent node.
// @Summary 注册或发现 Agent 节点
// @Description 使用节点 access token 重新注册,或使用全局 discovery token 发现新节点;请求头需携带 X-Agent-Token
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.NodePayload true "节点上报数据"
// @Success 200 {object} response.Any{data=agent.RegistrationResponse} "注册成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/nodes/register [post]
func RegisterHandler(c *gin.Context) {
var payload NodePayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
var (
result *RegistrationResponse
err error
)
if authNode, ok := NodeFromContext(c); ok {
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
} else {
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
}
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// HeartbeatHandler records agent heartbeat state.
// @Summary Agent 心跳上报
// @Description 上报节点状态、指标与健康事件,返回远程控制配置与活跃配置元信息
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.NodePayload true "心跳数据"
// @Success 200 {object} response.Any{data=agent.HeartbeatResponse} "心跳成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/nodes/heartbeat [post]
func HeartbeatHandler(c *gin.Context) {
var payload NodePayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := NodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
heartbeat, err := HeartbeatNode(c.Request.Context(), authNode, payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(heartbeat))
}
// GetActiveConfigHandler returns the active configuration version.
// @Summary 获取活跃配置版本
// @Description 返回当前生效的完整配置包,供 Agent 拉取并应用
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Success 200 {object} response.Any{data=agent.ConfigResponse} "活跃配置"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/config-versions/active [get]
func GetActiveConfigHandler(c *gin.Context) {
if _, ok := NodeFromContext(c); !ok {
response.AbortUnauthorized(c, errNodeMissingFromContext)
return
}
config, err := GetActiveConfig(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(config))
}
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
// @Summary 同步 WAF IP 组
// @Description 按 ID 与校验和增量同步 WAF IP 组定义
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.WAFIPGroupSyncInput true "同步请求"
// @Success 200 {object} response.Any{data=agent.WAFIPGroupSyncResult} "同步结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/waf/ip-groups/sync [post]
func SyncWAFIPGroupsHandler(c *gin.Context) {
var input WAFIPGroupSyncInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := SyncWAFIPGroups(c.Request.Context(), input)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// ReportApplyLogHandler records an agent apply log entry.
// @Summary 上报配置应用日志
// @Description 记录 Agent 配置下发与应用结果
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.ApplyLogPayload true "应用日志"
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "日志记录"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/apply-logs [post]
func ReportApplyLogHandler(c *gin.Context) {
var payload ApplyLogPayload
if !apiutil.BindJSON(c, &payload) {
return
}
if authNode, ok := NodeFromContext(c); ok {
payload.NodeID = authNode.NodeID
}
log, err := ReportApplyLog(c.Request.Context(), payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(log))
}
// GetPagesDeploymentHashHandler returns the upload SHA-256 hash for a Pages deployment package.
// @Summary 查询 Pages 部署包哈希
// @Description 返回 upload 框架记录的 SHA-256 哈希,供 Agent 对比本地缓存并按需拉取部署包(兼容旧路径)
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Param deployment_id path int true "部署 ID"
// @Success 200 {object} response.Any{data=protocol.PagesDeploymentHashResponse}
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/deployments/{deployment_id}/hash [get]
func GetPagesDeploymentHashHandler(c *gin.Context) {
deploymentID, ok := pagesUintParam(c, "deployment_id")
if !ok {
return
}
hash, err := pages.GetDeploymentPackageHash(c.Request.Context(), deploymentID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(protocol.PagesDeploymentHashResponse{
DeploymentID: deploymentID,
Hash: hash,
}))
}
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
// @Summary 下载 Pages 部署包
// @Description 流式下载指定部署的静态资源压缩包,供 Agent 边缘分发(兼容旧路径)
// @Tags openflare-agent
// @Produce application/octet-stream
// @Security AgentTokenAuth
// @Param deployment_id path int true "部署 ID"
// @Success 200 {file} binary "部署包文件"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/deployments/{deployment_id}/package [get]
func DownloadPagesPackageHandler(c *gin.Context) {
deploymentID, ok := pagesUintParam(c, "deployment_id")
if !ok {
return
}
packageObj, err := pages.OpenDeploymentPackage(c.Request.Context(), deploymentID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer func() { _ = packageObj.Body.Close() }()
c.Header("Content-Disposition", "attachment; filename="+packageObj.FileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
}
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
}
// GetPagesProjectLatestHashHandler returns the hash of a project's currently active deployment.
// @Summary 查询 Pages 项目最新激活部署哈希
// @Description 按项目 ID 返回当前激活部署的包哈希(类似 latest 指针),Agent 无需关心具体部署 ID
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Param project_id path int true "Pages 项目 ID"
// @Success 200 {object} response.Any{data=protocol.PagesProjectLatestHashResponse}
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/projects/{project_id}/latest/hash [get]
func GetPagesProjectLatestHashHandler(c *gin.Context) {
projectID, ok := pagesUintParam(c, "project_id")
if !ok {
return
}
metadata, err := pages.GetProjectLatestPackageMetadata(c.Request.Context(), projectID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: metadata.DeploymentID,
Hash: metadata.Hash,
PackageSize: metadata.PackageSize,
FileCount: metadata.FileCount,
TotalSize: metadata.TotalSize,
}))
}
// DownloadPagesProjectLatestPackageHandler streams the active deployment package for a project.
// @Summary 下载 Pages 项目最新激活部署包
// @Description 按项目 ID 下载当前激活部署的压缩包,供 Agent 边缘分发
// @Tags openflare-agent
// @Produce application/octet-stream
// @Security AgentTokenAuth
// @Param project_id path int true "Pages 项目 ID"
// @Success 200 {file} binary "部署包文件"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/projects/{project_id}/latest/package [get]
func DownloadPagesProjectLatestPackageHandler(c *gin.Context) {
projectID, ok := pagesUintParam(c, "project_id")
if !ok {
return
}
packageObj, err := pages.OpenProjectLatestPackage(c.Request.Context(), projectID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer func() { _ = packageObj.Body.Close() }()
c.Header("Content-Disposition", "attachment; filename="+packageObj.FileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
}
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
}
func pagesUintParam(c *gin.Context, name string) (uint, bool) {
raw := c.Param(name)
if raw == "" {
response.AbortBadRequest(c, "无效的 ID")
return 0, false
}
id64, err := strconv.ParseUint(raw, 10, 64)
if err != nil || id64 == 0 {
response.AbortBadRequest(c, "无效的 ID")
return 0, false
}
return uint(id64), true
}
// WebSocketHandler upgrades an authenticated agent websocket connection.
// @Summary Agent WebSocket 连接
// @Description 升级为 WebSocket 长连接,用于实时推送配置同步、WAF IP 组等指令;需携带 X-Agent-Token
// @Tags openflare-agent
// @Security AgentTokenAuth
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/ws [get]
func WebSocketHandler(c *gin.Context) {
authNode, ok := NodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
}
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"time"
"Wavelet/openflare/plugins/server/kernel/model"
)
const (
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFailed = "failed"
)
// RegistrationResponse is returned after agent registration.
// Server uses access_token; the agent client expects agent_token via RegisterNodeResponse.
type RegistrationResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"access_token"`
Name string `json:"name"`
}
// ConfigResponse is the full active config payload for agents.
// Server uses time.Time for CreatedAt; the agent client uses string via ActiveConfigResponse.
type ConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
SourceConfigJSON string `json:"source_config_json"`
SupportFiles []SupportFile `json:"support_files"`
CreatedAt time.Time `json:"created_at"`
}
// HeartbeatResponse is the heartbeat handler result.
type HeartbeatResponse struct {
Node *model.OpenFlareNode `json:"node"`
AgentSettings *Settings `json:"agent_settings"`
ActiveConfig *ActiveConfigMeta `json:"active_config"`
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
}
@@ -0,0 +1,263 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"slices"
"sort"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/share/protocol"
openrestyrender "Wavelet/openflare/share/render/openresty"
)
type activeConfigSnapshot struct {
WAF openrestyrender.WAFDocument `json:"waf"`
}
type runtimeIPMatchConfig struct {
IPs []string `json:"ips,omitempty"`
CIDRs []string `json:"cidrs,omitempty"`
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
}
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
return validatedAgentWAFIPGroups(ctx, ids, false)
}
// 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) {
groups, err := validatedAgentWAFIPGroups(ctx, ids, true)
if err != nil {
return nil, err
}
changed := make([]WAFIPGroup, 0, len(groups))
for _, group := range groups {
if strings.TrimSpace(checksums[strconv.FormatUint(uint64(group.ID), 10)]) == group.Checksum {
continue
}
changed = append(changed, group)
}
return changed, nil
}
func validatedAgentWAFIPGroups(ctx context.Context, ids []uint, fallbackToActive bool) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
if len(targetIDs) == 0 && fallbackToActive {
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
validationIDs := uniqueUintIDs(append(append([]uint{}, activeIDs...), targetIDs...))
allGroups, err := buildAgentWAFIPGroups(ctx, validationIDs)
if err != nil {
return nil, err
}
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(allGroups))
for _, group := range allGroups {
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = group
}
if err = protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups); err != nil {
return nil, err
}
targetSet := make(map[uint]struct{}, len(targetIDs))
for _, id := range targetIDs {
targetSet[id] = struct{}{}
}
result := make([]WAFIPGroup, 0, len(targetIDs))
for _, group := range allGroups {
if _, ok := targetSet[group.ID]; ok {
result = append(result, group)
}
}
return result, nil
}
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
return []WAFIPGroup{}, nil
}
slices.Sort(ids)
groups, err := repository.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 := repository.GetActiveConfigVersion(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 {
// Retain legacy flattened references while older active snapshots may
// still exist during a rolling Server upgrade.
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for nodeID, node := range group.Graph.Nodes {
if node.Type != "ip_match" {
continue
}
ids, err := runtimeIPMatchGroupIDs(node.Config)
if err != nil {
return nil, fmt.Errorf("活动配置 WAF 规则 %d 节点 %s 的 IP 匹配配置无效: %w", group.ID, nodeID, err)
}
for _, id := range ids {
if id > 0 {
idSet[id] = struct{}{}
}
}
}
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
ids = append(ids, id)
}
slices.Sort(ids)
return ids, nil
}
func runtimeIPMatchGroupIDs(raw json.RawMessage) ([]uint, error) {
var config runtimeIPMatchConfig
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&config); err != nil {
return nil, err
}
return config.IPGroupIDs, 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 = []openrestyrender.WAFRuleGroup{}
}
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,272 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strconv"
"strings"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/share/protocol"
db "Wavelet/plugins/infra/database"
"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{},
&model.ConfigVersion{},
))
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(&model.ConfigVersion{
Version: "20260618-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "test-checksum",
IsActive: true,
}).Error)
}
func seedActiveConfigWithWAFGraphIPGroup(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": "graph refs",
"enabled": true,
"graph": map[string]any{
"entry": "start",
"nodes": map[string]any{
"start": map[string]any{
"type": "start",
"config": map[string]any{},
"next": map[string]string{"next": "match"},
},
"match": map[string]any{
"type": "ip_match",
"config": map[string]any{
"ip_group_ids": []uint{ipGroupID},
},
"next": map[string]string{"true": "allow", "false": "allow"},
},
"allow": map[string]any{
"type": "allow",
"config": map[string]any{},
},
},
},
},
},
"bindings": []any{},
},
}
snapshotJSON, err := json.Marshal(snapshot)
require.NoError(t, err)
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "20260713-graph-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "graph-test-checksum",
IsActive: true,
}).Error)
}
func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "graph runtime group",
Type: "manual",
Enabled: true,
IPList: `["192.0.2.88"]`,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFGraphIPGroup(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, []string{"192.0.2.88"}, groups[0].IPList)
}
func TestChangedWAFIPGroupsForAgentRejectsMalformedIPMatchConfig(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "20260713-malformed-001",
SnapshotJSON: `{"waf":{"rule_groups":[{"id":7,"graph":{"entry":"match","nodes":{` +
`"match":{"type":"ip_match","config":{"ip_group_ids":"not-an-array"}}}}}],"bindings":[]}}`,
Checksum: "malformed-test-checksum",
IsActive: true,
}).Error)
_, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.ErrorContains(t, err, "规则 7 节点 match")
require.ErrorContains(t, err, "IP 匹配配置无效")
}
func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: strings.Repeat("x", protocol.MaxWAFIPGroupSnapshotBytes),
Type: "manual",
Enabled: true,
IPList: `[]`,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
agentGroup, err := buildAgentWAFIPGroup(ipGroup)
require.NoError(t, err)
_, err = ChangedWAFIPGroupsForAgent(ctx, []uint{ipGroup.ID}, map[string]string{
strconv.FormatUint(uint64(ipGroup.ID), 10): agentGroup.Checksum,
})
require.ErrorContains(t, err, "WAF IP 组快照大小")
require.ErrorContains(t, err, "超过上限")
}
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, repository.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, repository.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, repository.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, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
ipGroup.Enabled = false
require.NoError(t, repository.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)
}
@@ -0,0 +1,58 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"log/slog"
"Wavelet/openflare/plugins/server/kernel/repository"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
)
// HandleWSStatus processes an agent websocket status payload (replaces HTTP heartbeat in WS mode).
func HandleWSStatus(ctx context.Context, nodeID, remoteAddr string, rawPayload json.RawMessage) {
var payload NodePayload
if err := json.Unmarshal(rawPayload, &payload); err != nil {
slog.Debug("agent ws status payload decode failed", "node_id", nodeID, "error", err)
return
}
authNode, err := repository.GetOpenFlareNodeByNodeID(ctx, nodeID)
if err != nil {
slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err)
return
}
payload.IP = resolveReportedNodeIP(payload.IP, remoteAddr)
response, err := HeartbeatNode(ctx, authNode, payload)
if err != nil {
slog.Debug("agent ws status handling failed", "node_id", nodeID, "error", err)
return
}
settingsSent := false
if response.AgentSettings != nil {
settingsSent = ofws.SendAgentSettings(nodeID, response.AgentSettings)
}
activeConfigSent := false
if response.ActiveConfig != nil {
activeConfigSent = ofws.SendAgentActiveConfig(nodeID, response.ActiveConfig)
}
wafIPGroupsSent := false
if len(response.WAFIPGroups) > 0 {
wafIPGroupsSent = ofws.SendAgentWAFIPGroups(nodeID, response.WAFIPGroups)
}
slog.Debug("agent ws status processed",
"node_id", nodeID,
"current_version", payload.CurrentVersion,
"openresty_status", payload.OpenrestyStatus,
"settings_sent", settingsSent,
"active_config_sent", activeConfigSent,
"waf_ip_groups_sent", wafIPGroupsSent,
)
}
@@ -0,0 +1,170 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package fleet contains shared edge-fleet task metadata used by plugin registration.
package fleet
import (
"context"
"fmt"
"sync"
"time"
"Wavelet/openflare/plugins/server/domain/option/uptimekuma"
"Wavelet/openflare/plugins/server/domain/tls"
"Wavelet/openflare/plugins/server/domain/waf"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/task"
)
const (
// SSLRenewTask renews due ACME TLS certificates.
SSLRenewTask = "openflare:ssl_renew"
// TaskTypeSSLRenew is the admin task type for SSL renewal.
TaskTypeSSLRenew = "of_ssl_renew"
// WAFIPGroupSyncTask syncs due automatic/subscription WAF IP groups.
WAFIPGroupSyncTask = "openflare:waf_ip_group_sync"
// TaskTypeWAFIPGroupSync is the admin task type for WAF IP group sync.
TaskTypeWAFIPGroupSync = "of_waf_ip_group_sync"
// UptimeKumaSyncTask synchronizes proxy routes to Uptime Kuma monitors.
UptimeKumaSyncTask = "openflare:uptime_kuma_sync"
// TaskTypeUptimeKumaSync is the admin task type for Uptime Kuma sync.
TaskTypeUptimeKumaSync = "of_uptime_kuma_sync"
// LogDBSwitchTask 切换日志数据库任务标识。
LogDBSwitchTask = "openflare:log_db_switch"
// TaskTypeLogDBSwitch is the admin task type for log database switch.
TaskTypeLogDBSwitch = "of_log_db_switch"
)
var (
lastUptimeKumaSyncTime time.Time
uptimeKumaSyncMutex sync.Mutex
)
// SSLRenewMeta describes the SSL renewal task.
var SSLRenewMeta = task.TaskMeta{
Type: TaskTypeSSLRenew,
AsynqTask: SSLRenewTask,
Name: "OpenFlare SSL 自动续期",
Description: "扫描即将到期的 ACME 证书并触发自动续期",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// WAFIPGroupSyncMeta describes the WAF IP group sync task.
var WAFIPGroupSyncMeta = task.TaskMeta{
Type: TaskTypeWAFIPGroupSync,
AsynqTask: WAFIPGroupSyncTask,
Name: "OpenFlare WAF IP 组同步",
Description: "同步到期的自动规则与订阅类型 WAF IP 组",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// UptimeKumaSyncMeta describes the Uptime Kuma sync task.
var UptimeKumaSyncMeta = task.TaskMeta{
Type: TaskTypeUptimeKumaSync,
AsynqTask: UptimeKumaSyncTask,
Name: "OpenFlare Uptime Kuma 同步",
Description: "将启用的代理规则同步到 Uptime Kuma 监控",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// LogDBSwitchMeta 描述切换日志数据库任务。
var LogDBSwitchMeta = task.TaskMeta{
Type: TaskTypeLogDBSwitch,
AsynqTask: LogDBSwitchTask,
Name: "切换日志数据库",
Description: "复制迁移日志数据并在成功后切换日志主库(期间禁止日志写入)",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.TaskParam{
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
},
}
// SSLRenewHandler renews due TLS certificates.
type SSLRenewHandler struct{}
// Execute runs SSL certificate renewal for all due certificates.
func (h *SSLRenewHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
task.AppendLog(ctx, "开始扫描待续期证书")
if err := tls.RunSSLRenewJob(ctx); err != nil {
task.AppendLog(ctx, "SSL 自动续期失败: %v", err)
return nil, err
}
msg := "SSL 自动续期任务完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
// WAFIPGroupSyncHandler syncs due WAF IP groups to agents.
type WAFIPGroupSyncHandler struct{}
// Execute syncs all due automatic/subscription WAF IP groups.
func (h *WAFIPGroupSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
task.AppendLog(ctx, "开始同步到期的 WAF IP 组")
if err := waf.SyncDueWAFIPGroups(ctx); err != nil {
task.AppendLog(ctx, "WAF IP 组同步失败: %v", err)
return nil, err
}
msg := "WAF IP 组同步完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
// UptimeKumaSyncHandler synchronizes proxy routes to Uptime Kuma.
type UptimeKumaSyncHandler struct{}
// Execute runs Uptime Kuma sync when integration is enabled and the interval has elapsed.
func (h *UptimeKumaSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
// 从 SystemConfig 读取 UptimeKuma 配置
enabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyUptimeKumaEnabled)
if !enabled {
msg := "Uptime Kuma 集成未启用,跳过执行"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
interval, _ := repository.GetIntByKey(ctx, model.ConfigKeyUptimeKumaSyncInterval)
if interval <= 0 {
interval = 5
}
if time.Since(lastUptimeKumaSyncTime) < time.Duration(interval)*time.Minute {
msg := fmt.Sprintf("距上次同步不足 %d 分钟,跳过执行", interval)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
if !uptimeKumaSyncMutex.TryLock() {
msg := "Uptime Kuma 同步任务正在执行,跳过本次调度"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
defer uptimeKumaSyncMutex.Unlock()
task.AppendLog(ctx, "开始同步代理规则到 Uptime Kuma")
if err := uptimekuma.SyncToUptimeKuma(ctx); err != nil {
task.AppendLog(ctx, "Uptime Kuma 同步失败: %v", err)
return nil, err
}
lastUptimeKumaSyncTime = time.Now()
msg := "Uptime Kuma 同步完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
@@ -0,0 +1,36 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package fleet
import (
"context"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func TestUptimeKumaSyncHandlerSkipsWhenDisabled(t *testing.T) {
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{}))
db.SetDB(sqliteDB)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyUptimeKumaEnabled, "false"))
result, err := (&UptimeKumaSyncHandler{}).Execute(ctx, nil)
require.NoError(t, err)
require.NotNil(t, result)
assert.Contains(t, result.Message, "未启用")
}
@@ -0,0 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package flared defines shared error messages for tunnel client operations.
package flared
const (
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errTunnelNodeTypeMismatch = "此节点不是 TunnelClient 类型"
)
@@ -0,0 +1,134 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"fmt"
"net"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/relay"
"Wavelet/openflare/plugins/server/kernel/model"
)
const (
updateChannelStable = "stable"
defaultTunnelTargetPort = 80
)
func normalizeReleaseChannel(channel string) string {
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
return "preview"
}
return updateChannelStable
}
func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload {
payload.ClientVersion = strings.TrimSpace(payload.ClientVersion)
payload.FrpVersion = strings.TrimSpace(payload.FrpVersion)
payload.IP = strings.TrimSpace(payload.IP)
payload.TunnelStatus = strings.ToLower(strings.TrimSpace(payload.TunnelStatus))
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
payload.CurrentChecksum = strings.TrimSpace(payload.CurrentChecksum)
cleaned := make([]ConnectedRelay, 0, len(payload.ConnectedRelays))
for _, item := range payload.ConnectedRelays {
item.RelayNodeID = strings.TrimSpace(item.RelayNodeID)
item.Status = strings.ToLower(strings.TrimSpace(item.Status))
if item.RelayNodeID == "" {
continue
}
if item.Status == "" {
item.Status = "unknown"
}
cleaned = append(cleaned, item)
}
payload.ConnectedRelays = cleaned
return payload
}
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := repository.GetActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
return &ActiveConfigMeta{
Version: version.Version,
Checksum: version.Checksum,
}, nil
}
func listTunnelRelayNodes(ctx context.Context) ([]model.OpenFlareNode, error) {
nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil {
return nil, err
}
relays := make([]model.OpenFlareNode, 0)
for _, node := range nodes {
if node.NodeType == "tunnel_relay" {
relays = append(relays, node)
}
}
return relays, nil
}
func relayClientAddress(node *model.OpenFlareNode) string {
if node == nil {
return ""
}
port := node.RelayBindPort
if port <= 0 {
port = 7000
}
addr := strings.TrimSpace(node.RelayClientAccessAddr)
if addr == "" {
addr = strings.TrimSpace(node.IP)
}
if addr == "" {
return fmt.Sprintf("127.0.0.1:%d", port)
}
if _, _, err := net.SplitHostPort(addr); err == nil {
return addr
}
if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 {
return net.JoinHostPort(addr, strconv.Itoa(port))
}
return fmt.Sprintf("%s:%d", addr, port)
}
func parseTunnelTargetAddr(addr string) (string, int) {
addr = strings.TrimSpace(addr)
if addr == "" {
return "127.0.0.1", defaultTunnelTargetPort
}
host, portStr, err := net.SplitHostPort(addr)
if err != nil {
lastColon := strings.LastIndex(addr, ":")
if lastColon < 0 {
return addr, defaultTunnelTargetPort
}
host = addr[:lastColon]
portStr = addr[lastColon+1:]
}
port := defaultTunnelTargetPort
if _, scanErr := fmt.Sscanf(portStr, "%d", &port); scanErr != nil {
port = defaultTunnelTargetPort
}
if host == "" {
host = "127.0.0.1"
}
return host, port
}
func sanitizeProxyName(domain string) string {
return strings.ReplaceAll(strings.ReplaceAll(domain, ".", "-"), "*", "wildcard")
}
func buildTunnelSettings(ctx context.Context, node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *relay.Settings {
return relay.BuildSettings(ctx, node, updateNow, updateChannel, updateTag)
}
@@ -0,0 +1,223 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"errors"
"fmt"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/openflare/plugins/server/kernel/model"
"gorm.io/gorm"
)
const (
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFail = "failed"
maxApplyLogMessageLength = 16000
)
// Heartbeat processes an OpenFlared heartbeat and returns runtime settings.
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
if node == nil {
return nil, errors.New("tunnel client node is nil")
}
if node.NodeType != "tunnel_client" {
return nil, fmt.Errorf("node %s is not a tunnel_client", node.NodeID)
}
payload = normalizeFlaredHeartbeatPayload(payload)
previous := *node
updateNow := node.UpdateRequested
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
updateTag := strings.TrimSpace(node.UpdateTag)
now := time.Now().UTC()
changes := map[string]any{
"version": payload.ClientVersion,
"ext_version": payload.FrpVersion,
"current_version": payload.CurrentVersion,
"last_seen_at": now,
"status": nodeStatusOnline,
"update_requested": false,
"update_channel": updateChannelStable,
"update_tag": "",
}
if !previous.UpdateRequested {
delete(changes, "update_requested")
}
if previous.UpdateChannel == updateChannelStable {
delete(changes, "update_channel")
}
if previous.UpdateTag == "" {
delete(changes, "update_tag")
}
if !node.IPManualOverride && payload.IP != "" && previous.IP != payload.IP {
changes["ip"] = payload.IP
node.IP = payload.IP
}
node.Version = payload.ClientVersion
node.ExtVersion = payload.FrpVersion
node.CurrentVersion = payload.CurrentVersion
node.UpdateRequested = false
node.UpdateChannel = updateChannelStable
node.UpdateTag = ""
lastSeen := now
node.LastSeenAt = &lastSeen
node.Status = nodeStatusOnline
if err := repository.UpdateOpenFlareNodeColumns(ctx, node, changes); 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) {
return nil, err
}
return &HeartbeatResponse{
ActiveConfig: activeConfig,
TunnelSettings: buildTunnelSettings(ctx, node, updateNow, updateChannel, updateTag),
}, nil
}
// GetTunnelConfig builds the full tunnel routing config for an OpenFlared client.
func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelConfigResponse, error) {
if node == nil {
return nil, errors.New("node is nil")
}
activeVersion, err := getActiveConfigMeta(ctx)
if err != nil {
return nil, fmt.Errorf("no active config version: %w", err)
}
routes, err := repository.ListProxyRoutes(ctx)
if err != nil {
return nil, fmt.Errorf("get proxy routes: %w", err)
}
relayNodes, err := listTunnelRelayNodes(ctx)
if err != nil {
return nil, fmt.Errorf("get relay nodes: %w", err)
}
relays := make([]RelayInfo, 0, len(relayNodes))
for i := range relayNodes {
relayNode := relayNodes[i]
if relayNode.RelayStatus == "healthy" || relayNode.Status == nodeStatusOnline {
relays = append(relays, RelayInfo{
RelayNodeID: relayNode.NodeID,
Address: relayClientAddress(&relayNode),
AuthToken: relayNode.RelayAuthToken,
ProxyURL: strings.TrimSpace(relayNode.RelayClientProxyURL),
})
}
}
proxies := make([]ProxyEntry, 0)
for _, route := range routes {
if route == nil || route.UpstreamType != "tunnel" || route.TunnelNodeID == nil || *route.TunnelNodeID != node.ID {
continue
}
if !route.Enabled {
continue
}
zoneDomains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if domainErr != nil || len(zoneDomains) == 0 {
continue
}
localAddr, localPort := parseTunnelTargetAddr(route.TunnelTargetAddr)
proxies = append(proxies, ProxyEntry{
Name: fmt.Sprintf("%s-%s", node.NodeID, sanitizeProxyName(zoneDomains[0].Domain)),
Type: "http",
LocalAddr: localAddr,
LocalPort: localPort,
CustomDomains: zoneDomainNames(zoneDomains),
})
}
return &TunnelConfigResponse{
Version: activeVersion.Version,
Checksum: activeVersion.Checksum,
Relays: relays,
Proxies: proxies,
}, nil
}
func zoneDomainNames(domains []model.ZoneDomain) []string {
names := make([]string, 0, len(domains))
for _, domain := range domains {
names = append(names, domain.Domain)
}
return names
}
// ReportApplyLog records an apply result from OpenFlared.
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
now := time.Now().UTC()
payload = normalizeApplyLogPayload(payload)
if payload.NodeID == "" {
return nil, errors.New("node_id 不能为空")
}
if payload.Version == "" {
return nil, errors.New("version 不能为空")
}
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFail {
return nil, errors.New("result 仅支持 success、warning 或 failed")
}
latest, err := repository.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID)
if err != nil {
return nil, err
}
if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) {
if err := repository.UpdateOpenFlareNodeFromApplyResult(ctx, payload.NodeID, payload.Result, payload.Version, payload.Message, now); err != nil {
return nil, err
}
return latest, nil
}
log := &model.OpenFlareApplyLog{
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,
}
if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); err != nil {
return nil, err
}
return log, nil
}
func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
payload.NodeID = strings.TrimSpace(payload.NodeID)
payload.Version = strings.TrimSpace(payload.Version)
payload.Result = strings.ToLower(strings.TrimSpace(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 len(payload.Message) > maxApplyLogMessageLength {
payload.Message = payload.Message[:maxApplyLogMessageLength]
}
return payload
}
@@ -0,0 +1,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"strings"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
const ctxFlaredNodeKey = "flared_node"
// TunnelAuth authenticates flared requests using X-Tunnel-Token and verifies tunnel_client type.
func TunnelAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader("X-Tunnel-Token"))
node, err := agent.AuthenticateAccessToken(c.Request.Context(), token)
if err != nil {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
if node.NodeType != "tunnel_client" {
response.AbortForbidden(c, errTunnelNodeTypeMismatch)
return
}
c.Set(ctxFlaredNodeKey, node)
c.Next()
}
}
@@ -0,0 +1,112 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupFlaredMiddlewareTestDB(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{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func seedFlaredNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareNode {
t.Helper()
ctx := context.Background()
node := &model.OpenFlareNode{
NodeID: "flared-test-node",
Name: "flared-test",
Status: "pending",
NodeType: nodeType,
AccessToken: accessToken,
}
require.NoError(t, repository.CreateOpenFlareNode(ctx, node))
return node
}
func TestTunnelAuthMissingToken(t *testing.T) {
cleanup := setupFlaredMiddlewareTestDB(t)
defer cleanup()
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) {
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodGet, "/flared/test", nil)
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusUnauthorized, rec.Code)
}
func TestTunnelAuthRejectsWrongNodeType(t *testing.T) {
cleanup := setupFlaredMiddlewareTestDB(t)
defer cleanup()
seedFlaredNode(t, "edge_node", "edge-token-flared")
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) {
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodGet, "/flared/test", nil)
req.Header.Set("X-Tunnel-Token", "edge-token-flared")
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusForbidden, rec.Code)
}
func TestTunnelAuthAcceptsTunnelClient(t *testing.T) {
cleanup := setupFlaredMiddlewareTestDB(t)
defer cleanup()
node := seedFlaredNode(t, "tunnel_client", "tunnel-token-valid")
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) {
authNode, ok := c.Get(ctxFlaredNodeKey)
require.True(t, ok)
assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID)
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodGet, "/flared/test", nil)
req.Header.Set("X-Tunnel-Token", "tunnel-token-valid")
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusOK, rec.Code)
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"fmt"
"strings"
"time"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"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, ","),
},
})
}
if err := agent.ReconcileScopedNodeHealthEvents(ctx, 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,73 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"context"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"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{},
&model.SystemConfig{},
&model.ConfigVersion{},
))
db.SetDB(sqliteDB)
agent.ResetAuthCacheForTest()
return func() {
db.SetDB(nil)
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 := repository.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)
}
@@ -0,0 +1,30 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import pkgprotocol "Wavelet/openflare/share/protocol"
// HeartbeatPayload is an alias for FlaredHeartbeatPayload.
type HeartbeatPayload = pkgprotocol.FlaredHeartbeatPayload
// ConnectedRelay is an alias for FlaredConnectedRelay.
type ConnectedRelay = pkgprotocol.FlaredConnectedRelay
// ActiveConfigMeta is an alias for ActiveConfigMeta.
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
// HeartbeatResponse is an alias for FlaredHeartbeatResponse.
type HeartbeatResponse = pkgprotocol.FlaredHeartbeatResponse
// TunnelConfigResponse is an alias for FlaredTunnelConfigResponse.
type TunnelConfigResponse = pkgprotocol.FlaredTunnelConfigResponse
// RelayInfo is an alias for FlaredRelayInfo.
type RelayInfo = pkgprotocol.FlaredRelayInfo
// ProxyEntry is an alias for FlaredProxyEntry.
type ProxyEntry = pkgprotocol.FlaredProxyEntry
// ApplyLogPayload is an alias for ApplyLogPayload.
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
@@ -0,0 +1,135 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package flared
import (
"net/http"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// PostHeartbeat handles POST /tunnel/heartbeat.
// @Summary 上报 Tunnel 心跳
// @Description Tunnel 客户端定期上报运行状态与中继连接信息,返回活跃配置元数据与隧道设置
// @Tags openflare-tunnel
// @Accept json
// @Produce json
// @Security TunnelTokenAuth
// @Param body body flared.HeartbeatPayload true "心跳载荷"
// @Success 200 {object} response.Any{data=flared.HeartbeatResponse} "心跳响应"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Tunnel Token 无效"
// @Failure 403 {object} response.Any "节点类型不匹配"
// @Router /api/v1/tunnel/heartbeat [post]
func PostHeartbeat(c *gin.Context) {
var payload HeartbeatPayload
if !apiutil.BindJSON(c, &payload) {
return
}
authNode, ok := c.Get(ctxFlaredNodeKey)
if !ok {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
node, ok := authNode.(*model.OpenFlareNode)
if !ok {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
result, err := Heartbeat(c.Request.Context(), node, payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// GetActiveConfig handles GET /tunnel/config/active.
// @Summary 获取活跃隧道配置
// @Description 返回 Tunnel 客户端当前应应用的完整路由配置(含中继列表与代理定义)
// @Tags openflare-tunnel
// @Produce json
// @Security TunnelTokenAuth
// @Success 200 {object} response.Any{data=flared.TunnelConfigResponse} "隧道配置"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Tunnel Token 无效"
// @Failure 403 {object} response.Any "节点类型不匹配"
// @Router /api/v1/tunnel/config/active [get]
func GetActiveConfig(c *gin.Context) {
authNode, ok := c.Get(ctxFlaredNodeKey)
if !ok {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
node, ok := authNode.(*model.OpenFlareNode)
if !ok {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
config, err := GetTunnelConfig(c.Request.Context(), node)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(config))
}
// PostApplyLog handles POST /tunnel/apply-log.
// @Summary 上报 Tunnel 配置下发结果
// @Description Tunnel 客户端上报配置应用结果,服务端记录下发日志
// @Tags openflare-tunnel
// @Accept json
// @Produce json
// @Security TunnelTokenAuth
// @Param body body flared.ApplyLogPayload true "下发结果载荷"
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "下发日志记录"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Tunnel Token 无效"
// @Failure 403 {object} response.Any "节点类型不匹配"
// @Router /api/v1/tunnel/apply-log [post]
func PostApplyLog(c *gin.Context) {
var payload ApplyLogPayload
if !apiutil.BindJSON(c, &payload) {
return
}
if authNode, ok := c.Get(ctxFlaredNodeKey); ok {
if node, ok := authNode.(*model.OpenFlareNode); ok {
payload.NodeID = node.NodeID
}
}
log, err := ReportApplyLog(c.Request.Context(), payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(log))
}
// GetWebSocket handles GET /tunnel/ws.
// @Summary 升级 Tunnel WebSocket 连接
// @Description 将已认证的 Tunnel 客户端连接升级为 WebSocket 长连接,用于配置推送
// @Tags openflare-tunnel
// @Security TunnelTokenAuth
// @Failure 401 {object} response.Any "Tunnel Token 无效"
// @Failure 403 {object} response.Any "节点类型不匹配"
// @Router /api/v1/tunnel/ws [get]
func GetWebSocket(c *gin.Context) {
authNode, ok := c.Get(ctxFlaredNodeKey)
if !ok {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
node, ok := authNode.(*model.OpenFlareNode)
if !ok {
response.AbortUnauthorized(c, errTunnelTokenInvalid)
return
}
ofws.ServeFlared(c, node.NodeID)
}
@@ -0,0 +1,198 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package integration
import (
"context"
"net/http"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
ofnode "Wavelet/openflare/plugins/server/domain/fleet/node"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/testhelper"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupProtocolTestEnv(t *testing.T) (*gin.Engine, 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.SystemConfig{},
&model.OpenFlareApplyLog{},
&model.OpenFlareNodeSystemProfile{},
&model.OpenFlareHealthEvent{},
&model.ConfigVersion{},
))
db.SetDB(sqliteDB)
agent.ResetAuthCacheForTest()
testhelper.SetupLogStoresForTest(t)
engine := testhelper.NewTestGinEngine()
mountOpenFlareTestRoutes(engine)
cleanup := func() {
db.SetDB(nil)
agent.ResetAuthCacheForTest()
}
return engine, cleanup
}
func TestAgentRelayFlaredProtocol(t *testing.T) {
engine, cleanup := setupProtocolTestEnv(t)
defer cleanup()
ctx := context.Background()
t.Run("create edge node and heartbeat with X-Agent-Token", func(t *testing.T) {
edge, err := ofnode.CreateNode(ctx, ofnode.Input{
Name: "edge-1",
IP: "10.0.0.1",
})
require.NoError(t, err)
require.NotEmpty(t, edge.AccessToken)
assert.Equal(t, "edge_node", edge.NodeType)
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/nodes/heartbeat", map[string]any{
"name": "edge-1",
"ip": "203.0.113.10",
"version": "0.1.0",
}, map[string]string{
"X-Agent-Token": edge.AccessToken,
})
assert.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.NotNil(t, data["agent_settings"])
})
t.Run("create tunnel_relay node and relay heartbeat", func(t *testing.T) {
relayNode, err := ofnode.CreateNode(ctx, ofnode.Input{
Name: "relay-1",
NodeType: "tunnel_relay",
})
require.NoError(t, err)
require.NotEmpty(t, relayNode.AccessToken)
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/relay/heartbeat", map[string]any{
"version": "v0.1.0",
"frp_version": "0.61.0",
"relay_status": "healthy",
"name": "relay-1",
"ip": "203.0.113.20",
}, map[string]string{
"X-Agent-Token": relayNode.AccessToken,
})
assert.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
var heartbeatData struct {
RelayConfig map[string]any `json:"relay_config"`
RelaySettings map[string]any `json:"relay_settings"`
}
unmarshalAPIData(t, resp.Data, &heartbeatData)
assert.NotNil(t, heartbeatData.RelayConfig)
assert.NotNil(t, heartbeatData.RelaySettings)
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID)
require.NoError(t, err)
assert.Equal(t, "online", stored.Status)
assert.Equal(t, "healthy", stored.RelayStatus)
})
t.Run("create tunnel_client node and flared heartbeat with X-Tunnel-Token", func(t *testing.T) {
clientNode, err := ofnode.CreateNode(ctx, ofnode.Input{
Name: "client-1",
NodeType: "tunnel_client",
})
require.NoError(t, err)
require.NotEmpty(t, clientNode.AccessToken)
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/tunnel/heartbeat", map[string]any{
"client_version": "v0.2.0",
"frp_version": "0.61.0",
"tunnel_status": "running",
}, map[string]string{
"X-Tunnel-Token": clientNode.AccessToken,
})
assert.Equal(t, http.StatusOK, rec.Code)
requireAPIOK(t, rec)
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID)
require.NoError(t, err)
assert.Equal(t, "online", stored.Status)
assert.Equal(t, "v0.2.0", stored.Version)
})
t.Run("agent register with discovery token from options", func(t *testing.T) {
bootstrap, err := ofnode.GetBootstrapToken(ctx)
require.NoError(t, err)
require.NotEmpty(t, bootstrap.DiscoveryToken)
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/nodes/register", map[string]any{
"name": "discovered-edge",
"ip": "203.0.113.30",
"version": "0.2.0",
}, map[string]string{
"X-Agent-Token": bootstrap.DiscoveryToken,
})
assert.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
var registration agent.RegistrationResponse
unmarshalAPIData(t, resp.Data, &registration)
assert.NotEmpty(t, registration.NodeID)
assert.NotEmpty(t, registration.AccessToken)
assert.Equal(t, "discovered-edge", registration.Name)
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, registration.NodeID)
require.NoError(t, err)
assert.Equal(t, "online", stored.Status)
assert.Equal(t, registration.AccessToken, stored.AccessToken)
})
t.Run("POST agent apply-logs", func(t *testing.T) {
edge, err := ofnode.CreateNode(ctx, ofnode.Input{
Name: "edge-apply",
IP: "10.0.0.2",
})
require.NoError(t, err)
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/apply-logs", map[string]any{
"version": "20260618-001",
"result": "success",
"message": "apply ok",
}, map[string]string{
"X-Agent-Token": edge.AccessToken,
})
assert.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
var applyLog model.OpenFlareApplyLog
unmarshalAPIData(t, resp.Data, &applyLog)
assert.Equal(t, edge.NodeID, applyLog.NodeID)
assert.Equal(t, "success", applyLog.Result)
assert.Equal(t, "20260618-001", applyLog.Version)
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, edge.NodeID)
require.NoError(t, err)
assert.Equal(t, "online", stored.Status)
assert.Equal(t, "20260618-001", stored.CurrentVersion)
})
}
@@ -0,0 +1,178 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package integration
import (
"bytes"
"context"
"net/http"
"net/http/httptest"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/pkg/idgen"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
type statusPayload struct {
Version string `json:"version"`
ServerAddress string `json:"server_address"`
}
func setupAuthOptionIntegration(t *testing.T) (*gorm.DB, *gin.Engine) {
t.Helper()
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
t.Cleanup(cleanup)
require.NoError(t, dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeyCapLoginEnabled).
Update("value", "false").Error)
require.NoError(t, repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled))
runtimeconfig.SetSessionSecret("test_openflare_session_secret")
store := cookie.NewStore([]byte("test_openflare_session_secret"))
r := testhelper.NewTestGinEngine(sessions.Sessions("test_openflare_session", store))
mountOpenFlareTestRoutes(r)
return dbConn, r
}
func seedUser(t *testing.T, dbConn *gorm.DB, username, password string, isAdmin bool) *model.User {
t.Helper()
user := &model.User{
ID: idgen.NextUint64ID(),
Username: username,
Nickname: username,
Email: username + "@openflare.test",
IsActive: true,
IsAdmin: isAdmin,
}
require.NoError(t, user.SetEncryptedPassword(password))
require.NoError(t, dbConn.Create(user).Error)
return user
}
func seedUserWithAccessToken(t *testing.T, dbConn *gorm.DB, username, password string, isAdmin bool) string {
t.Helper()
user := seedUser(t, dbConn, username, password, isAdmin)
token, err := model.GenerateTokenString()
require.NoError(t, err)
tokenRecord := model.AccessToken{
UserID: user.ID,
Name: username + "-integration-token",
TokenHash: model.HashToken(token),
MaskedToken: model.MaskTokenString(token),
IsAdmin: isAdmin,
}
require.NoError(t, dbConn.Create(&tokenRecord).Error)
return token
}
func TestGETStatusReturnsSuccessEnvelope(t *testing.T) {
_, r := setupAuthOptionIntegration(t)
w := performJSONRequest(t, r, http.MethodGet, apiPath("/status"), nil, nil)
assert.Equal(t, http.StatusOK, w.Code)
resp := requireAPIOK(t, w)
var status statusPayload
unmarshalAPIData(t, resp.Data, &status)
assert.NotEmpty(t, status.Version)
}
func TestGETOptionRequiresAdminAuth(t *testing.T) {
dbConn, r := setupAuthOptionIntegration(t)
commonToken := seedUserWithAccessToken(t, dbConn, "commonuser", "password123", false)
adminToken := seedUserWithAccessToken(t, dbConn, "adminuser", "password123", true)
t.Run("unauthenticated", func(t *testing.T) {
t.Skip("console auth is owned by Wavelet auth plugin")
})
t.Run("non-admin user forbidden", func(t *testing.T) {
t.Skip("console auth is owned by Wavelet auth plugin")
_ = commonToken
})
t.Run("admin user allowed", func(t *testing.T) {
w := performJSONRequest(t, r, http.MethodGet, apiPath("/option/"), nil, adminAuthHeaders(adminToken))
assert.Equal(t, http.StatusOK, w.Code)
requireAPIOK(t, w)
})
}
func TestPOSTOptionUpdateRejectsInvalidParams(t *testing.T) {
dbConn, r := setupAuthOptionIntegration(t)
adminToken := seedUserWithAccessToken(t, dbConn, "adminuser", "password123", true)
req := httptest.NewRequest(http.MethodPost, apiPath("/option/update"), bytes.NewReader([]byte("{invalid")))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Access-Token", adminToken)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
resp := decodeAPIResponse(t, w)
assert.NotEmpty(t, resp.ErrorMsg)
}
func TestGETNodesWithAccessToken(t *testing.T) {
dbConn, r := setupAuthOptionIntegration(t)
require.NoError(t, dbConn.AutoMigrate(&model.OpenFlareNode{}))
adminToken := seedUserWithAccessToken(t, dbConn, "admin", "password123", true)
w := performJSONRequest(t, r, http.MethodGet, apiPath("/nodes/"), nil, adminAuthHeaders(adminToken))
assert.Equal(t, http.StatusOK, w.Code)
requireAPIOK(t, w)
}
func TestOptionUpdatePersistsAndReflectsInStatus(t *testing.T) {
dbConn, r := setupAuthOptionIntegration(t)
adminToken := seedUserWithAccessToken(t, dbConn, "admin", "password123", true)
updateResp := performJSONRequest(t, r, http.MethodPost, apiPath("/option/update"), map[string]string{
"key": model.ConfigKeyServerAddress,
"value": "https://hotreload.openflare.test",
}, adminAuthHeaders(adminToken))
assert.Equal(t, http.StatusOK, updateResp.Code)
requireAPIOK(t, updateResp)
statusAfter := getStatusServerAddress(t, r, nil)
assert.Equal(t, "https://hotreload.openflare.test", statusAfter)
// 验证已持久化到 SystemConfig
ctx := context.Background()
saved, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
require.NoError(t, err)
assert.Equal(t, "https://hotreload.openflare.test", saved.Value)
}
func getStatusServerAddress(t *testing.T, r http.Handler, headers map[string]string) string {
t.Helper()
w := performJSONRequest(t, r, http.MethodGet, apiPath("/status"), nil, headers)
require.Equal(t, http.StatusOK, w.Code)
resp := requireAPIOK(t, w)
var status statusPayload
unmarshalAPIData(t, resp.Data, &status)
return status.ServerAddress
}
@@ -0,0 +1,288 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package integration
import (
"context"
"net/http"
"testing"
"time"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/testhelper"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
const (
adminUserID = uint64(1001)
adminUsername = "openflare-admin"
)
type adminSeed struct {
User model.User
Token string
TokenHash string
}
func setupCoreChainTest(t *testing.T) (*gin.Engine, adminSeed, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.User{},
&model.AccessToken{},
&model.Origin{},
&model.ProxyRoute{},
&model.ConfigVersion{},
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFRuleGroupBinding{},
&model.OpenFlareWAFIPGroup{},
&model.OpenFlareNode{},
&model.SystemConfig{},
&model.OpenFlareApplyLog{},
&model.Zone{},
&model.ZoneDomain{},
))
db.SetDB(sqliteDB)
agent.ResetAuthCacheForTest()
seed, err := seedAdminWithAccessToken(sqliteDB)
require.NoError(t, err)
engine := testhelper.NewTestGinEngine()
mountOpenFlareTestRoutes(engine)
cleanup := func() {
db.SetDB(nil)
agent.ResetAuthCacheForTest()
}
return engine, seed, cleanup
}
func seedAdminWithAccessToken(conn *gorm.DB) (adminSeed, error) {
now := time.Now().UTC()
admin := model.User{
ID: adminUserID,
Username: adminUsername,
Nickname: "OpenFlare Admin",
IsActive: true,
IsAdmin: true,
LastLoginAt: now,
}
if err := conn.Create(&admin).Error; err != nil {
return adminSeed{}, err
}
token, err := model.GenerateTokenString()
if err != nil {
return adminSeed{}, err
}
tokenHash := model.HashToken(token)
tokenRecord := model.AccessToken{
UserID: adminUserID,
Name: "integration-admin-token",
TokenHash: tokenHash,
MaskedToken: model.MaskTokenString(token),
IsAdmin: true,
}
if err := conn.Create(&tokenRecord).Error; err != nil {
return adminSeed{}, err
}
return adminSeed{
User: admin,
Token: token,
TokenHash: tokenHash,
}, nil
}
func TestCoreChainMigrationFlow(t *testing.T) {
engine, seed, cleanup := setupCoreChainTest(t)
defer cleanup()
var (
originID uint
proxyRouteID uint
configVersion string
configChecksum string
nodeID uint
nodePublicID string
agentToken string
)
t.Run("create origin", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/origins/"), map[string]any{
"name": "Primary Origin",
"address": "origin.core-chain.internal",
"remark": "integration upstream",
}, map[string]string{
"X-Access-Token": seed.Token,
})
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
originID = uint(data["id"].(float64))
assert.NotZero(t, originID)
assert.Equal(t, "Primary Origin", data["name"])
assert.Equal(t, "origin.core-chain.internal", data["address"])
})
t.Run("create proxy route linked to origin", func(t *testing.T) {
// Create Zone and ZoneDomain directly in the DB
zone := model.Zone{Domain: "example.com"}
require.NoError(t, db.DB(context.Background()).Create(&zone).Error)
zoneDomain := model.ZoneDomain{
ZoneID: zone.ID,
Domain: "core-chain.example.com",
}
require.NoError(t, db.DB(context.Background()).Create(&zoneDomain).Error)
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{
"site_name": "core-chain-site",
"zone_domain_ids": []uint{zoneDomain.ID},
"origin_id": originID,
"origin_scheme": "http",
"origin_port": "8080",
"enabled": true,
}, map[string]string{
"X-Access-Token": seed.Token,
})
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
proxyRouteID = uint(data["id"].(float64))
assert.NotZero(t, proxyRouteID)
assert.Equal(t, "core-chain-site", data["site_name"])
assert.NotEmpty(t, data["zone_domains"])
zoneDomains := data["zone_domains"].([]any)
assert.Len(t, zoneDomains, 1)
assert.Equal(t, "core-chain.example.com", zoneDomains[0].(map[string]any)["domain"])
assert.InDelta(t, float64(originID), data["origin_id"], 1e-9)
assert.Equal(t, "http://origin.core-chain.internal:8080", data["origin_url"])
})
t.Run("publish config version", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/config-versions/publish"), nil, map[string]string{
"X-Access-Token": seed.Token,
})
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
configVersion, _ = data["version"].(string)
configChecksum, _ = data["checksum"].(string)
assert.NotEmpty(t, configVersion)
assert.NotEmpty(t, configChecksum)
assert.Equal(t, true, data["is_active"])
activeRec := performJSONRequest(t, engine, http.MethodGet, apiPath("/config-versions/active"), nil, map[string]string{
"X-Access-Token": seed.Token,
})
require.Equal(t, http.StatusOK, activeRec.Code)
activeResp := requireAPIOK(t, activeRec)
activeData := unmarshalAPIMap(t, activeResp.Data)
assert.Equal(t, configVersion, activeData["version"])
assert.Equal(t, configChecksum, activeData["checksum"])
})
t.Run("create node", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/nodes/"), map[string]any{
"name": "edge-core-chain",
"ip": "10.10.0.1",
"auto_update_enabled": true,
}, map[string]string{
"X-Access-Token": seed.Token,
})
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
nodeID = uint(data["id"].(float64))
nodePublicID, _ = data["node_id"].(string)
agentToken, _ = data["access_token"].(string)
assert.NotZero(t, nodeID)
assert.NotEmpty(t, nodePublicID)
assert.Len(t, agentToken, 32)
})
t.Run("create apply log for node", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/apply-logs", map[string]any{
"version": configVersion,
"result": "success",
"message": "config applied",
"checksum": configChecksum,
"main_config_checksum": "main-checksum",
"route_config_checksum": "route-checksum",
"support_file_count": 2,
}, map[string]string{
"X-Agent-Token": agentToken,
})
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.Equal(t, nodePublicID, data["node_id"])
assert.Equal(t, configVersion, data["version"])
assert.Equal(t, "success", data["result"])
assert.Equal(t, configChecksum, data["checksum"])
})
t.Run("verify apply log listing and node metadata", func(t *testing.T) {
listRec := performJSONRequest(
t,
engine,
http.MethodGet,
apiPath("/apply-logs/?node_id="+nodePublicID+"&pageNo=1&pageSize=10"),
nil,
map[string]string{
"X-Access-Token": seed.Token,
},
)
require.Equal(t, http.StatusOK, listRec.Code)
listResp := requireAPIOK(t, listRec)
listData := unmarshalAPIMap(t, listResp.Data)
assert.InDelta(t, float64(1), listData["total"], 1e-9)
rows, ok := listData["rows"].([]any)
require.True(t, ok)
require.Len(t, rows, 1)
row, ok := rows[0].(map[string]any)
require.True(t, ok)
assert.Equal(t, nodePublicID, row["node_id"])
assert.Equal(t, configVersion, row["version"])
assert.Equal(t, "success", row["result"])
nodeRec := performJSONRequest(t, engine, http.MethodGet, apiPath("/nodes/"), nil, map[string]string{
"X-Access-Token": seed.Token,
})
require.Equal(t, http.StatusOK, nodeRec.Code)
nodeResp := requireAPIOK(t, nodeRec)
nodes := unmarshalAPISlice(t, nodeResp.Data)
require.Len(t, nodes, 1)
nodeView, ok := nodes[0].(map[string]any)
require.True(t, ok)
assert.InDelta(t, float64(nodeID), nodeView["id"], 1e-9)
assert.Equal(t, nodePublicID, nodeView["node_id"])
assert.Equal(t, "success", nodeView["latest_apply_result"])
assert.Equal(t, configChecksum, nodeView["latest_apply_checksum"])
assert.InDelta(t, float64(2), nodeView["latest_support_file_count"], 1e-9)
})
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package integration
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server"
ofrouter "Wavelet/openflare/plugins/server/httpapi"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func decodeAPIResponse(t *testing.T, rec *httptest.ResponseRecorder) response.Any {
t.Helper()
var resp response.Any
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
return resp
}
func requireAPIOK(t *testing.T, rec *httptest.ResponseRecorder) response.Any {
t.Helper()
resp := decodeAPIResponse(t, rec)
require.Empty(t, resp.ErrorMsg, "unexpected API error: %s", resp.ErrorMsg)
return resp
}
func unmarshalAPIData(t *testing.T, data any, target any) {
t.Helper()
payload, err := json.Marshal(data)
require.NoError(t, err)
require.NoError(t, json.Unmarshal(payload, target))
}
func unmarshalAPIMap(t *testing.T, data any) map[string]any {
t.Helper()
var result map[string]any
unmarshalAPIData(t, data, &result)
return result
}
func unmarshalAPISlice(t *testing.T, data any) []any {
t.Helper()
var result []any
unmarshalAPIData(t, data, &result)
return result
}
// mountOpenFlareTestRoutes 复刻 driver_http 的挂载方式:先由 server 插件经内核
// 路由注册表声明路由,再把每条 (方法, 路径, 中间件+处理链) 挂到测试引擎上。
func mountOpenFlareTestRoutes(engine *gin.Engine) {
ctx := core.NewContext(context.Background())
core.Provide[contracts.AuthService](ctx, testhelper.StubAuth{})
if err := server.New().Apply(ctx); err != nil {
panic(err)
}
for _, rd := range ctx.Router().Routes() {
chain := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers))
for _, item := range append(append([]any{}, rd.Middlewares...), rd.Handlers...) {
chain = append(chain, toGinHandler(item))
}
engine.Handle(rd.Method, rd.Path, chain...)
}
}
// toGinHandler 与 driver_http 接受的处理函数形态保持一致。
func toGinHandler(item any) gin.HandlerFunc {
switch fn := item.(type) {
case gin.HandlerFunc:
return fn
case func(*gin.Context):
return gin.HandlerFunc(fn)
default:
panic("unexpected handler type")
}
}
func apiPath(subpath string) string {
return ofrouter.V1BasePath + subpath
}
func performJSONRequest(
t *testing.T,
engine http.Handler,
method, path string,
body any,
headers map[string]string,
) *httptest.ResponseRecorder {
t.Helper()
var payload []byte
if body != nil {
var err error
payload, err = json.Marshal(body)
require.NoError(t, err)
}
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
for key, value := range headers {
req.Header.Set(key, value)
}
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
return rec
}
func adminAuthHeaders(token string) map[string]string {
return map[string]string{
"X-Access-Token": token,
}
}
@@ -0,0 +1,374 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package integration
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"fmt"
"math/big"
"net/http"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
"Wavelet/openflare/plugins/server/kernel/testhelper"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupSecurityTest(t *testing.T) (*gin.Engine, adminSeed, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.User{},
&model.AccessToken{},
&model.Origin{},
&model.ProxyRoute{},
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFRuleGroupBinding{},
&model.OpenFlareWAFIPGroup{},
&model.TLSCertificate{},
&model.Zone{},
&model.ZoneDomain{},
&model.DNSAccount{},
&model.AcmeAccount{},
&model.SystemConfig{},
))
db.SetDB(sqliteDB)
seed, err := seedAdminWithAccessToken(sqliteDB)
require.NoError(t, err)
previous := runtimeconfig.Get()
runtimeconfig.SetSessionSecret("test_session_secret_for_security_integration")
engine := testhelper.NewTestGinEngine()
mountOpenFlareTestRoutes(engine)
cleanup := func() {
runtimeconfig.Set(previous)
db.SetDB(nil)
}
return engine, seed, cleanup
}
func generateSelfSignedCertificatePair(t *testing.T, dnsNames []string) (string, string) {
t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
template := &x509.Certificate{
SerialNumber: big.NewInt(time.Now().UnixNano()),
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},
}
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
require.NoError(t, 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)
}
func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
engine, seed, cleanup := setupSecurityTest(t)
defer cleanup()
var (
ruleGroupID uint
ipGroupID uint
proxyRouteID uint
certID uint
domainID uint
dnsAccountID uint
)
t.Run("WAF rule group create", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/rule-groups"), map[string]any{
"name": "edge-security",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
ruleGroupID = uint(data["id"].(float64))
assert.NotZero(t, ruleGroupID)
assert.Equal(t, "edge-security", data["name"])
assert.Equal(t, false, data["is_global"])
assert.InDelta(t, float64(1), data["revision"], 1e-9)
assert.NotNil(t, data["graph"])
})
t.Run("WAF rule group list includes global and custom groups", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodGet, apiPath("/waf/rule-groups"), nil, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
groups := unmarshalAPISlice(t, resp.Data)
require.GreaterOrEqual(t, len(groups), 2)
foundCustom := false
foundGlobal := false
for _, item := range groups {
group, ok := item.(map[string]any)
require.True(t, ok)
if group["is_global"] == true {
foundGlobal = true
}
if uint(group["id"].(float64)) == ruleGroupID {
foundCustom = true
assert.Equal(t, "edge-security", group["name"])
}
}
assert.True(t, foundGlobal)
assert.True(t, foundCustom)
})
t.Run("WAF rule group get detail", func(t *testing.T) {
rec := performJSONRequest(
t,
engine,
http.MethodGet,
fmt.Sprintf("%s/waf/rule-groups/%d", apiPath(""), ruleGroupID),
nil,
adminAuthHeaders(seed.Token),
)
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.InDelta(t, float64(ruleGroupID), data["id"], 1e-9)
assert.Equal(t, "edge-security", data["name"])
})
t.Run("WAF rule group update", func(t *testing.T) {
rec := performJSONRequest(
t,
engine,
http.MethodPost,
fmt.Sprintf("%s/waf/rule-groups/%d/meta", apiPath(""), ruleGroupID),
map[string]any{
"name": "edge-security-updated", "enabled": true,
},
adminAuthHeaders(seed.Token),
)
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.Equal(t, "edge-security-updated", data["name"])
assert.Equal(t, true, data["enabled"])
})
t.Run("WAF IP group create", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/ip-groups"), map[string]any{
"name": "blocked-ips",
"type": "manual",
"enabled": true,
"ip_list": []string{"203.0.113.0/24", "198.51.100.10"},
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
ipGroupID = uint(data["id"].(float64))
assert.NotZero(t, ipGroupID)
assert.Equal(t, "blocked-ips", data["name"])
assert.Equal(t, "manual", data["type"])
})
t.Run("create proxy route for WAF binding", func(t *testing.T) {
// Create Zone and ZoneDomain directly in the DB
routeZone := model.Zone{Domain: "example-route.com"}
require.NoError(t, db.DB(context.Background()).Create(&routeZone).Error)
routeZoneDomain := model.ZoneDomain{
ZoneID: routeZone.ID,
Domain: "route.example-route.com",
}
require.NoError(t, db.DB(context.Background()).Create(&routeZoneDomain).Error)
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{
"site_name": "security-site",
"zone_domain_ids": []uint{routeZoneDomain.ID},
"origin_url": "http://origin.security.internal:8080",
"enabled": true,
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
proxyRouteID = uint(data["id"].(float64))
assert.NotZero(t, proxyRouteID)
assert.NotEmpty(t, data["zone_domains"])
zoneDomains := data["zone_domains"].([]any)
assert.Len(t, zoneDomains, 1)
assert.Equal(t, "route.example-route.com", zoneDomains[0].(map[string]any)["domain"])
})
t.Run("bind WAF rule group to proxy route", func(t *testing.T) {
rec := performJSONRequest(
t,
engine,
http.MethodPost,
fmt.Sprintf("%s/waf/sites/%d/rule-groups", apiPath(""), proxyRouteID),
map[string]any{
"ids": []uint{ruleGroupID},
},
adminAuthHeaders(seed.Token),
)
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.InDelta(t, float64(proxyRouteID), data["route_id"], 1e-9)
appliedIDs, ok := data["applied_ids"].([]any)
require.True(t, ok)
require.Len(t, appliedIDs, 1)
assert.InDelta(t, float64(ruleGroupID), appliedIDs[0], 1e-9)
})
t.Run("verify site rule groups binding", func(t *testing.T) {
rec := performJSONRequest(
t,
engine,
http.MethodGet,
fmt.Sprintf("%s/waf/sites/%d/rule-groups", apiPath(""), proxyRouteID),
nil,
adminAuthHeaders(seed.Token),
)
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.NotNil(t, data["global_rule_group"])
appliedGroups, ok := data["applied_rule_groups"].([]any)
require.True(t, ok)
require.Len(t, appliedGroups, 1)
group, ok := appliedGroups[0].(map[string]any)
require.True(t, ok)
assert.InDelta(t, float64(ruleGroupID), group["id"], 1e-9)
})
t.Run("create TLS certificate with PEM", func(t *testing.T) {
certPEM, keyPEM := generateSelfSignedCertificatePair(t, []string{"security.example.com"})
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/tls-certificates/"), map[string]any{
"name": "security-cert",
"cert_pem": certPEM,
"key_pem": keyPEM,
"remark": "self-signed integration cert",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
certID = uint(data["id"].(float64))
assert.NotZero(t, certID)
assert.Equal(t, "security-cert", data["name"])
assert.Equal(t, "upload", data["provider"])
})
t.Run("create Zone domain", func(t *testing.T) {
zoneRec := performJSONRequest(t, engine, http.MethodPost, apiPath("/zones/"), map[string]any{
"domain": "example.com",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, zoneRec.Code)
zoneData := unmarshalAPIMap(t, requireAPIOK(t, zoneRec).Data)
zoneID := uint(zoneData["id"].(float64))
rec := performJSONRequest(t, engine, http.MethodPost, fmt.Sprintf("%s/zones/%d/domains", apiPath(""), zoneID), map[string]any{
"domain": "*.example.com",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusBadRequest, rec.Code)
errResp := decodeAPIResponse(t, rec)
assert.NotEmpty(t, errResp.ErrorMsg)
rec = performJSONRequest(t, engine, http.MethodPost, fmt.Sprintf("%s/zones/%d/domains", apiPath(""), zoneID), map[string]any{
"domain": "security.example.com", "cert_id": certID,
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
domainID = uint(data["id"].(float64))
assert.NotZero(t, domainID)
assert.Equal(t, "security.example.com", data["domain"])
assert.InDelta(t, float64(certID), data["cert_id"], 1e-9)
})
t.Run("create DNS account", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/dns-accounts/"), map[string]any{
"name": "cloudflare-dns",
"type": "cloudflare",
"authorization": "test-api-token-value",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
dnsAccountID = uint(data["id"].(float64))
assert.NotZero(t, dnsAccountID)
assert.Equal(t, "cloudflare-dns", data["name"])
assert.Equal(t, "cloudflare", data["type"])
// API 响应会脱敏 authorization,不应回显明文凭证。
if auth, ok := data["authorization"]; ok {
assert.NotEqual(t, "test-api-token-value", auth)
}
})
t.Run("WAF rule group delete", func(t *testing.T) {
rec := performJSONRequest(
t,
engine,
http.MethodPost,
fmt.Sprintf("%s/waf/rule-groups/%d/delete", apiPath(""), ruleGroupID),
nil,
adminAuthHeaders(seed.Token),
)
require.Equal(t, http.StatusOK, rec.Code)
requireAPIOK(t, rec)
detailRec := performJSONRequest(
t,
engine,
http.MethodGet,
fmt.Sprintf("%s/waf/rule-groups/%d", apiPath(""), ruleGroupID),
nil,
adminAuthHeaders(seed.Token),
)
require.Equal(t, http.StatusNotFound, detailRec.Code)
detailResp := decodeAPIResponse(t, detailRec)
assert.NotEmpty(t, detailResp.ErrorMsg)
})
_ = ipGroupID
_ = domainID
_ = dnsAccountID
}
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package node defines node validation and management error messages.
package node
const (
errNodeNameRequired = "节点名不能为空"
errNodeIPTooLong = "节点 IP 不能超过 64 个字符"
errNodeIPInvalid = "节点 IP 格式无效"
errNodeIPManualRequired = "锁定节点 IP 时必须填写节点 IP"
errNodeGeoNameTooLong = "节点位置名不能超过 128 个字符"
errNodeGeoCoordinateMismatch = "地图坐标必须同时填写纬度和经度"
errNodeGeoLatitudeInvalid = "纬度必须在 -90 到 90 之间"
errNodeGeoLongitudeInvalid = "经度必须在 -180 到 180 之间"
errNodeIDConflict = "节点标识生成冲突,请重试"
errNodeNotFound = "节点不存在"
errNodeForceSyncFailed = "节点不在线或通过 WebSocket 发送同步指令失败"
errNoActiveConfigVersion = "当前没有激活版本"
errAgentPreviewTagInvalid = "指定版本不是 preview 发布"
errAgentStableTagInvalid = "正式版更新不能选择 preview 发布"
)
@@ -0,0 +1,414 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package node
import (
"context"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/kernel/model"
)
const (
nodeStatusOnline = "online"
nodeStatusOffline = "offline"
nodeStatusPending = "pending"
openrestyStatusHealthy = "healthy"
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
githubReleasesAPIBase = "https://api.github.com/repos/%s/releases"
nodeTypeTunnelRelay = "tunnel_relay"
nodeTypeTunnelClient = "tunnel_client"
nodeTypeEdgeNode = "edge_node"
nodeTokenByteLength = 16
maxNodeIPLength = 64
maxNodeGeoNameLength = 128
)
type releaseChannel string
const (
releaseChannelStable releaseChannel = "stable"
releaseChannelPreview releaseChannel = "preview"
)
var releaseHTTPClient = &http.Client{Timeout: 30 * time.Second}
type githubReleaseResponse struct {
TagName string `json:"tag_name"`
Body string `json:"body"`
HTMLURL string `json:"html_url"`
PublishedAt string `json:"published_at"`
Prerelease bool `json:"prerelease"`
Draft bool `json:"draft"`
}
func newRandomToken() (string, error) {
buf := make([]byte, nodeTokenByteLength)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return hex.EncodeToString(buf), nil
}
func tokenEqual(got, want string) bool {
sumGot := sha256.Sum256([]byte(got))
sumWant := sha256.Sum256([]byte(want))
return subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) == 1
}
func newServerNodeID() (string, error) {
token, err := newRandomToken()
if err != nil {
return "", err
}
return "node-" + token, nil
}
func normalizeNodeType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case nodeTypeTunnelRelay:
return nodeTypeTunnelRelay
case nodeTypeTunnelClient:
return nodeTypeTunnelClient
default:
return nodeTypeEdgeNode
}
}
func normalizeRelayPort(port int, defaultPort int) int {
if port <= 0 || port > 65535 {
return defaultPort
}
return port
}
func normalizeReleaseChannel(channel string) releaseChannel {
switch strings.ToLower(strings.TrimSpace(channel)) {
case string(releaseChannelPreview):
return releaseChannelPreview
default:
return releaseChannelStable
}
}
func (channel releaseChannel) String() string {
if channel == releaseChannelPreview {
return string(releaseChannelPreview)
}
return string(releaseChannelStable)
}
func normalizeOpenrestyStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case openrestyStatusHealthy:
return openrestyStatusHealthy
case openrestyStatusUnhealthy:
return openrestyStatusUnhealthy
default:
return openrestyStatusUnknown
}
}
func cloneCoordinate(value *float64) *float64 {
if value == nil {
return nil
}
cloned := *value
return &cloned
}
func resolveNodeIPManualOverride(input Input, existing *model.OpenFlareNode, normalizedIP string) bool {
if input.IPManualOverride != nil {
return *input.IPManualOverride
}
if existing == nil {
return strings.TrimSpace(normalizedIP) != ""
}
if existing.IPManualOverride {
return true
}
return strings.TrimSpace(normalizedIP) != "" && strings.TrimSpace(normalizedIP) != strings.TrimSpace(existing.IP)
}
func normalizeNodeInput(input Input) (string, string, string, *float64, *float64, bool, error) {
name := strings.TrimSpace(input.Name)
ip := strings.TrimSpace(input.IP)
geoName := strings.TrimSpace(input.GeoName)
if err := validateNodeIPInput(input, ip); err != nil {
return "", "", "", nil, nil, false, err
}
if len(geoName) > maxNodeGeoNameLength {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoNameTooLong)
}
geoLatitude := cloneCoordinate(input.GeoLatitude)
geoLongitude := cloneCoordinate(input.GeoLongitude)
if err := validateNodeGeoCoordinates(geoLatitude, geoLongitude); err != nil {
return "", "", "", nil, nil, false, err
}
manualOverride := input.GeoManualOverride || geoName != "" || geoLatitude != nil || geoLongitude != nil
if !manualOverride || (geoLatitude == nil && geoLongitude == nil && geoName == "") {
return name, ip, "", nil, nil, false, nil
}
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
}
func validateNodeIPInput(input Input, ip string) error {
if len(ip) > maxNodeIPLength {
return fmt.Errorf("%s", errNodeIPTooLong)
}
if ip != "" && net.ParseIP(ip) == nil {
return fmt.Errorf("%s", errNodeIPInvalid)
}
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
return fmt.Errorf("%s", errNodeIPManualRequired)
}
return nil
}
func validateNodeGeoCoordinates(geoLatitude, geoLongitude *float64) error {
if (geoLatitude == nil) != (geoLongitude == nil) {
return fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
}
if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) {
return fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
}
if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) {
return fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
}
return nil
}
func computeNodeStatus(node *model.OpenFlareNode) string {
if node == nil {
return nodeStatusOffline
}
if node.LastSeenAt == nil || node.LastSeenAt.IsZero() {
return nodeStatusPending
}
// 默认离线阈值 60 秒(与 node_offline_threshold 默认一致),避免在这里读取配置
// 实际阈值会在需要精确判断的地方通过 getNodeOfflineThreshold 读取
threshold := 60 * time.Second
if time.Since(*node.LastSeenAt) > threshold {
return nodeStatusOffline
}
return nodeStatusOnline
}
func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
if node == nil {
return time.Time{}
}
nodeType := strings.TrimSpace(node.NodeType)
if nodeType == "" {
nodeType = nodeTypeEdgeNode
}
if nodeType == nodeTypeTunnelRelay && ofws.IsRelayConnected(node.NodeID) {
return ofws.RelayWSConnectedLastSeenValue
}
if nodeType == nodeTypeTunnelClient && ofws.IsFlaredConnected(node.NodeID) {
return ofws.FlaredWSConnectedLastSeenValue
}
if ofws.IsAgentConnected(node.NodeID) {
return ofws.AgentWSConnectedLastSeenValue
}
if node.LastSeenAt == nil {
return time.Time{}
}
return *node.LastSeenAt
}
func buildNodeView(node *model.OpenFlareNode) *View {
if node == nil {
return nil
}
status := computeNodeStatus(node)
view := &View{
ID: node.ID,
NodeID: node.NodeID,
Name: node.Name,
IP: node.IP,
IPManualOverride: node.IPManualOverride,
GeoName: strings.TrimSpace(node.GeoName),
GeoLatitude: node.GeoLatitude,
GeoLongitude: node.GeoLongitude,
GeoManualOverride: node.GeoManualOverride,
AccessToken: node.AccessToken,
UpdateChannel: strings.TrimSpace(node.UpdateChannel),
UpdateTag: strings.TrimSpace(node.UpdateTag),
RestartOpenrestyRequested: node.RestartOpenrestyRequested,
Version: node.Version,
ExtVersion: node.ExtVersion,
OpenrestyStatus: normalizeOpenrestyStatus(node.OpenrestyStatus),
OpenrestyMessage: strings.TrimSpace(node.OpenrestyMessage),
Status: status,
CurrentVersion: node.CurrentVersion,
LastSeenAt: nodeViewLastSeenAt(node),
LastError: node.LastError,
CreatedAt: node.CreatedAt,
UpdatedAt: node.UpdatedAt,
AutoUpdateEnabled: node.AutoUpdateEnabled,
UpdateRequested: node.UpdateRequested,
NodeType: node.NodeType,
RelayBindPort: node.RelayBindPort,
RelayVhostHTTPPort: node.RelayVhostHTTPPort,
RelayAgentAccessAddr: node.RelayAgentAccessAddr,
RelayClientAccessAddr: node.RelayClientAccessAddr,
RelayClientProxyURL: node.RelayClientProxyURL,
RelayStatus: node.RelayStatus,
RelayWebServerEnabled: node.RelayWebServerEnabled,
}
if view.UpdateChannel == "" {
view.UpdateChannel = releaseChannelStable.String()
}
if view.NodeType == "" {
view.NodeType = nodeTypeEdgeNode
}
return view
}
func buildNodeAgentReleaseView(node *model.OpenFlareNode, release *githubReleaseResponse, channel releaseChannel) *AgentReleaseInfo {
currentVersion := strings.TrimSpace(node.Version)
view := &AgentReleaseInfo{
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 isVersionNewer(current string, latest string) bool {
return compareVersions(current, latest) < 0
}
func fetchLatestGitHubRelease(ctx context.Context, repo string, channel releaseChannel) (*githubReleaseResponse, error) {
switch normalizeReleaseChannel(string(channel)) {
case releaseChannelPreview:
return fetchLatestPreviewGitHubRelease(ctx, repo)
case releaseChannelStable:
return fetchLatestStableGitHubRelease(ctx, repo)
default:
return fetchLatestStableGitHubRelease(ctx, repo)
}
}
func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) {
url := fmt.Sprintf(githubReleasesAPIBase+"/latest", strings.TrimSpace(repo))
req, err := newGitHubReleaseRequest(ctx, url)
if err != nil {
return nil, errors.New("创建更新请求失败")
}
resp, err := releaseHTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("获取最新版本失败: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
}
return decodeGitHubRelease(resp.Body)
}
func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) {
url := fmt.Sprintf(githubReleasesAPIBase+"?per_page=20", strings.TrimSpace(repo))
req, err := newGitHubReleaseRequest(ctx, url)
if err != nil {
return nil, errors.New("创建更新请求失败")
}
resp, err := releaseHTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("获取 preview 版本失败: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
}
var releases []githubReleaseResponse
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
return nil, errors.New("解析 preview 版本信息失败")
}
for _, release := range releases {
if release.Draft || !release.Prerelease {
continue
}
releaseCopy := release
return &releaseCopy, nil
}
return nil, errors.New("当前没有可用的 preview 发布")
}
func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*githubReleaseResponse, error) {
tag = strings.TrimSpace(tag)
if tag == "" {
return nil, errors.New("缺少发布版本号")
}
url := fmt.Sprintf(githubReleasesAPIBase+"/tags/%s", strings.TrimSpace(repo), tag)
req, err := newGitHubReleaseRequest(ctx, url)
if err != nil {
return nil, errors.New("创建更新请求失败")
}
resp, err := releaseHTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("获取指定版本失败: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode == http.StatusNotFound {
return nil, fmt.Errorf("未找到指定版本: %s", tag)
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
}
return decodeGitHubRelease(resp.Body)
}
func newGitHubReleaseRequest(ctx context.Context, url string) (*http.Request, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
req.Header.Set("User-Agent", "OpenFlare-Server")
return req, nil
}
func decodeGitHubRelease(reader io.Reader) (*githubReleaseResponse, error) {
var release githubReleaseResponse
if err := json.NewDecoder(reader).Decode(&release); err != nil {
return nil, errors.New("解析版本信息失败")
}
return &release, nil
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
@@ -0,0 +1,428 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package node
import (
"context"
"errors"
"fmt"
"strings"
"time"
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/domain/observability"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/logger"
"gorm.io/gorm"
)
const (
defaultRelayBindPort = 7000
defaultRelayVhostHTTPPort = 8080
)
// getAgentUpdateRepo 从 SystemConfig 读取 Agent 更新仓库配置
func getAgentUpdateRepo(ctx context.Context) string {
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentUpdateRepo)
if err != nil || strings.TrimSpace(config.Value) == "" {
return "Rain-kl/OpenFlare" // 默认值
}
return strings.TrimSpace(config.Value)
}
// Input is the create/update node payload.
type Input struct {
Name string `json:"name"`
IP string `json:"ip"`
IPManualOverride *bool `json:"ip_manual_override"`
AutoUpdateEnabled bool `json:"auto_update_enabled"`
GeoName string `json:"geo_name"`
GeoLatitude *float64 `json:"geo_latitude"`
GeoLongitude *float64 `json:"geo_longitude"`
GeoManualOverride bool `json:"geo_manual_override"`
NodeType string `json:"node_type"`
RelayBindPort int `json:"relay_bind_port"`
RelayVhostHTTPPort int `json:"relay_vhost_http_port"`
RelayAgentAccessAddr string `json:"relay_agent_access_addr"`
RelayClientAccessAddr string `json:"relay_client_access_addr"`
RelayClientProxyURL string `json:"relay_client_proxy_url"`
RelayWebServerEnabled bool `json:"relay_web_server_enabled"`
}
// AgentUpdateInput requests an agent self-update on a node.
type AgentUpdateInput struct {
Channel string `json:"channel"`
TagName string `json:"tag_name"`
}
// AgentReleaseInfo describes the latest agent release for a node.
type AgentReleaseInfo 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"`
}
// BootstrapView exposes the global discovery token.
type BootstrapView struct {
DiscoveryToken string `json:"discovery_token"`
}
// View is the admin-facing node representation.
type View struct {
ID uint `json:"id"`
NodeID string `json:"node_id"`
Name string `json:"name"`
IP string `json:"ip"`
IPManualOverride bool `json:"ip_manual_override"`
GeoName string `json:"geo_name"`
GeoLatitude *float64 `json:"geo_latitude"`
GeoLongitude *float64 `json:"geo_longitude"`
GeoManualOverride bool `json:"geo_manual_override"`
AccessToken string `json:"access_token"`
AutoUpdateEnabled bool `json:"auto_update_enabled"`
UpdateRequested bool `json:"update_requested"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
RestartOpenrestyRequested bool `json:"restart_openresty_requested"`
Version string `json:"version"`
ExtVersion string `json:"ext_version"`
OpenrestyStatus string `json:"openresty_status"`
OpenrestyMessage string `json:"openresty_message"`
Status string `json:"status"`
CurrentVersion string `json:"current_version"`
LastSeenAt any `json:"last_seen_at"`
LastError string `json:"last_error"`
LatestApplyResult string `json:"latest_apply_result"`
LatestApplyMessage string `json:"latest_apply_message"`
LatestApplyChecksum string `json:"latest_apply_checksum"`
LatestMainConfigChecksum string `json:"latest_main_config_checksum"`
LatestRouteConfigChecksum string `json:"latest_route_config_checksum"`
LatestSupportFileCount int `json:"latest_support_file_count"`
LatestApplyAt *time.Time `json:"latest_apply_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
NodeType string `json:"node_type"`
RelayBindPort int `json:"relay_bind_port"`
RelayVhostHTTPPort int `json:"relay_vhost_http_port"`
RelayAgentAccessAddr string `json:"relay_agent_access_addr"`
RelayClientAccessAddr string `json:"relay_client_access_addr"`
RelayClientProxyURL string `json:"relay_client_proxy_url"`
RelayStatus string `json:"relay_status"`
RelayWebServerEnabled bool `json:"relay_web_server_enabled"`
}
// ObservabilityQuery filters node observability data.
type ObservabilityQuery struct {
Hours int `json:"hours"`
Limit int `json:"limit"`
}
// ObservabilityView is the node observability API response.
type ObservabilityView = observability.NodeView
// HealthEventCleanupResult reports health event cleanup outcome.
type HealthEventCleanupResult = observability.HealthEventCleanupResult
// ListNodes returns all node views with latest apply log metadata.
func ListNodes(ctx context.Context) ([]*View, error) {
nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil {
return nil, err
}
nodeIDs := make([]string, 0, len(nodes))
for _, node := range nodes {
nodeIDs = append(nodeIDs, node.NodeID)
}
latestLogs, err := repository.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs)
if err != nil {
return nil, err
}
views := make([]*View, 0, len(nodes))
for _, node := range nodes {
view := buildNodeView(&node)
view.Status = computeNodeStatus(&node)
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
}
// CreateNode creates a reserved node with generated node_id and access_token.
func CreateNode(ctx context.Context, input Input) (*View, error) {
name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input)
if err != nil {
return nil, err
}
if name == "" {
return nil, errors.New(errNodeNameRequired)
}
ipManualOverride := resolveNodeIPManualOverride(input, nil, ip)
node := &model.OpenFlareNode{
Name: name,
IP: ip,
IPManualOverride: ipManualOverride,
GeoName: geoName,
GeoLatitude: geoLatitude,
GeoLongitude: geoLongitude,
GeoManualOverride: geoManualOverride,
Version: "",
ExtVersion: "",
Status: nodeStatusPending,
AutoUpdateEnabled: input.AutoUpdateEnabled,
NodeType: normalizeNodeType(input.NodeType),
CapabilitiesJSON: "[]",
}
node.NodeID, err = newServerNodeID()
if err != nil {
return nil, err
}
node.AccessToken, err = newRandomToken()
if err != nil {
return nil, err
}
if node.NodeType == "tunnel_relay" {
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, defaultRelayBindPort)
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, defaultRelayVhostHTTPPort)
node.RelayAuthToken, err = newRandomToken()
if err != nil {
return nil, err
}
node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr)
node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr)
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
node.RelayWebServerEnabled = input.RelayWebServerEnabled
}
if err = repository.CreateOpenFlareNode(ctx, node); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errNodeIDConflict)
}
return nil, err
}
return buildNodeView(node), nil
}
// UpdateNode updates an existing node.
func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) {
name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input)
if err != nil {
return nil, err
}
if name == "" {
return nil, errors.New(errNodeNameRequired)
}
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
ipManualOverride := resolveNodeIPManualOverride(input, node, ip)
previousIP := node.IP
node.Name = name
node.IP = ip
node.IPManualOverride = ipManualOverride
node.GeoName = geoName
node.GeoLatitude = geoLatitude
node.GeoLongitude = geoLongitude
node.GeoManualOverride = geoManualOverride
node.AutoUpdateEnabled = input.AutoUpdateEnabled
if node.NodeType == "tunnel_relay" {
node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr)
node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr)
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
node.RelayWebServerEnabled = input.RelayWebServerEnabled
if input.RelayBindPort > 0 {
node.RelayBindPort = input.RelayBindPort
}
if input.RelayVhostHTTPPort > 0 {
node.RelayVhostHTTPPort = input.RelayVhostHTTPPort
}
}
if err = repository.SaveOpenFlareNode(ctx, node); err != nil {
return nil, err
}
if strings.TrimSpace(previousIP) != strings.TrimSpace(node.IP) {
if _, dispatchErr := cf.DispatchNodeSync(ctx, node.ID, "cloudflare_node_ip_update"); dispatchErr != nil {
logger.ErrorF(ctx, "[Cloudflare] enqueue node sync failed: node_id=%d error=%v", node.ID, dispatchErr)
}
}
return buildNodeView(node), nil
}
// DeleteNode removes a node by id.
func DeleteNode(ctx context.Context, id uint) error {
if _, err := repository.GetOpenFlareNodeByID(ctx, id); err != nil {
return err
}
return repository.DeleteOpenFlareNode(ctx, id)
}
// GetBootstrapToken returns the global discovery token, creating one if missing.
func GetBootstrapToken(ctx context.Context) (*BootstrapView, error) {
token, err := ensureGlobalDiscoveryToken(ctx)
if err != nil {
return nil, err
}
return &BootstrapView{DiscoveryToken: token}, nil
}
// RotateBootstrapToken rotates the global discovery token.
func RotateBootstrapToken(ctx context.Context) (*BootstrapView, error) {
token, err := newRandomToken()
if err != nil {
return nil, err
}
if err = repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyAgentDiscoveryToken, token); err != nil {
return nil, err
}
return &BootstrapView{DiscoveryToken: token}, nil
}
// GetAgentRelease checks the latest agent release for a node.
func GetAgentRelease(ctx context.Context, id uint, channel string) (*AgentReleaseInfo, error) {
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
release, err := fetchLatestGitHubRelease(ctx, getAgentUpdateRepo(ctx), normalizeReleaseChannel(channel))
if err != nil {
return nil, err
}
return buildNodeAgentReleaseView(node, release, normalizeReleaseChannel(channel)), nil
}
// RequestAgentUpdate marks a node for manual agent update.
func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*View, error) {
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
channel := normalizeReleaseChannel(input.Channel)
tagName := strings.TrimSpace(input.TagName)
if tagName != "" {
release, releaseErr := fetchGitHubReleaseByTag(ctx, getAgentUpdateRepo(ctx), tagName)
if releaseErr != nil {
return nil, releaseErr
}
if channel == releaseChannelPreview && !release.Prerelease {
return nil, errors.New(errAgentPreviewTagInvalid)
}
if channel == releaseChannelStable && release.Prerelease {
return nil, errors.New(errAgentStableTagInvalid)
}
}
node.UpdateRequested = true
node.UpdateChannel = channel.String()
node.UpdateTag = tagName
if err = repository.UpdateOpenFlareNodeFields(ctx, node, "update_requested", "update_channel", "update_tag"); err != nil {
return nil, err
}
return buildNodeView(node), nil
}
// RequestOpenrestyRestart marks a node for openresty restart.
func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) {
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
node.RestartOpenrestyRequested = true
if err = repository.UpdateOpenFlareNodeFields(ctx, node, "restart_openresty_requested"); err != nil {
return nil, err
}
return buildNodeView(node), nil
}
// RequestForceSync pushes force_sync_config to a connected agent websocket.
func RequestForceSync(ctx context.Context, id uint) (*View, error) {
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
activeConfig, err := repository.GetActiveConfigVersion(ctx)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion)
}
return nil, fmt.Errorf("无法获取当前激活的配置版本:%w", err)
}
if !ofws.SendForceSyncConfig(node.NodeID, forceSyncConfigPayload{
Version: activeConfig.Version,
Checksum: activeConfig.Checksum,
}) {
return nil, errors.New(errNodeForceSyncFailed)
}
return buildNodeView(node), nil
}
// GetObservability returns observability details for a node.
func GetObservability(ctx context.Context, id uint, query ObservabilityQuery) (*ObservabilityView, error) {
return observability.GetNodeObservability(ctx, id, observability.NodeQuery{
Hours: query.Hours,
Limit: query.Limit,
})
}
// CleanupHealthEvents removes all health events for a node.
func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) {
return observability.CleanupHealthEvents(ctx, id)
}
func ensureGlobalDiscoveryToken(ctx context.Context) (string, error) {
// 从 SystemConfig 读取 Agent 发现令牌
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentDiscoveryToken)
if err == nil && strings.TrimSpace(config.Value) != "" {
return strings.TrimSpace(config.Value), nil
}
// 如果不存在,生成新令牌并保存
token, err := newRandomToken()
if err != nil {
return "", err
}
// 更新到 SystemConfig
if err = repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyAgentDiscoveryToken, token); err != nil {
return "", err
}
return token, nil
}
type forceSyncConfigPayload struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
// ValidateDiscoveryToken validates the global discovery token.
func ValidateDiscoveryToken(ctx context.Context, token string) error {
token = strings.TrimSpace(token)
if token == "" {
return errors.New("缺少 Discovery Token")
}
discoveryToken, err := ensureGlobalDiscoveryToken(ctx)
if err != nil {
return err
}
if !tokenEqual(token, discoveryToken) {
return errors.New("discovery Token 无效") // error 消息首字母小写
}
return nil
}
@@ -0,0 +1,383 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package node
import (
"context"
"io"
"net/http"
"strings"
"testing"
"time"
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setReleaseHTTPClientForTest(client *http.Client) *http.Client {
previous := releaseHTTPClient
releaseHTTPClient = client
return previous
}
func setupNodeTestDB(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.SystemConfig{},
&model.OpenFlareApplyLog{},
))
db.SetDB(sqliteDB)
testhelper.SetupLogStoresForTest(t)
return func() {
db.SetDB(nil)
}
}
func TestCreateEdgeNode(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
view, err := CreateNode(ctx, Input{
Name: "edge-1",
IP: "10.0.0.1",
AutoUpdateEnabled: true,
})
require.NoError(t, err)
assert.NotZero(t, view.ID)
assert.True(t, strings.HasPrefix(view.NodeID, "node-"))
assert.Len(t, view.AccessToken, 32)
assert.Equal(t, "edge_node", view.NodeType)
assert.Equal(t, nodeStatusPending, view.Status)
assert.True(t, view.AutoUpdateEnabled)
}
func TestCreateTunnelRelayNode(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
view, err := CreateNode(ctx, Input{
Name: "relay-1",
NodeType: "tunnel_relay",
})
require.NoError(t, err)
assert.Equal(t, "tunnel_relay", view.NodeType)
assert.Equal(t, 7000, view.RelayBindPort)
assert.Equal(t, 8080, view.RelayVhostHTTPPort)
stored, err := repository.GetOpenFlareNodeByID(ctx, view.ID)
require.NoError(t, err)
assert.NotEmpty(t, stored.RelayAuthToken)
}
func TestCreateTunnelClientNode(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
view, err := CreateNode(ctx, Input{
Name: "client-1",
NodeType: "tunnel_client",
})
require.NoError(t, err)
assert.Equal(t, "tunnel_client", view.NodeType)
}
func TestCreateNodeRequiresName(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
_, err := CreateNode(ctx, Input{IP: "10.0.0.2"})
require.Error(t, err)
assert.Equal(t, errNodeNameRequired, err.Error())
}
func TestUpdateNode(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-update"})
require.NoError(t, err)
updated, err := UpdateNode(ctx, created.ID, Input{
Name: "edge-updated",
IP: "192.168.1.10",
AutoUpdateEnabled: true,
})
require.NoError(t, err)
assert.Equal(t, "edge-updated", updated.Name)
assert.Equal(t, "192.168.1.10", updated.IP)
assert.True(t, updated.AutoUpdateEnabled)
}
func TestUpdateNodeDispatchesCloudflareSyncWhenIPChanges(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-update", IP: "192.0.2.10"})
require.NoError(t, err)
var dispatchedNodeID uint
restore := cf.SetDispatchTaskForTest(func(_ context.Context, taskType string, payload []byte, _ string) (string, error) {
assert.Equal(t, cf.TaskTypeSyncByNode, taskType)
assert.Contains(t, string(payload), `"node_id":`)
dispatchedNodeID = created.ID
return "task-1", nil
})
defer restore()
_, err = UpdateNode(ctx, created.ID, Input{Name: "edge-update", IP: "192.0.2.11"})
require.NoError(t, err)
assert.Equal(t, created.ID, dispatchedNodeID)
}
func TestDeleteNode(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-delete"})
require.NoError(t, err)
require.NoError(t, DeleteNode(ctx, created.ID))
_, err = repository.GetOpenFlareNodeByID(ctx, created.ID)
require.Error(t, err)
assert.ErrorIs(t, err, gorm.ErrRecordNotFound)
}
func TestListNodesWithApplyLogMetadata(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-list"})
require.NoError(t, err)
applyAt := time.Now().UTC().Truncate(time.Second)
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{
NodeID: created.NodeID,
Version: "20260618-001",
Result: "success",
Message: "ok",
Checksum: "checksum-1",
MainConfigChecksum: "main-1",
RouteConfigChecksum: "route-1",
SupportFileCount: 3,
CreatedAt: applyAt,
}).Error)
views, err := ListNodes(ctx)
require.NoError(t, err)
require.Len(t, views, 1)
assert.Equal(t, "success", views[0].LatestApplyResult)
assert.Equal(t, "checksum-1", views[0].LatestApplyChecksum)
assert.Equal(t, 3, views[0].LatestSupportFileCount)
require.NotNil(t, views[0].LatestApplyAt)
assert.Equal(t, applyAt, views[0].LatestApplyAt.UTC())
}
func TestBootstrapTokenLifecycle(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
first, err := GetBootstrapToken(ctx)
require.NoError(t, err)
assert.Len(t, first.DiscoveryToken, 32)
second, err := GetBootstrapToken(ctx)
require.NoError(t, err)
assert.Equal(t, first.DiscoveryToken, second.DiscoveryToken)
rotated, err := RotateBootstrapToken(ctx)
require.NoError(t, err)
assert.NotEqual(t, first.DiscoveryToken, rotated.DiscoveryToken)
// 验证令牌已保存到 SystemConfig
savedToken, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentDiscoveryToken)
require.NoError(t, err)
assert.Equal(t, rotated.DiscoveryToken, savedToken.Value)
}
func TestValidateDiscoveryToken(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
bootstrap, err := GetBootstrapToken(ctx)
require.NoError(t, err)
require.NoError(t, ValidateDiscoveryToken(ctx, bootstrap.DiscoveryToken))
require.Error(t, ValidateDiscoveryToken(ctx, "invalid-token"))
require.Error(t, ValidateDiscoveryToken(ctx, ""))
require.Error(t, ValidateDiscoveryToken(ctx, bootstrap.DiscoveryToken[:len(bootstrap.DiscoveryToken)-1]+"x"))
}
func TestRequestAgentUpdateWithPreviewTag(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-update-agent"})
require.NoError(t, err)
originalClient := setReleaseHTTPClientForTest(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
expected := "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v0.5.0-rc.1"
require.Equal(t, expected, req.URL.String())
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v0.5.0-rc.1","prerelease":true}`)),
}, nil
}),
})
t.Cleanup(func() {
setReleaseHTTPClientForTest(originalClient)
})
updated, err := RequestAgentUpdate(ctx, created.ID, AgentUpdateInput{
Channel: "preview",
TagName: "v0.5.0-rc.1",
})
require.NoError(t, err)
assert.True(t, updated.UpdateRequested)
assert.Equal(t, "preview", updated.UpdateChannel)
assert.Equal(t, "v0.5.0-rc.1", updated.UpdateTag)
}
func TestRequestOpenrestyRestart(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-restart"})
require.NoError(t, err)
updated, err := RequestOpenrestyRestart(ctx, created.ID)
require.NoError(t, err)
assert.True(t, updated.RestartOpenrestyRequested)
}
func seedActiveConfigVersion(t *testing.T, ctx context.Context) {
t.Helper()
conn := db.DB(ctx)
require.NotNil(t, conn)
require.NoError(t, conn.AutoMigrate(&model.ConfigVersion{}))
require.NoError(t, conn.Create(&model.ConfigVersion{
Version: "20260618-001",
SnapshotJSON: `{}`,
RenderedConfig: `server {}`,
Checksum: "abc123",
IsActive: true,
CreatedBy: "test",
}).Error)
}
func TestRequestForceSyncRequiresWebSocket(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
seedActiveConfigVersion(t, ctx)
created, err := CreateNode(ctx, Input{Name: "edge-sync"})
require.NoError(t, err)
_, err = RequestForceSync(ctx, created.ID)
require.Error(t, err)
assert.Equal(t, errNodeForceSyncFailed, err.Error())
}
func TestRequestForceSyncRequiresActiveConfig(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
conn := db.DB(ctx)
require.NotNil(t, conn)
require.NoError(t, conn.AutoMigrate(&model.ConfigVersion{}))
created, err := CreateNode(ctx, Input{Name: "edge-sync-active"})
require.NoError(t, err)
_, err = RequestForceSync(ctx, created.ID)
require.Error(t, err)
assert.Contains(t, err.Error(), errNoActiveConfigVersion)
}
func TestGetObservabilityStub(t *testing.T) {
cleanup := setupNodeTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateNode(ctx, Input{Name: "edge-obs"})
require.NoError(t, err)
view, err := GetObservability(ctx, created.ID, ObservabilityQuery{Hours: 24, Limit: 50})
require.NoError(t, err)
assert.Equal(t, created.NodeID, view.NodeID)
assert.Empty(t, view.MetricSnapshots)
}
func TestComputeNodeStatus(t *testing.T) {
now := time.Now()
pending := &model.OpenFlareNode{}
assert.Equal(t, nodeStatusPending, computeNodeStatus(pending))
online := &model.OpenFlareNode{LastSeenAt: &now}
assert.Equal(t, nodeStatusOnline, computeNodeStatus(online))
// computeNodeStatus 使用默认阈值 60 秒
offlineAt := now.Add(-61 * time.Second)
offline := &model.OpenFlareNode{LastSeenAt: &offlineAt}
assert.Equal(t, nodeStatusOffline, computeNodeStatus(offline))
}
type roundTripFunc func(req *http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestCompareVersions(t *testing.T) {
tests := []struct {
local string
remote string
expected int
}{
{"v3.0.0-beta", "v3.0.0-beta.1", -1},
{"v3.0.0-beta", "v3.0.0", -1},
{"v3.0.0-beta.1", "v3.0.0", -1},
{"dev", "v3.0.0", -1},
{"v3.0.0", "v3.0.0", 0},
{"v3.0.0", "v2.9.9", 1},
{"v3.0.0", "v3.0.1", -1},
{"v3.0.0-beta.1", "v3.0.0-beta.2", -1},
{"v3.0.0-beta.11", "v3.0.0-beta.2", 1},
}
for _, tt := range tests {
t.Run(tt.local+"_vs_"+tt.remote, func(t *testing.T) {
res := compareVersions(tt.local, tt.remote)
assert.Equal(t, tt.expected, res)
})
}
}
@@ -0,0 +1,324 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package node
import (
"encoding/json"
"errors"
"io"
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
func handleLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
return apiutil.AbortNotFoundIfMissing(c, err, errNodeNotFound)
}
// ListNodesHandler lists all nodes.
// @Summary 获取节点列表
// @Description 返回所有节点及最新配置下发记录,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]node.View} "节点列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/nodes [get]
func ListNodesHandler(c *gin.Context) {
nodes, err := ListNodes(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(nodes))
}
// CreateNodeHandler creates a node.
// @Summary 创建节点
// @Description 创建新的边缘节点记录,需要管理员权限
// @Tags openflare-node
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body node.Input true "节点参数"
// @Success 200 {object} response.Any{data=node.View} "创建成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/nodes [post]
func CreateNodeHandler(c *gin.Context) {
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
view, err := CreateNode(c.Request.Context(), input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// UpdateNodeHandler updates a node.
// @Summary 更新节点
// @Description 更新指定节点的配置信息,需要管理员权限
// @Tags openflare-node
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Param body body node.Input true "节点参数"
// @Success 200 {object} response.Any{data=node.View} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/update [post]
func UpdateNodeHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
view, err := UpdateNode(c.Request.Context(), id, input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// DeleteNodeHandler deletes a node.
// @Summary 删除节点
// @Description 删除指定节点记录,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/delete [post]
func DeleteNodeHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := DeleteNode(c.Request.Context(), id); handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// GetBootstrapTokenHandler returns the global discovery token.
// @Summary 获取引导令牌
// @Description 返回全局节点发现引导令牌,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=node.BootstrapView} "引导令牌"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/nodes/bootstrap-token [get]
func GetBootstrapTokenHandler(c *gin.Context) {
view, err := GetBootstrapToken(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// RotateBootstrapTokenHandler rotates the global discovery token.
// @Summary 轮换引导令牌
// @Description 重新生成全局节点发现引导令牌,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=node.BootstrapView} "新引导令牌"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/nodes/bootstrap-token/rotate [post]
func RotateBootstrapTokenHandler(c *gin.Context) {
view, err := RotateBootstrapToken(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// GetAgentReleaseHandler returns the latest agent release for a node.
// @Summary 获取 Agent 发布信息
// @Description 返回指定节点可用的最新 Agent 版本信息,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Param channel query string false "发布渠道"
// @Success 200 {object} response.Any{data=node.AgentReleaseInfo} "Agent 发布信息"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/agent-release [get]
func GetAgentReleaseHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
release, err := GetAgentRelease(c.Request.Context(), id, c.Query("channel"))
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(release))
}
// RequestAgentUpdateHandler requests agent self-update on a node.
// @Summary 请求 Agent 更新
// @Description 向指定节点下发 Agent 自更新指令,需要管理员权限
// @Tags openflare-node
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Param body body node.AgentUpdateInput false "更新参数(可选)"
// @Success 200 {object} response.Any{data=node.View} "更新请求已下发"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/agent-update [post]
func RequestAgentUpdateHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var request AgentUpdateInput
if c.Request.ContentLength > 0 {
if err := bindOptionalJSON(c.Request.Body, &request); err != nil {
response.AbortBadRequest(c, "参数错误")
return
}
}
view, err := RequestAgentUpdate(c.Request.Context(), id, request)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// RequestOpenrestyRestartHandler requests openresty restart on a node.
// @Summary 请求重启 OpenResty
// @Description 向指定节点下发 OpenResty 重启指令,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Success 200 {object} response.Any{data=node.View} "重启请求已下发"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/openresty-restart [post]
func RequestOpenrestyRestartHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
view, err := RequestOpenrestyRestart(c.Request.Context(), id)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// RequestForceSyncHandler requests force sync on a node.
// @Summary 请求强制同步配置
// @Description 向指定节点下发强制同步当前活跃配置的指令,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Success 200 {object} response.Any{data=node.View} "同步请求已下发"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/force-sync [post]
func RequestForceSyncHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
view, err := RequestForceSync(c.Request.Context(), id)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// GetObservabilityHandler returns node observability details.
// @Summary 获取节点可观测性数据
// @Description 返回指定节点的指标、健康事件与流量分析数据,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Param hours query int false "统计时间范围(小时)"
// @Param limit query int false "返回记录数量上限"
// @Success 200 {object} response.Any{data=node.ObservabilityView} "可观测性数据"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/observability [get]
func GetObservabilityHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var query ObservabilityQuery
if err := c.ShouldBindQuery(&query); err != nil {
response.AbortBadRequest(c, "参数错误")
return
}
view, err := GetObservability(c.Request.Context(), id, query)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// CleanupHealthEventsHandler cleans up node health events.
// @Summary 清理节点健康事件
// @Description 清理指定节点的历史健康事件记录,需要管理员权限
// @Tags openflare-node
// @Produce json
// @Security SessionCookie
// @Param id path int true "节点 ID"
// @Success 200 {object} response.Any{data=node.HealthEventCleanupResult} "清理结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或节点不存在"
// @Router /api/v1/d/nodes/{id}/observability/cleanup [post]
func CleanupHealthEventsHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
result, err := CleanupHealthEvents(c.Request.Context(), id)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
func bindOptionalJSON(body io.Reader, target any) error {
if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) {
return err
}
return nil
}
@@ -0,0 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package node
import "Wavelet/openflare/share/ofutil"
func compareVersions(local, remote string) int {
return ofutil.CompareVersions(local, remote)
}
@@ -0,0 +1,11 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package relay provides relay node management and authentication for the OpenFlare platform.
package relay
const (
//nolint:gosec // error message text, not a credential
errAgentTokenInvalid = "无权进行此操作,Agent Token 无效"
errRelayNodeTypeMismatch = "此节点不是 TunnelRelay 类型"
)
@@ -0,0 +1,139 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"net"
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
)
const (
relayStatusUnhealthy = "unhealthy"
releaseChannelStable = "stable"
defaultAgentHeartbeatInterval = 3000 // 默认心跳间隔 3 秒(毫秒)
defaultAgentUpdateRepo = "Rain-kl/OpenFlare"
)
func normalizeRelayStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case "healthy":
return "healthy"
case relayStatusUnhealthy:
return relayStatusUnhealthy
default:
return "unknown"
}
}
func normalizeReleaseChannel(channel string) string {
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
return "preview"
}
return releaseChannelStable
}
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
reported := normalizeNodeIP(reportedIP)
remote := normalizeRemoteAddr(remoteAddr)
if reported == "" {
return remote
}
if isPublicNodeIP(reported) {
return reported
}
if isPublicNodeIP(remote) {
return remote
}
return reported
}
func normalizeNodeIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if host, _, err := net.SplitHostPort(raw); err == nil {
raw = host
}
raw = strings.Trim(raw, "[]")
return raw
}
func normalizeRemoteAddr(remoteAddr string) string {
remoteAddr = strings.TrimSpace(remoteAddr)
if remoteAddr == "" {
return ""
}
host, _, err := net.SplitHostPort(remoteAddr)
if err != nil {
return normalizeNodeIP(remoteAddr)
}
return normalizeNodeIP(host)
}
func isPublicNodeIP(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
if ip == nil {
return false
}
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func buildRelayConfig(ctx context.Context, node *model.OpenFlareNode) *Config {
if node == nil {
return nil
}
webServerPort, err := repository.GetIntByKey(ctx, model.ConfigKeyRelayFRPSWebUIPort)
if err != nil || webServerPort <= 0 {
webServerPort = node.RelayBindPort + 500
}
return &Config{
BindPort: node.RelayBindPort,
VhostHTTPPort: node.RelayVhostHTTPPort,
AuthToken: node.RelayAuthToken,
LogLevel: "info",
WebServerEnabled: node.RelayWebServerEnabled,
WebServerPort: webServerPort,
}
}
// BuildSettings returns runtime settings shared by relay and flared clients.
func BuildSettings(ctx context.Context, node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *Settings {
autoUpdate := false
if node != nil {
autoUpdate = node.AutoUpdateEnabled
}
if strings.TrimSpace(updateChannel) == "" {
updateChannel = releaseChannelStable
}
// 从 SystemConfig 读取配置,使用默认值作为降级
heartbeatInterval, _ := repository.GetIntByKey(ctx, model.ConfigKeyAgentHeartbeatInterval)
if heartbeatInterval <= 0 {
heartbeatInterval = defaultAgentHeartbeatInterval
}
wsUpgradeEnabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyAgentWebsocketUpgradeEnabled)
updateRepo, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentUpdateRepo)
if strings.TrimSpace(updateRepo.Value) == "" {
updateRepo.Value = defaultAgentUpdateRepo
}
return &Settings{
HeartbeatInterval: heartbeatInterval,
WebsocketUpgradeEnabled: wsUpgradeEnabled,
AutoUpdate: autoUpdate,
UpdateRepo: updateRepo.Value,
UpdateNow: updateNow,
UpdateChannel: updateChannel,
UpdateTag: strings.TrimSpace(updateTag),
}
}
@@ -0,0 +1,112 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"errors"
"fmt"
"strings"
"time"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
)
const nodeStatusOnline = "online"
// Heartbeat processes a relay heartbeat, updates node status, and returns config.
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
if node == nil {
return nil, errors.New("relay node is nil")
}
payload.Version = strings.TrimSpace(payload.Version)
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
payload.RelayStatus = normalizeRelayStatus(payload.RelayStatus)
payload.Name = strings.TrimSpace(payload.Name)
payload.IP = strings.TrimSpace(payload.IP)
previous := *node
updateNow := node.UpdateRequested
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
updateTag := strings.TrimSpace(node.UpdateTag)
now := time.Now().UTC()
changes := map[string]any{
"version": payload.Version,
"ext_version": payload.ExtVersion,
"relay_status": payload.RelayStatus,
"last_seen_at": now,
"status": nodeStatusOnline,
"update_requested": false,
"update_channel": releaseChannelStable,
"update_tag": "",
}
if payload.Name != "" && strings.TrimSpace(node.Name) == "" {
changes["name"] = payload.Name
node.Name = payload.Name
}
if payload.IP != "" && !node.IPManualOverride {
changes["ip"] = payload.IP
node.IP = payload.IP
}
if !node.GeoManualOverride {
beforeGeo := node.GeoName
beforeLat := node.GeoLatitude
beforeLon := node.GeoLongitude
ofgeoip.ApplyNodeGeoFromIP(ctx, node, node.IP)
if node.GeoName != beforeGeo {
changes["geo_name"] = node.GeoName
}
if !coordinatesEqual(beforeLat, node.GeoLatitude) {
changes["geo_latitude"] = node.GeoLatitude
}
if !coordinatesEqual(beforeLon, node.GeoLongitude) {
changes["geo_longitude"] = node.GeoLongitude
}
}
if !previous.UpdateRequested {
delete(changes, "update_requested")
}
if previous.UpdateChannel == releaseChannelStable {
delete(changes, "update_channel")
}
if previous.UpdateTag == "" {
delete(changes, "update_tag")
}
node.Version = payload.Version
node.ExtVersion = payload.ExtVersion
node.RelayStatus = payload.RelayStatus
node.UpdateRequested = false
node.UpdateChannel = releaseChannelStable
node.UpdateTag = ""
lastSeen := now
node.LastSeenAt = &lastSeen
node.Status = nodeStatusOnline
if err := repository.UpdateOpenFlareNodeColumns(ctx, node, changes); 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(ctx, node),
RelaySettings: BuildSettings(ctx, node, updateNow, updateChannel, updateTag),
}, nil
}
func coordinatesEqual(before *float64, after *float64) bool {
if before == nil || after == nil {
return before == after
}
return *before == *after
}
@@ -0,0 +1,176 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"encoding/json"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"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.SystemConfig{},
&model.OpenFlareNodeSystemProfile{},
&model.OpenFlareMetricSnapshot{},
&model.OpenFlareHealthEvent{},
&model.OpenFlareNodeObservationFrps{},
))
db.SetDB(sqliteDB)
agent.ResetAuthCacheForTest()
testhelper.SetupLogStoresForTest(t)
return func() {
db.SetDB(nil)
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,
DiskReadBytes: 100,
DiskWriteBytes: 200,
},
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 := repository.GetOpenFlareNodeSystemProfile(ctx, node.NodeID)
require.NoError(t, err)
assert.Equal(t, "relay-runtime", profile.Hostname)
assert.Equal(t, "Ubuntu", profile.OSName)
snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10)
require.NoError(t, err)
require.Len(t, snapshots, 1)
assert.InDelta(t, 12.5, snapshots[0].CPUUsagePercent, 1e-9)
frpsObs, err := repository.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 := repository.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 = repository.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,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"strings"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
const ctxRelayNodeKey = "relay_node"
// Auth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
func Auth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader("X-Agent-Token"))
node, err := agent.AuthenticateAccessToken(c.Request.Context(), token)
if err != nil {
response.AbortUnauthorized(c, errAgentTokenInvalid)
return
}
if node.NodeType != "tunnel_relay" {
response.AbortForbidden(c, errRelayNodeTypeMismatch)
return
}
c.Set(ctxRelayNodeKey, node)
c.Next()
}
}
@@ -0,0 +1,112 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupRelayMiddlewareTestDB(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{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func seedRelayNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareNode {
t.Helper()
ctx := context.Background()
node := &model.OpenFlareNode{
NodeID: "relay-test-node",
Name: "relay-test",
Status: "pending",
NodeType: nodeType,
AccessToken: accessToken,
}
require.NoError(t, repository.CreateOpenFlareNode(ctx, node))
return node
}
func TestRelayAuthMissingToken(t *testing.T) {
cleanup := setupRelayMiddlewareTestDB(t)
defer cleanup()
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodGet, "/relay/test", nil)
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusUnauthorized, rec.Code)
}
func TestRelayAuthRejectsWrongNodeType(t *testing.T) {
cleanup := setupRelayMiddlewareTestDB(t)
defer cleanup()
seedRelayNode(t, "edge_node", "edge-token-relay")
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodGet, "/relay/test", nil)
req.Header.Set("X-Agent-Token", "edge-token-relay")
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusForbidden, rec.Code)
}
func TestRelayAuthAcceptsTunnelRelay(t *testing.T) {
cleanup := setupRelayMiddlewareTestDB(t)
defer cleanup()
node := seedRelayNode(t, "tunnel_relay", "relay-token-valid")
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
authNode, ok := c.Get(ctxRelayNodeKey)
require.True(t, ok)
assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID)
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodGet, "/relay/test", nil)
req.Header.Set("X-Agent-Token", "relay-token-valid")
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusOK, rec.Code)
}
@@ -0,0 +1,60 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"context"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/openflare/plugins/server/kernel/model"
"go.uber.org/zap"
)
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 == relayStatusUnhealthy {
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,
},
})
}
return agent.ReconcileScopedNodeHealthEvents(ctx, 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,
HostMetrics: payload.Snapshot,
HealthEvents: payload.HealthEvents,
}, reportedAt)
frpsObs := &model.OpenFlareNodeObservationFrps{
NodeID: nodeID,
CapturedAt: reportedAt,
FrpsConnections: payload.FrpsConnCount,
FrpsProxyCount: payload.FrpsProxyCount,
FrpsClientCount: payload.FrpsClientCount,
FrpsProxies: agent.MarshalJSON(payload.FrpsProxies),
}
if err := repository.InsertOpenFlareNodeObservationFrps(ctx, frpsObs); err != nil {
zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err))
}
}
@@ -0,0 +1,21 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import pkgprotocol "Wavelet/openflare/share/protocol"
// ProxyStat is an alias for protocol.RelayProxyStat.
type ProxyStat = pkgprotocol.RelayProxyStat
// HeartbeatPayload is an alias for protocol.RelayHeartbeatPayload.
type HeartbeatPayload = pkgprotocol.RelayHeartbeatPayload
// Config is an alias for protocol.RelayConfig.
type Config = pkgprotocol.RelayConfig
// Settings is an alias for protocol.RelaySettings.
type Settings = pkgprotocol.RelaySettings
// HeartbeatResponse is an alias for protocol.RelayHeartbeatResponse.
type HeartbeatResponse = pkgprotocol.RelayHeartbeatResponse
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package relay
import (
"net/http"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// PostHeartbeat handles POST /relay/heartbeat.
// @Summary 上报 Relay 心跳
// @Description Relay 节点定期上报运行状态与 frps 观测数据,返回运行时配置
// @Tags openflare-relay
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body relay.HeartbeatPayload true "心跳载荷"
// @Success 200 {object} response.Any{data=relay.HeartbeatResponse} "心跳响应"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Agent Token 无效"
// @Failure 403 {object} response.Any "节点类型不匹配"
// @Router /api/v1/relay/heartbeat [post]
func PostHeartbeat(c *gin.Context) {
var payload HeartbeatPayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := c.Get(ctxRelayNodeKey)
if !ok {
response.AbortUnauthorized(c, errAgentTokenInvalid)
return
}
node, ok := authNode.(*model.OpenFlareNode)
if !ok {
response.AbortUnauthorized(c, errAgentTokenInvalid)
return
}
result, err := Heartbeat(c.Request.Context(), node, payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// GetWebSocket handles GET /relay/ws.
// @Summary 升级 Relay WebSocket 连接
// @Description 将已认证的 Relay 连接升级为 WebSocket 长连接,用于配置推送
// @Tags openflare-relay
// @Security AgentTokenAuth
// @Failure 401 {object} response.Any "Agent Token 无效"
// @Failure 403 {object} response.Any "节点类型不匹配"
// @Router /api/v1/relay/ws [get]
func GetWebSocket(c *gin.Context) {
authNode, ok := c.Get(ctxRelayNodeKey)
if !ok {
response.AbortUnauthorized(c, errAgentTokenInvalid)
return
}
node, ok := authNode.(*model.OpenFlareNode)
if !ok {
response.AbortUnauthorized(c, errAgentTokenInvalid)
return
}
ofws.ServeRelay(c, node.NodeID)
}
@@ -0,0 +1,211 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package websocket manages persistent WebSocket connections between the OpenFlare server and its agents.
package websocket
import (
"context"
"encoding/json"
"log/slog"
"sync"
"time"
"github.com/gin-gonic/gin"
)
const (
// AgentWSConnectedLastSeenValue is the sentinel last_seen_at value when agent WS is connected.
AgentWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__"
agentMessageTypeStatus = "status"
agentMessageTypeSettings = "settings"
agentMessageTypeActiveConfig = "active_config"
agentMessageTypeForceSyncConfig = "force_sync_config"
agentMessageTypeWAFIPGroups = "waf_ip_groups"
)
// AgentStatusHandler processes inbound agent websocket status payloads.
type AgentStatusHandler func(ctx context.Context, nodeID, remoteAddr string, payload json.RawMessage)
type agentClient struct {
wsClientCore
remoteAddr string
onStatus AgentStatusHandler
}
type agentHub struct {
mu sync.RWMutex
clients map[string]*agentClient
}
var defaultAgentHub = &agentHub{clients: make(map[string]*agentClient)}
// ServeAgent handles an upgraded agent websocket connection.
func ServeAgent(c *gin.Context, nodeID string, onStatus AgentStatusHandler) {
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
slog.Debug("agent ws upgrade failed", "node_id", nodeID, "error", err)
return
}
client := &agentClient{
wsClientCore: wsClientCore{
nodeID: nodeID,
conn: conn,
send: make(chan Message, wsChannelBuf),
done: make(chan struct{}),
},
remoteAddr: c.Request.RemoteAddr,
onStatus: onStatus,
}
defaultAgentHub.register(client)
defer defaultAgentHub.unregister(client)
slog.Debug("agent ws connected", "node_id", nodeID, "remote", client.remoteAddr)
go client.writePump()
client.readPump()
}
func (h *agentHub) register(client *agentClient) {
h.mu.Lock()
if existing := h.clients[client.nodeID]; existing != nil {
existing.close()
}
h.clients[client.nodeID] = client
h.mu.Unlock()
}
func (h *agentHub) unregister(client *agentClient) {
h.mu.Lock()
if current := h.clients[client.nodeID]; current == client {
delete(h.clients, client.nodeID)
}
h.mu.Unlock()
client.close()
}
// IsAgentConnected reports whether an agent websocket is active.
func IsAgentConnected(nodeID string) bool {
defaultAgentHub.mu.RLock()
client := defaultAgentHub.clients[nodeID]
defaultAgentHub.mu.RUnlock()
if client == nil {
return false
}
select {
case <-client.done:
return false
default:
return true
}
}
// SendAgentSettings pushes agent settings to a connected agent.
func SendAgentSettings(nodeID string, payload any) bool {
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeSettings, Payload: payload})
}
// SendAgentActiveConfig pushes active config metadata to a connected agent.
func SendAgentActiveConfig(nodeID string, payload any) bool {
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeActiveConfig, Payload: payload})
}
// SendAgentWAFIPGroups pushes WAF IP group updates to a connected agent.
func SendAgentWAFIPGroups(nodeID string, payload any) bool {
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeWAFIPGroups, Payload: payload})
}
// BroadcastWAFIPGroups pushes changed WAF IP groups to all connected agents.
func BroadcastWAFIPGroups(payload any) int {
return broadcastAgent(agentMessageTypeWAFIPGroups, payload)
}
// BroadcastActiveConfig pushes active config metadata to all connected agents.
func BroadcastActiveConfig(payload any) int {
return broadcastAgent(agentMessageTypeActiveConfig, payload)
}
func broadcastAgent(messageType string, payload any) int {
if payload == nil {
return 0
}
message := Message{Type: messageType, Payload: payload}
defaultAgentHub.mu.RLock()
clients := make([]*agentClient, 0, len(defaultAgentHub.clients))
for _, client := range defaultAgentHub.clients {
clients = append(clients, client)
}
defaultAgentHub.mu.RUnlock()
success := 0
for _, client := range clients {
if client.enqueue(message) {
success++
}
}
return success
}
// SendForceSyncConfig notifies an agent to force sync configuration.
func SendForceSyncConfig(nodeID string, payload any) bool {
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeForceSyncConfig, Payload: payload})
}
func sendAgentMessage(nodeID string, message Message) bool {
defaultAgentHub.mu.RLock()
client := defaultAgentHub.clients[nodeID]
defaultAgentHub.mu.RUnlock()
if client == nil {
return false
}
return client.enqueue(message)
}
func (c *agentClient) readPump() {
defer c.close()
for {
_ = c.conn.SetReadDeadline(time.Now().Add(agentWSReadTimeout()))
_, data, err := c.conn.ReadMessage()
if err != nil {
slog.Debug("agent ws read closed", "node_id", c.nodeID, "error", err)
return
}
var inbound struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
if err = json.Unmarshal(data, &inbound); err != nil {
slog.Debug("agent ws invalid message", "node_id", c.nodeID, "error", err)
continue
}
slog.Debug("agent ws message received", "node_id", c.nodeID, "type", inbound.Type)
switch inbound.Type {
case agentMessageTypeStatus:
if c.onStatus != nil {
c.onStatus(context.Background(), c.nodeID, c.remoteAddr, inbound.Payload)
}
case messageTypePing:
_ = c.enqueue(Message{Type: messageTypePong})
case messageTypePong:
default:
slog.Debug("agent ws unsupported message type", "node_id", c.nodeID, "type", inbound.Type)
}
}
}
func agentWSReadTimeout() time.Duration {
timeout := wsReadDeadline
if timeout < minAgentWSReadTimeout {
return minAgentWSReadTimeout
}
return timeout
}
func (c *agentClient) writePump() {
runWritePump(c.nodeID, c.conn, c.done, c.send, c.close, "agent ws")
}
@@ -0,0 +1,52 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"sync"
"github.com/gorilla/websocket"
)
// wsClientCore holds the state and lifecycle shared by all WebSocket client
// variants (agent/relay/flared). Embed it; call close exactly-once semantics
// are guaranteed via once.
type wsClientCore struct {
nodeID string
conn *websocket.Conn
send chan Message
done chan struct{}
once sync.Once
}
// close tears down the connection at most once.
func (c *wsClientCore) close() {
if c == nil {
return
}
c.once.Do(func() {
close(c.done)
if c.conn != nil {
_ = c.conn.Close()
}
})
}
// enqueue best-effort delivers message; it never blocks and fails fast when
// the client is closed or its send buffer is full.
func (c *wsClientCore) enqueue(message Message) bool {
// 先确定性检查 closed:若与发送合并在同一个 select,两个 case 同时就绪时
// Go 会随机选择,close 后仍可能投递成功。
select {
case <-c.done:
return false
default:
}
select {
case c.send <- message:
return true
default:
return false
}
}
@@ -0,0 +1,61 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"sync"
"testing"
)
func TestWSClientCoreCloseIsIdempotent(t *testing.T) {
core := &wsClientCore{
send: make(chan Message, 1),
done: make(chan struct{}),
}
var wg sync.WaitGroup
for range 8 {
wg.Add(1)
go func() {
defer wg.Done()
core.close()
}()
}
wg.Wait()
select {
case <-core.done:
default:
t.Fatal("close did not signal done")
}
}
func TestWSClientCoreEnqueueFailsAfterClose(t *testing.T) {
core := &wsClientCore{
send: make(chan Message, 1),
done: make(chan struct{}),
}
core.close()
// 循环多次:若 close 检查与发送合并在同一个 select,两 case 同时就绪时
// Go 随机选择,单次调用可能碰巧通过。
for range 50 {
if core.enqueue(Message{Type: messageTypePing}) {
t.Fatal("enqueue must fail after close")
}
}
}
func TestWSClientCoreEnqueueNeverBlocks(t *testing.T) {
core := &wsClientCore{
send: make(chan Message, 1), // 缓冲小于消息数,验证不阻塞
done: make(chan struct{}),
}
defer core.close()
for range 3 {
if !core.enqueue(Message{Type: messageTypePing}) && len(core.send) == 0 {
t.Fatal("enqueue failed with empty buffer")
}
}
if core.enqueue(Message{Type: messageTypePing}) {
t.Fatal("enqueue must fail when buffer full")
}
}
@@ -0,0 +1,26 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"net/http"
"github.com/gorilla/websocket"
)
const (
messageTypePing = "ping"
messageTypePong = "pong"
messageTypeNotify = "notify"
)
// Message is a JSON-framed websocket payload.
type Message struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}
var upgrader = websocket.Upgrader{
CheckOrigin: func(_ *http.Request) bool { return true },
}
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import "time"
const (
wsChannelBuf = 16
wsPingInterval = 30 * time.Second
wsReadDeadline = 90 * time.Second
wsWriteDeadline = 10 * time.Second
minAgentWSReadTimeout = 30 * time.Second
)
@@ -0,0 +1,124 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"log/slog"
"sync"
"github.com/gin-gonic/gin"
)
const (
// FlaredWSConnectedLastSeenValue is the sentinel last_seen_at value when flared WS is connected.
FlaredWSConnectedLastSeenValue = "__OPENFLARE_FLARED_WS_CONNECTED__"
flaredMessageTypeActiveConfig = "active_config"
flaredMessageTypeForceSync = "force_sync"
flaredMessageTypePong = "pong"
)
type flaredClient struct {
wsClientCore
}
type flaredHub struct {
mu sync.RWMutex
clients map[string]*flaredClient
}
var defaultFlaredHub = &flaredHub{clients: make(map[string]*flaredClient)}
// ServeFlared handles an upgraded flared websocket connection.
func ServeFlared(c *gin.Context, nodeID string) {
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
slog.Debug("flared ws upgrade failed", "node_id", nodeID, "error", err)
return
}
client := &flaredClient{
wsClientCore: wsClientCore{
nodeID: nodeID,
conn: conn,
send: make(chan Message, wsChannelBuf),
done: make(chan struct{}),
},
}
defaultFlaredHub.register(client)
defer defaultFlaredHub.unregister(client)
slog.Debug("flared ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr)
go client.writePump()
client.readPump()
}
func (h *flaredHub) register(client *flaredClient) {
h.mu.Lock()
if existing := h.clients[client.nodeID]; existing != nil {
existing.close()
}
h.clients[client.nodeID] = client
h.mu.Unlock()
}
func (h *flaredHub) unregister(client *flaredClient) {
h.mu.Lock()
if current := h.clients[client.nodeID]; current == client {
delete(h.clients, client.nodeID)
}
h.mu.Unlock()
client.close()
}
// DisconnectFlaredClient forcefully disconnects a flared websocket client.
func DisconnectFlaredClient(nodeID string) {
defaultFlaredHub.mu.Lock()
client := defaultFlaredHub.clients[nodeID]
if client != nil {
delete(defaultFlaredHub.clients, nodeID)
}
defaultFlaredHub.mu.Unlock()
if client != nil {
client.close()
}
}
// IsFlaredConnected reports whether a flared websocket is active.
func IsFlaredConnected(nodeID string) bool {
defaultFlaredHub.mu.RLock()
client := defaultFlaredHub.clients[nodeID]
defaultFlaredHub.mu.RUnlock()
if client == nil {
return false
}
select {
case <-client.done:
return false
default:
return true
}
}
// SendFlaredPong enqueues a pong message for the flared node.
func SendFlaredPong(nodeID string) bool {
defaultFlaredHub.mu.RLock()
client := defaultFlaredHub.clients[nodeID]
defaultFlaredHub.mu.RUnlock()
if client == nil {
return false
}
// 委托 enqueue:closed 检查与发送不能合并在同一个 select(两 case 同时
// 就绪时 Go 随机选择,close 后仍可能投递成功)。
return client.enqueue(Message{Type: flaredMessageTypePong})
}
func (c *flaredClient) readPump() {
runReadPump(c.nodeID, c.conn, c.close, "flared ws", SendFlaredPong, flaredMessageTypePong)
}
func (c *flaredClient) writePump() {
runWritePump(c.nodeID, c.conn, c.done, c.send, c.close, "flared ws")
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"encoding/json"
"log/slog"
"time"
"github.com/gorilla/websocket"
)
func runReadPump(
nodeID string,
conn *websocket.Conn,
closeFn func(),
logLabel string,
sendPong func(string) bool,
clientPongType string,
) {
defer closeFn()
_ = conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
})
for {
_, data, err := conn.ReadMessage()
if err != nil {
slog.Debug(logLabel+" read closed", "node_id", nodeID, "error", err)
return
}
var message Message
if err = json.Unmarshal(data, &message); err != nil {
slog.Debug(logLabel+" invalid message", "node_id", nodeID, "error", err)
continue
}
switch message.Type {
case messageTypePing:
_ = sendPong(nodeID)
case clientPongType:
// Refresh read deadline when the client replies with a JSON pong.
// This keeps the connection alive when the WebSocket is proxied
// through Cloudflare, which enforces a 100-second idle timeout on
// the TCP stream. Without this refresh, the server's 90-second read
// deadline expires and terminates the connection even though the
// client is actively responding to pings.
_ = conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
default:
slog.Debug(logLabel+" unsupported message", "node_id", nodeID, "type", message.Type)
}
}
}
@@ -0,0 +1,105 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"log/slog"
"sync"
"github.com/gin-gonic/gin"
)
// RelayWSConnectedLastSeenValue is the sentinel last_seen_at value when relay WS is connected.
const RelayWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__"
type relayClient struct {
wsClientCore
}
type relayHub struct {
mu sync.RWMutex
clients map[string]*relayClient
}
var defaultRelayHub = &relayHub{clients: make(map[string]*relayClient)}
// ServeRelay handles an upgraded relay websocket connection.
func ServeRelay(c *gin.Context, nodeID string) {
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
slog.Debug("relay ws upgrade failed", "node_id", nodeID, "error", err)
return
}
client := &relayClient{
wsClientCore: wsClientCore{
nodeID: nodeID,
conn: conn,
send: make(chan Message, wsChannelBuf),
done: make(chan struct{}),
},
}
defaultRelayHub.register(client)
defer defaultRelayHub.unregister(client)
slog.Debug("relay ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr)
go client.writePump()
client.readPump()
}
func (h *relayHub) register(client *relayClient) {
h.mu.Lock()
if existing := h.clients[client.nodeID]; existing != nil {
existing.close()
}
h.clients[client.nodeID] = client
h.mu.Unlock()
}
func (h *relayHub) unregister(client *relayClient) {
h.mu.Lock()
if current := h.clients[client.nodeID]; current == client {
delete(h.clients, client.nodeID)
}
h.mu.Unlock()
client.close()
}
// IsRelayConnected reports whether a relay websocket is active.
func IsRelayConnected(nodeID string) bool {
defaultRelayHub.mu.RLock()
client := defaultRelayHub.clients[nodeID]
defaultRelayHub.mu.RUnlock()
if client == nil {
return false
}
select {
case <-client.done:
return false
default:
return true
}
}
// SendRelayPong enqueues a pong message for the relay node.
func SendRelayPong(nodeID string) bool {
defaultRelayHub.mu.RLock()
client := defaultRelayHub.clients[nodeID]
defaultRelayHub.mu.RUnlock()
if client == nil {
return false
}
// 委托 enqueue:closed 检查与发送不能合并在同一个 select(两 case 同时
// 就绪时 Go 随机选择,close 后仍可能投递成功)。
return client.enqueue(Message{Type: messageTypePong})
}
func (c *relayClient) readPump() {
runReadPump(c.nodeID, c.conn, c.close, "relay ws", SendRelayPong, messageTypePong)
}
func (c *relayClient) writePump() {
runWritePump(c.nodeID, c.conn, c.done, c.send, c.close, "relay ws")
}
@@ -0,0 +1,47 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"log/slog"
"time"
"github.com/gorilla/websocket"
)
// runWritePump drains send onto conn until done is closed, emitting
// JSON pings at wsPingInterval. Shared by agent/relay/flared clients;
// closeFn must be idempotent.
func runWritePump(
nodeID string,
conn *websocket.Conn,
done <-chan struct{},
send chan Message,
closeFn func(),
logLabel string,
) {
ticker := time.NewTicker(wsPingInterval)
defer ticker.Stop()
for {
select {
case <-done:
return
case message := <-send:
_ = conn.SetWriteDeadline(time.Now().Add(wsWriteDeadline))
if err := conn.WriteJSON(message); err != nil {
slog.Debug(logLabel+" write failed", "node_id", nodeID, "error", err)
closeFn()
return
}
case <-ticker.C:
select {
case <-done:
return
case send <- Message{Type: messageTypePing}:
default:
}
}
}
}
@@ -0,0 +1,85 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestResolveAccessLogIPSummaryWindowHours(t *testing.T) {
t.Parallel()
since, until, hours, err := resolveAccessLogIPSummaryWindow("", "", 0)
require.NoError(t, err)
assert.Equal(t, defaultAccessLogQueryDays*24, hours)
assert.WithinDuration(t, time.Now().UTC(), until, 2*time.Second)
assert.WithinDuration(t, until.Add(-time.Duration(hours)*time.Hour), since, time.Second)
_, _, hours, err = resolveAccessLogIPSummaryWindow("", "", 24)
require.NoError(t, err)
assert.Equal(t, 24, hours)
_, _, hours, err = resolveAccessLogIPSummaryWindow("", "", 9999)
require.NoError(t, err)
assert.Equal(t, maxAccessLogOverviewHours, hours)
}
func TestResolveAccessLogIPSummaryWindowCustomRange(t *testing.T) {
t.Parallel()
start := time.Date(2026, 7, 1, 0, 0, 0, 0, time.UTC)
end := start.Add(72 * time.Hour)
since, until, hours, err := resolveAccessLogIPSummaryWindow(
start.Format(time.RFC3339),
end.Format(time.RFC3339),
24,
)
require.NoError(t, err)
assert.True(t, since.Equal(start))
assert.True(t, until.Equal(end))
assert.Equal(t, 72, hours)
}
func TestResolveAccessLogIPSummaryWindowErrors(t *testing.T) {
t.Parallel()
_, _, _, err := resolveAccessLogIPSummaryWindow("2026-07-01T00:00:00Z", "", 24)
require.Error(t, err)
_, _, _, err = resolveAccessLogIPSummaryWindow("bad", "2026-07-02T00:00:00Z", 24)
require.Error(t, err)
start := time.Date(2026, 7, 2, 0, 0, 0, 0, time.UTC)
end := start.Add(-time.Hour)
_, _, _, err = resolveAccessLogIPSummaryWindow(
start.Format(time.RFC3339),
end.Format(time.RFC3339),
24,
)
require.Error(t, err)
start = time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)
end = start.Add(40 * 24 * time.Hour)
_, _, _, err = resolveAccessLogIPSummaryWindow(
start.Format(time.RFC3339),
end.Format(time.RFC3339),
24,
)
require.Error(t, err)
}
func TestNormalizeIPSummarySortBy(t *testing.T) {
t.Parallel()
assert.Equal(t, "total_requests", normalizeIPSummarySortBy(""))
assert.Equal(t, "request_length", normalizeIPSummarySortBy("bytes_received"))
assert.Equal(t, "request_length", normalizeIPSummarySortBy("request_length"))
assert.Equal(t, "bytes_sent", normalizeIPSummarySortBy("bytes_sent"))
assert.Equal(t, "success_ratio", normalizeIPSummarySortBy("success_ratio"))
assert.Equal(t, "last_seen_at", normalizeIPSummarySortBy("last_seen_at"))
assert.Equal(t, "remote_addr", normalizeIPSummarySortBy("remote_addr"))
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,697 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"context"
"sort"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
)
const observabilityTrendBuckets = 24
const unknownTrendNodeKey = "__unknown__"
const (
healthEventStatusActive = "active"
healthEventStatusResolved = "resolved"
healthSeverityCritical = "critical"
healthSeverityWarning = "warning"
percentageMultiplier = 100
sortOrderAsc = "asc"
)
// DistributionItem is a key/value distribution entry.
type DistributionItem struct {
Key string `json:"key"`
Value int64 `json:"value"`
}
// TrafficDistributions groups traffic distribution charts.
type TrafficDistributions struct {
StatusCodes []DistributionItem `json:"status_codes"`
TopDomains []DistributionItem `json:"top_domains"`
SourceCountries []DistributionItem `json:"source_countries"`
}
const metricSnapshotEdgeHealthMatchWindow = 2 * time.Minute
// NodeMetricSnapshotView is a metric snapshot enriched with edge health connections.
type NodeMetricSnapshotView struct {
ID uint `json:"id,omitempty"`
NodeID string `json:"node_id,omitempty"`
CapturedAt time.Time `json:"captured_at"`
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"`
OpenrestyConnections int64 `json:"openresty_connections"`
}
// TrafficWindowSummary summarizes a traffic reporting window.
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"`
}
// HealthSummary summarizes node health alerts and risks.
type HealthSummary 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"`
}
// TrafficTrendPoint is a traffic trend bucket.
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"`
Status2xxCount int64 `json:"status_2xx_count"`
Status4xxCount int64 `json:"status_4xx_count"`
Status5xxCount int64 `json:"status_5xx_count"`
}
// CapacityTrendPoint is a capacity trend bucket.
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"`
}
// NetworkTrendPoint is a business-byte trend bucket from access logs (L1).
// Host NIC trends are intentionally not exposed.
type NetworkTrendPoint struct {
BucketStartedAt time.Time `json:"bucket_started_at"`
BytesReceived int64 `json:"bytes_received"` // sum(request_length)
BytesProvided int64 `json:"bytes_provided"` // sum(bytes_sent)
ReportedNodes int `json:"reported_nodes"`
}
// DiskIOTrendPoint is a disk IO trend bucket.
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 distributionAccumulator map[string]int64
type capacityTrendAccumulator struct {
cpuSum float64
cpuCount int
memSum float64
memCount int
nodes map[string]struct{}
}
type snapshotTrendAccumulator struct {
nodes map[string]struct{}
}
type diskCounterState struct {
read int64
write int64
seen bool
}
func buildTrafficWindowSummaryFromAccessLogs(
ctx context.Context,
nodeID string,
since, until time.Time,
) *TrafficWindowSummary {
row, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{
NodeID: nodeID,
Since: since,
Until: until,
})
if err != nil || row.RequestCount <= 0 {
return nil
}
summary := &TrafficWindowSummary{
WindowStartedAt: since.UTC(),
WindowEndedAt: until.UTC(),
RequestCount: row.RequestCount,
UniqueVisitorCount: row.UniqueIPCount,
ErrorCount: row.ErrorCount,
}
if duration := until.Sub(since).Seconds(); duration > 0 {
summary.EstimatedQPS = float64(row.RequestCount) / duration
}
if row.RequestCount > 0 {
summary.ErrorRatePercent = (float64(row.ErrorCount) / float64(row.RequestCount)) * 100
}
return summary
}
// BuildMetricSnapshotViews merges metric snapshots with edge health connections for API responses.
func BuildMetricSnapshotViews(
snapshots []*model.OpenFlareMetricSnapshot,
edgeHealth []*model.OpenFlareEdgeHealth,
) []*NodeMetricSnapshotView {
if len(snapshots) == 0 {
return []*NodeMetricSnapshotView{}
}
views := make([]*NodeMetricSnapshotView, 0, len(snapshots))
for _, snapshot := range snapshots {
if snapshot == nil {
continue
}
view := &NodeMetricSnapshotView{
ID: snapshot.ID,
NodeID: snapshot.NodeID,
CapturedAt: snapshot.CapturedAt,
CPUUsagePercent: snapshot.CPUUsagePercent,
MemoryUsedBytes: snapshot.MemoryUsedBytes,
MemoryTotalBytes: snapshot.MemoryTotalBytes,
StorageUsedBytes: snapshot.StorageUsedBytes,
StorageTotalBytes: snapshot.StorageTotalBytes,
DiskReadBytes: snapshot.DiskReadBytes,
DiskWriteBytes: snapshot.DiskWriteBytes,
}
if matched := matchEdgeHealth(snapshot.CapturedAt, edgeHealth); matched != nil {
view.OpenrestyConnections = matched.Connections
}
views = append(views, view)
}
return views
}
// BuildTrafficDistributionsFromAccessLogs builds distributions from access logs (L1).
func BuildTrafficDistributionsFromAccessLogs(
ctx context.Context,
since, until time.Time,
limit int,
accessLogRegions []*model.OpenFlareAccessLogRegionCount,
) TrafficDistributions {
statusCodes := make(distributionAccumulator)
topDomains := make(distributionAccumulator)
query := model.OpenFlareAccessLogQuery{Since: since, Until: until}
if statusRows, err := repository.ValueCountsOpenFlareAccessLogs(ctx, query, "status_code", limit); err == nil {
for _, row := range statusRows {
if strings.TrimSpace(row.Value) == "" || row.Count <= 0 {
continue
}
statusCodes[row.Value] = row.Count
}
}
if hostRows, err := repository.ValueCountsOpenFlareAccessLogs(ctx, query, "host", limit); err == nil {
for _, row := range hostRows {
if strings.TrimSpace(row.Value) == "" || row.Count <= 0 {
continue
}
topDomains[row.Value] = row.Count
}
}
sourceCountries := make(distributionAccumulator)
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 buildHealthSummary(
snapshot *model.OpenFlareMetricSnapshot,
traffic *TrafficWindowSummary,
events []*model.OpenFlareHealthEvent,
) HealthSummary {
summary := HealthSummary{}
for _, event := range events {
if event == nil {
continue
}
if event.Status == healthEventStatusResolved {
summary.ResolvedAlerts++
continue
}
summary.ActiveAlerts++
switch event.Severity {
case healthSeverityCritical:
summary.CriticalAlerts++
case healthSeverityWarning:
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 traffic != nil && traffic.RequestCount >= 100 {
summary.HasTrafficRisk = (float64(traffic.ErrorCount) / float64(traffic.RequestCount)) >= 0.05
}
summary.HasRuntimeRisk = summary.ActiveAlerts > 0 || summary.HasCapacityRisk || summary.HasTrafficRisk
return summary
}
// BuildNodeTrends builds 24h trend series.
// Business traffic (requests/errors and provided/received bytes) comes from access logs.
// Host capacity/disk come from metric snapshots (hourly when available). Host NIC is not tracked.
func BuildNodeTrends(
ctx context.Context,
now time.Time,
nodeID string,
snapshots []*model.OpenFlareMetricSnapshot,
) NodeTrends {
trendSince := now.Add(-24 * time.Hour)
trafficTrend := BuildTrafficTrendPointsFromAccessLogs(ctx, now, nodeID, trendSince)
capacityTrend := BuildCapacityTrendPoints(now, snapshots)
networkTrend := emptyNetworkTrendPoints(now)
applyAccessLogBytesToNetworkTrend(ctx, now, nodeID, trendSince, networkTrend)
diskIOTrend := BuildDiskIOTrendPoints(now, snapshots)
metricHourly, metricErr := repository.ListOpenFlareMetricHourlySince(ctx, nodeID, trendSince)
if metricErr == nil && len(metricHourly) > 0 {
capacityTrend = BuildCapacityTrendPointsFromHourly(now, metricHourly)
diskIOTrend = BuildDiskIOTrendPointsFromHourly(now, metricHourly)
}
return NodeTrends{
Traffic24h: trafficTrend,
Capacity24h: capacityTrend,
Network24h: networkTrend,
DiskIO24h: diskIOTrend,
}
}
// BuildTrafficTrendPointsFromAccessLogs builds 24h request/error/status buckets from access logs.
// Uses raw bucket aggregates: the hourly rollup (of_access_log_hourly) has no per-status counts,
// and the 24h window on the dashboard is cached, so the raw scan is acceptable.
// UniqueVisitorCount from buckets is exact (uniqExact on raw); TrafficSummary is used elsewhere for UV.
func BuildTrafficTrendPointsFromAccessLogs(ctx context.Context, now time.Time, nodeID string, since time.Time) []TrafficTrendPoint {
start := trendWindowStart(now)
points := make([]TrafficTrendPoint, observabilityTrendBuckets)
for index := range points {
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
}
buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{
NodeID: nodeID,
Since: since,
Until: now,
FoldMinutes: 60,
SortBy: defaultAccessLogSortBy,
SortOrder: sortOrderAsc,
})
if err != nil || len(buckets) == 0 {
return points
}
byEpoch := make(map[int64]*model.OpenFlareAccessLogBucketRow, len(buckets))
for _, row := range buckets {
if row == nil {
continue
}
byEpoch[row.BucketEpoch] = row
}
for index := range points {
epoch := points[index].BucketStartedAt.Unix()
if row, ok := byEpoch[epoch]; ok {
points[index].RequestCount = row.RequestCount
points[index].ErrorCount = row.ServerErrorCount
points[index].UniqueVisitorCount = row.UniqueIPCount
points[index].Status2xxCount = row.Status2xxCount
points[index].Status4xxCount = row.Status4xxCount
points[index].Status5xxCount = row.Status5xxCount
}
}
return points
}
func applyAccessLogBytesToNetworkTrend(ctx context.Context, now time.Time, nodeID string, since time.Time, points []NetworkTrendPoint) {
if len(points) == 0 {
return
}
// Prefer of_access_log_hourly (summed across hosts).
if hourly, err := analyticsListAccessLogHourlyBytes(ctx, nodeID, since); err == nil && len(hourly) > 0 {
for hourUnix, totals := range hourly {
for index := range points {
if points[index].BucketStartedAt.Unix() == hourUnix {
points[index].BytesProvided = totals.provided
points[index].BytesReceived = totals.received
}
}
}
return
}
buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{
NodeID: nodeID,
Since: since,
Until: now,
FoldMinutes: 60,
SortBy: defaultAccessLogSortBy,
SortOrder: sortOrderAsc,
})
if err != nil || len(buckets) == 0 {
return
}
byEpoch := make(map[int64]*model.OpenFlareAccessLogBucketRow, len(buckets))
for _, row := range buckets {
if row == nil {
continue
}
byEpoch[row.BucketEpoch] = row
}
for index := range points {
epoch := points[index].BucketStartedAt.Unix()
if row, ok := byEpoch[epoch]; ok {
points[index].BytesProvided = row.BytesSent
points[index].BytesReceived = row.RequestLength
}
}
}
type accessLogHourBytes struct {
provided int64
received int64
}
func analyticsListAccessLogHourlyBytes(ctx context.Context, nodeID string, since time.Time) (map[int64]accessLogHourBytes, error) {
rows, err := repository.ListOpenFlareAccessLogHourlySince(ctx, nodeID, since)
if err != nil {
return nil, err
}
out := make(map[int64]accessLogHourBytes)
for _, row := range rows {
if row == nil {
continue
}
key := row.Hour.UTC().Truncate(time.Hour).Unix()
cur := out[key]
cur.provided += row.BytesSent
cur.received += row.RequestLength
out[key] = cur
}
return out, nil
}
// BuildTrafficTrendPointsFromHourly builds 24h traffic trend buckets from hourly rollups.
// UniqueVisitorCount is left at 0: hourly UV is not summed (use TrafficSummary for exact UV).
func BuildTrafficTrendPointsFromHourly(now time.Time, hourly []*model.OpenFlareTrafficHourly) []TrafficTrendPoint {
start := trendWindowStart(now)
points := make([]TrafficTrendPoint, observabilityTrendBuckets)
for index := range points {
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
}
for _, row := range hourly {
if row == nil {
continue
}
index, ok := trendBucketIndex(row.Hour, start)
if !ok {
continue
}
points[index].RequestCount += row.RequestCount
points[index].ErrorCount += row.ErrorCount
}
return points
}
// BuildCapacityTrendPoints builds 24h capacity trend buckets.
func BuildCapacityTrendPoints(now time.Time, snapshots []*model.OpenFlareMetricSnapshot) []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
}
// BuildCapacityTrendPointsFromHourly builds 24h capacity trend buckets from hourly aggregates.
func BuildCapacityTrendPointsFromHourly(now time.Time, hourly []*model.OpenFlareMetricHourly) []CapacityTrendPoint {
start := trendWindowStart(now)
points := make([]CapacityTrendPoint, observabilityTrendBuckets)
for index := range points {
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
}
for _, row := range hourly {
if row == nil {
continue
}
index, ok := trendBucketIndex(row.Hour, start)
if !ok {
continue
}
points[index].AverageCPUUsagePercent = row.AverageCPUUsagePercent
points[index].AverageMemoryUsagePercent = row.AverageMemoryUsagePercent
points[index].ReportedNodes = row.ReportedNodes
}
return points
}
func emptyNetworkTrendPoints(now time.Time) []NetworkTrendPoint {
start := trendWindowStart(now)
points := make([]NetworkTrendPoint, observabilityTrendBuckets)
for index := range points {
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
}
return points
}
// BuildDiskIOTrendPoints builds 24h disk IO trend buckets.
func BuildDiskIOTrendPoints(now time.Time, snapshots []*model.OpenFlareMetricSnapshot) []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)
})
previousByNode := make(map[string]diskCounterState, len(snapshots))
for _, snapshot := range snapshots {
nodeKey := snapshot.NodeID
if nodeKey == "" {
nodeKey = unknownTrendNodeKey
}
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
}
points[index].DiskReadBytes += nonNegativeDelta(snapshot.DiskReadBytes, previous.read)
points[index].DiskWriteBytes += nonNegativeDelta(snapshot.DiskWriteBytes, previous.write)
if snapshot.NodeID != "" {
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
}
}
for index := range points {
points[index].ReportedNodes = len(accumulators[index].nodes)
}
return points
}
// BuildDiskIOTrendPointsFromHourly builds 24h disk IO trend buckets from hourly aggregates.
func BuildDiskIOTrendPointsFromHourly(now time.Time, hourly []*model.OpenFlareMetricHourly) []DiskIOTrendPoint {
start := trendWindowStart(now)
points := make([]DiskIOTrendPoint, observabilityTrendBuckets)
for index := range points {
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
}
for _, row := range hourly {
if row == nil {
continue
}
index, ok := trendBucketIndex(row.Hour, start)
if !ok {
continue
}
points[index].DiskReadBytes += row.DiskReadBytes
points[index].DiskWriteBytes += row.DiskWriteBytes
points[index].ReportedNodes = row.ReportedNodes
}
return points
}
func nonNegativeDelta(current int64, previous int64) int64 {
delta := current - previous
if delta < 0 {
return 0
}
return delta
}
func latestMetricSnapshot(snapshots []*model.OpenFlareMetricSnapshot) *model.OpenFlareMetricSnapshot {
var latest *model.OpenFlareMetricSnapshot
for _, snapshot := range snapshots {
if snapshot == nil {
continue
}
if latest == nil || snapshot.CapturedAt.After(latest.CapturedAt) {
latest = snapshot
}
}
return latest
}
func matchEdgeHealth(
capturedAt time.Time,
health []*model.OpenFlareEdgeHealth,
) *model.OpenFlareEdgeHealth {
var matched *model.OpenFlareEdgeHealth
bestDelta := metricSnapshotEdgeHealthMatchWindow + time.Second
for _, row := range health {
if row == nil {
continue
}
delta := capturedAt.Sub(row.CapturedAt)
if delta < 0 {
delta = -delta
}
if delta > metricSnapshotEdgeHealthMatchWindow {
continue
}
if matched == nil || delta < bestDelta {
matched = row
bestDelta = delta
}
}
return matched
}
// LatestMetricSnapshotsByNode returns the latest snapshot per node.
func LatestMetricSnapshotsByNode(snapshots []*model.OpenFlareMetricSnapshot) map[string]*model.OpenFlareMetricSnapshot {
result := make(map[string]*model.OpenFlareMetricSnapshot, 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
}
// ActiveHealthEventsByNode groups active health events by node id.
func ActiveHealthEventsByNode(events []*model.OpenFlareHealthEvent) map[string][]*model.OpenFlareHealthEvent {
result := make(map[string][]*model.OpenFlareHealthEvent)
for _, event := range events {
if event == nil || event.NodeID == "" {
continue
}
result[event.NodeID] = append(result[event.NodeID], event)
}
return result
}
// Percentage returns used/total as a percentage.
func Percentage(used int64, total int64) float64 {
if used <= 0 || total <= 0 {
return 0
}
return (float64(used) / float64(total)) * percentageMultiplier
}
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
}
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,168 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
)
func TestBuildTrafficTrendPointsFromHourlyBucketsByHour(t *testing.T) {
now := time.Date(2026, 7, 2, 15, 30, 0, 0, time.UTC)
hourly := []*model.OpenFlareTrafficHourly{
{
NodeID: "node-a",
Hour: now.Add(-2 * time.Hour).Truncate(time.Hour),
RequestCount: 12,
ErrorCount: 1,
UniqueVisitorCount: 4,
},
}
points := BuildTrafficTrendPointsFromHourly(now, hourly)
if len(points) != observabilityTrendBuckets {
t.Fatalf("BuildTrafficTrendPointsFromHourly() len = %d, want %d", len(points), observabilityTrendBuckets)
}
// Hourly UV must not be summed into trend points.
for _, point := range points {
if point.UniqueVisitorCount != 0 {
t.Fatalf("UniqueVisitorCount = %d, want 0 on hourly path", point.UniqueVisitorCount)
}
}
index, ok := trendBucketIndex(now.Add(-2*time.Hour).Truncate(time.Hour), trendWindowStart(now))
if !ok {
t.Fatal("expected valid bucket index")
}
if points[index].RequestCount != 12 {
t.Fatalf("request_count = %d, want 12", points[index].RequestCount)
}
if points[index].ErrorCount != 1 {
t.Fatalf("error_count = %d, want 1", points[index].ErrorCount)
}
}
func TestBuildMetricSnapshotViewsMergesEdgeHealthConnections(t *testing.T) {
t.Parallel()
capturedAt := time.Date(2026, 6, 19, 12, 0, 0, 0, time.UTC)
snapshots := []*model.OpenFlareMetricSnapshot{
{
ID: 1,
NodeID: "node-a",
CapturedAt: capturedAt,
CPUUsagePercent: 12.5,
},
}
edgeHealth := []*model.OpenFlareEdgeHealth{
{
NodeID: "node-a",
CapturedAt: capturedAt.Add(5 * time.Second),
Status: "healthy",
Connections: 7,
},
}
views := BuildMetricSnapshotViews(snapshots, edgeHealth)
if len(views) != 1 {
t.Fatalf("BuildMetricSnapshotViews() len = %d, want 1", len(views))
}
if views[0].OpenrestyConnections != 7 {
t.Fatalf("OpenrestyConnections = %d, want 7", views[0].OpenrestyConnections)
}
}
func TestBuildTrafficWindowSummaryFromAccessLogsNilWithoutData(t *testing.T) {
t.Parallel()
// Without an access-log store / data, summary is nil.
if summary := buildTrafficWindowSummaryFromAccessLogs(t.Context(), "missing", time.Now().Add(-time.Hour), time.Now()); summary != nil {
t.Fatalf("buildTrafficWindowSummaryFromAccessLogs() = %#v, want nil", summary)
}
}
func TestBuildCapacityTrendPointsFromHourlyFillsBuckets(t *testing.T) {
t.Parallel()
now := time.Date(2026, 7, 10, 9, 30, 0, 0, time.UTC)
hourly := []*model.OpenFlareMetricHourly{
{
Hour: now.Add(-3 * time.Hour).Truncate(time.Hour),
AverageCPUUsagePercent: 42.5,
AverageMemoryUsagePercent: 61.2,
ReportedNodes: 1,
},
{
Hour: now.Truncate(time.Hour),
AverageCPUUsagePercent: 12.0,
AverageMemoryUsagePercent: 50.0,
ReportedNodes: 2,
},
}
points := BuildCapacityTrendPointsFromHourly(now, hourly)
if len(points) != observabilityTrendBuckets {
t.Fatalf("len = %d, want %d", len(points), observabilityTrendBuckets)
}
if points[len(points)-4].AverageCPUUsagePercent != 42.5 {
t.Fatalf("hour-3 cpu = %v, want 42.5", points[len(points)-4].AverageCPUUsagePercent)
}
if points[len(points)-1].ReportedNodes != 2 {
t.Fatalf("current hour reported_nodes = %d, want 2", points[len(points)-1].ReportedNodes)
}
}
func TestEmptyNetworkTrendPointsHas24Buckets(t *testing.T) {
t.Parallel()
now := time.Date(2026, 7, 10, 9, 30, 0, 0, time.UTC)
points := emptyNetworkTrendPoints(now)
if len(points) != observabilityTrendBuckets {
t.Fatalf("len(points) = %d, want %d", len(points), observabilityTrendBuckets)
}
if !points[0].BucketStartedAt.Before(points[len(points)-1].BucketStartedAt) {
t.Fatalf("bucket order invalid: first=%v last=%v", points[0].BucketStartedAt, points[len(points)-1].BucketStartedAt)
}
}
func TestBuildHealthSummaryUsesTrafficSummary(t *testing.T) {
t.Parallel()
snapshot := &model.OpenFlareMetricSnapshot{
CPUUsagePercent: 10,
MemoryUsedBytes: 1,
MemoryTotalBytes: 10,
}
traffic := &TrafficWindowSummary{
RequestCount: 200,
ErrorCount: 20, // 10% error rate
}
summary := buildHealthSummary(snapshot, traffic, nil)
if !summary.HasTrafficRisk {
t.Fatal("HasTrafficRisk = false, want true for 10% error rate with >=100 requests")
}
if summary.HasCapacityRisk {
t.Fatal("HasCapacityRisk = true, want false")
}
}
func TestBuildDiskIOTrendPointsFromHourlyFillsBuckets(t *testing.T) {
t.Parallel()
now := time.Date(2026, 7, 10, 9, 30, 0, 0, time.UTC)
hourly := []*model.OpenFlareMetricHourly{
{
Hour: now.Add(-1 * time.Hour).Truncate(time.Hour),
DiskReadBytes: 1024,
DiskWriteBytes: 2048,
ReportedNodes: 1,
},
}
points := BuildDiskIOTrendPointsFromHourly(now, hourly)
prev := points[len(points)-2]
if prev.DiskReadBytes != 1024 || prev.DiskWriteBytes != 2048 {
t.Fatalf("previous hour disk io = %#v, want read=1024 write=2048", prev)
}
}
@@ -0,0 +1,65 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package chwriter
import (
"sync"
"time"
)
const dedupTTL = 2 * time.Minute
type dedupSet struct {
mu sync.Mutex
keys map[string]time.Time
lastCleanup time.Time
}
func newDedupSet() *dedupSet {
return &dedupSet{
keys: make(map[string]time.Time),
lastCleanup: time.Now(),
}
}
// markIfNew records key when it has not been seen within dedupTTL.
func (s *dedupSet) markIfNew(key string) bool {
if s == nil || key == "" {
return false
}
now := time.Now()
s.mu.Lock()
defer s.mu.Unlock()
s.cleanupExpiredLocked(now)
if expiresAt, exists := s.keys[key]; exists && now.Before(expiresAt) {
return false
}
s.keys[key] = now.Add(dedupTTL)
return true
}
// unmark removes a key so a later enqueue or flush retry may accept it again.
func (s *dedupSet) unmark(key string) {
if s == nil || key == "" {
return
}
s.mu.Lock()
defer s.mu.Unlock()
delete(s.keys, key)
}
func (s *dedupSet) cleanupExpiredLocked(now time.Time) {
if now.Sub(s.lastCleanup) < 30*time.Second {
return
}
for existing, expiresAt := range s.keys {
if now.After(expiresAt) {
delete(s.keys, existing)
}
}
s.lastCleanup = now
}
@@ -0,0 +1,186 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package chwriter
import (
"context"
"errors"
"sync"
"testing"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/pkg/batchwriter"
)
func TestDedupSetMarkIfNew(t *testing.T) {
t.Parallel()
set := newDedupSet()
if !set.markIfNew("node-a|1") {
t.Fatal("markIfNew() = false, want true on first key")
}
if set.markIfNew("node-a|1") {
t.Fatal("markIfNew() = true, want false on duplicate key")
}
if !set.markIfNew("node-b|1") {
t.Fatal("markIfNew() = false, want true on different key")
}
if set.markIfNew("") {
t.Fatal("markIfNew() = true, want false on empty key")
}
}
func TestDedupSetUnmarkAllowsRetry(t *testing.T) {
t.Parallel()
set := newDedupSet()
if !set.markIfNew("k") {
t.Fatal("markIfNew() = false, want true")
}
set.unmark("k")
if !set.markIfNew("k") {
t.Fatal("markIfNew() after unmark = false, want true")
}
}
func TestQueueWithDedupDoesNotMarkWhenEnqueueFails(t *testing.T) {
t.Parallel()
cfg := batchwriter.DefaultConfig()
cfg.QueueSize = 1
cfg.MaxBatchSize = 10
cfg.FlushInterval = time.Hour
// Block the worker so the queue stays full after one enqueue.
block := make(chan struct{})
writer, err := batchwriter.New[int](cfg, func(context.Context, []int) error {
<-block
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
close(block)
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
// Fill the channel buffer (and the worker's current receive slot may empty one).
// Keep enqueueing until full so subsequent queueWithDedup fails.
for i := 0; i < cfg.QueueSize+2; i++ {
_ = writer.TryEnqueue(i)
if writer.IsFull() {
break
}
}
if !writer.IsFull() {
t.Fatal("writer not full after filling; cannot test enqueue failure path")
}
dedup := newDedupSet()
queueWithDedup(writer, dedup, "dedup-key", 99)
// Key must not remain marked after failed enqueue.
if !dedup.markIfNew("dedup-key") {
t.Fatal("dedup key still marked after failed enqueue; want unmark")
}
}
func TestQueueWithDedupMarksOnlyOnSuccess(t *testing.T) {
t.Parallel()
cfg := batchwriter.DefaultConfig()
cfg.MaxBatchSize = 100
cfg.FlushInterval = time.Hour
writer, err := batchwriter.New[int](cfg, func(context.Context, []int) error { return nil })
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
dedup := newDedupSet()
queueWithDedup(writer, dedup, "ok-key", 1)
if dedup.markIfNew("ok-key") {
t.Fatal("markIfNew() = true after successful enqueue, want false (key marked)")
}
}
func TestFlushErrorHandlerUnmarksKeys(t *testing.T) {
t.Parallel()
dedup := newDedupSet()
flushErr := errors.New("ch down")
var (
mu sync.Mutex
errCount int
)
cfg := batchwriter.Config{
Name: "test_obs",
QueueSize: 10,
MaxBatchSize: 1,
FlushInterval: time.Hour,
}
keyFn := func(s analyticsmodel.NodeMetricSnapshot) string {
return metricSnapshotKey(s)
}
writer, err := batchwriter.New(
cfg,
func(context.Context, []analyticsmodel.NodeMetricSnapshot) error { return flushErr },
batchwriter.WithFlushErrorHandler[analyticsmodel.NodeMetricSnapshot](func(_ context.Context, items []analyticsmodel.NodeMetricSnapshot, err error) {
mu.Lock()
errCount++
mu.Unlock()
for _, item := range items {
dedup.unmark(keyFn(item))
}
}),
)
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
item := analyticsmodel.NodeMetricSnapshot{
NodeID: "n1",
CapturedAt: time.Unix(1, 0).UTC(),
}
key := keyFn(item)
if !dedup.markIfNew(key) {
t.Fatal("markIfNew failed")
}
if !writer.TryEnqueue(item) {
t.Fatal("TryEnqueue failed")
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := errCount >= 1
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
if !dedup.markIfNew(key) {
t.Fatal("key still marked after flush error unmark; want available for retry")
}
}
@@ -0,0 +1,88 @@
//go:build live_ch
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package chwriter_test
import (
"context"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/observability/chwriter"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// Run with Docker ClickHouse + config.yaml:
//
// go test -tags live_ch ./internal/apps/openflare/chwriter -run TestLiveAppWritePath -count=1 -timeout 2m
func TestLiveAppWritePath(t *testing.T) {
if db.ChConn == nil {
t.Skip("ClickHouse connection not ready")
}
ctx := context.Background()
chwriter.Init(ctx)
now := time.Now().UTC()
nodeID := "e2e-app-write-" + now.Format("150405")
if err := repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: nodeID,
CapturedAt: now,
CPUUsagePercent: 33.3,
MemoryUsedBytes: 111,
MemoryTotalBytes: 1000,
StorageUsedBytes: 222,
StorageTotalBytes: 2000,
DiskReadBytes: 10,
DiskWriteBytes: 20,
NetworkRxBytes: 30,
NetworkTxBytes: 40,
}); err != nil {
t.Fatalf("InsertOpenFlareMetricSnapshot: %v", err)
}
deadline := time.Now().Add(45 * time.Second)
var found bool
for time.Now().Before(deadline) {
rows, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, nodeID, now.Add(-time.Minute), 10)
if err != nil {
t.Fatalf("ListOpenFlareMetricSnapshotsSince: %v", err)
}
if len(rows) > 0 {
found = true
t.Logf("found snapshot id=%d cpu=%.1f after flush", rows[0].ID, rows[0].CPUUsagePercent)
break
}
time.Sleep(2 * time.Second)
}
if !found {
t.Fatal("metric snapshot not visible in ClickHouse after flush wait")
}
latest, err := repository.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", now.Add(-time.Hour))
if err != nil {
t.Fatalf("ListOpenFlareLatestMetricSnapshotsSince: %v", err)
}
var latestOK bool
for _, row := range latest {
if row != nil && row.NodeID == nodeID {
latestOK = true
break
}
}
if !latestOK {
t.Fatalf("latest-per-node query missing node %s (rows=%d)", nodeID, len(latest))
}
stats := chwriter.WriterStats()
if len(stats) == 0 {
t.Fatal("WriterStats empty after Init")
}
for _, s := range stats {
t.Logf("writer %s running=%v depth=%d drops=%d flush_err=%d", s.Name, s.Running, s.Depth, s.Drops, s.FlushErrors)
}
}
@@ -0,0 +1,400 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package chwriter queues OpenFlare ClickHouse writes and flushes them through
// internal/infra/persistence/batchwriter with per-table writer instances.
package chwriter
import (
"context"
"fmt"
"sync"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
"Wavelet/pkg/batchwriter"
"Wavelet/pkg/logger"
)
const (
// Observability traffic is sparse (heartbeat ~10s/node). Prefer larger batches to
// cut ClickHouse parts/merges; MaxFlushWait bounds visibility lag for single-node labs.
observabilityQueueSize = 5_000
observabilityMaxBatchSize = 500
observabilityMinBatchSize = 20
observabilityFlushEvery = 10 * time.Second
observabilityMaxFlushWait = 30 * time.Second
nodeAccessLogQueueSize = 10_000
nodeAccessLogMaxBatchSize = 1_000
nodeAccessLogMinBatchSize = 50
nodeAccessLogFlushEvery = 2 * time.Second
nodeAccessLogMaxFlushWait = 5 * time.Second
// flushAttempts is total tries (1 initial + short retries) before giving up a batch.
flushAttempts = 2
flushRetryBackoff = 50 * time.Millisecond
)
var (
initOnce sync.Once
metricSnapshotWriter *batchwriter.Writer[analyticsmodel.NodeMetricSnapshot]
edgeHealthWriter *batchwriter.Writer[analyticsmodel.NodeEdgeHealth]
frpsWriter *batchwriter.Writer[analyticsmodel.NodeObsFrps]
frpcWriter *batchwriter.Writer[analyticsmodel.NodeObsFrpc]
nodeAccessLogWriter *batchwriter.Writer[analyticsmodel.NodeAccessLog]
metricSnapshotDedup *dedupSet
edgeHealthDedup *dedupSet
frpsDedup *dedupSet
frpcDedup *dedupSet
)
// Init starts OpenFlare log batch writers. Safe to call multiple times.
// Writers always initialize regardless of ClickHouse.enabled; the active log
// store is resolved via logstore at flush time (PG/SQLite when CH is not active).
func Init(ctx context.Context) {
initOnce.Do(func() {
metricSnapshotDedup = newDedupSet()
edgeHealthDedup = newDedupSet()
frpsDedup = newDedupSet()
frpcDedup = newDedupSet()
metricSnapshotWriter = mustNewObservabilityWriter(
"metric_snapshots",
withFlushRetries(flushNodeMetricSnapshots),
metricSnapshotDedup,
metricSnapshotKey,
)
edgeHealthWriter = mustNewObservabilityWriter(
"edge_health",
withFlushRetries(flushNodeEdgeHealth),
edgeHealthDedup,
edgeHealthKey,
)
frpsWriter = mustNewObservabilityWriter(
"frps_obs",
withFlushRetries(flushNodeObsFrps),
frpsDedup,
frpsKey,
)
frpcWriter = mustNewObservabilityWriter(
"frpc_obs",
withFlushRetries(flushNodeObsFrpc),
frpcDedup,
frpcKey,
)
nodeAccessLogWriter = mustNewNodeAccessLogWriter()
metricSnapshotWriter.Start(ctx)
edgeHealthWriter.Start(ctx)
frpsWriter.Start(ctx)
frpcWriter.Start(ctx)
nodeAccessLogWriter.Start(ctx)
wireModelInsertHooks()
})
}
// Stop drains all OpenFlare ClickHouse writers.
func Stop(ctx context.Context) error {
if !running() {
return nil
}
var firstErr error
for _, writer := range []batchStopper{
metricSnapshotWriter,
edgeHealthWriter,
frpsWriter,
frpcWriter,
nodeAccessLogWriter,
} {
if writer == nil {
continue
}
if err := writer.Stop(ctx); err != nil && firstErr == nil {
firstErr = err
}
}
return firstErr
}
// Drain 等待所有 OpenFlare 日志 writer 的在途批次落库:轮询队列 Depth 归零后
// 再保持一个最大 flush 周期(observabilityFlushEvery)持续为空才返回;
// 不停止 writer(迁移冻结后由 ensureWritable 拒绝新写入)。未初始化时直接返回 nil。
func Drain(ctx context.Context) error {
return drainWriters(ctx, WriterStats, observabilityFlushEvery)
}
// drainWriters 轮询 stats 直至所有队列 Depth=0 并持续 quietPeriod 无新积压。
func drainWriters(ctx context.Context, stats func() []batchwriter.Stats, quietPeriod time.Duration) error {
if !running() {
return nil
}
ticker := time.NewTicker(drainPollInterval)
defer ticker.Stop()
var quietSince time.Time
for {
if allDepthZero(stats()) {
if quietSince.IsZero() {
quietSince = time.Now()
} else if time.Since(quietSince) >= quietPeriod {
return nil
}
} else {
quietSince = time.Time{}
}
select {
case <-ctx.Done():
return ctx.Err()
case <-ticker.C:
}
}
}
// drainPollInterval 队列轮询间隔。
const drainPollInterval = 50 * time.Millisecond
func allDepthZero(stats []batchwriter.Stats) bool {
for _, s := range stats {
if s.Depth > 0 {
return false
}
}
return true
}
// WriterStats returns queue depth and failure counters for all OpenFlare writers.
func WriterStats() []batchwriter.Stats {
writers := []statsProvider{
metricSnapshotWriter,
edgeHealthWriter,
frpsWriter,
frpcWriter,
nodeAccessLogWriter,
}
out := make([]batchwriter.Stats, 0, len(writers))
for _, w := range writers {
if w == nil {
continue
}
out = append(out, w.Stats())
}
return out
}
// QueueMetricSnapshot enqueues a metric snapshot for asynchronous flush.
func QueueMetricSnapshot(snapshot analyticsmodel.NodeMetricSnapshot) {
queueWithDedup(metricSnapshotWriter, metricSnapshotDedup, metricSnapshotKey(snapshot), snapshot)
}
// QueueEdgeHealth enqueues an L2 edge health snapshot for asynchronous flush.
func QueueEdgeHealth(row analyticsmodel.NodeEdgeHealth) {
queueWithDedup(edgeHealthWriter, edgeHealthDedup, edgeHealthKey(row), row)
}
// QueueFrpsObservation enqueues an FRPS observation for asynchronous flush.
func QueueFrpsObservation(observation analyticsmodel.NodeObsFrps) {
queueWithDedup(frpsWriter, frpsDedup, frpsKey(observation), observation)
}
// QueueFrpcObservation enqueues an FRPC observation for asynchronous flush.
func QueueFrpcObservation(observation analyticsmodel.NodeObsFrpc) {
queueWithDedup(frpcWriter, frpcDedup, frpcKey(observation), observation)
}
// QueueNodeAccessLogs enqueues node access logs for asynchronous flush.
func QueueNodeAccessLogs(logs []analyticsmodel.NodeAccessLog) {
if nodeAccessLogWriter == nil || len(logs) == 0 {
return
}
for _, logItem := range logs {
nodeAccessLogWriter.TryEnqueue(logItem)
}
}
func queueWithDedup[T any](writer *batchwriter.Writer[T], dedup *dedupSet, key string, item T) {
if writer == nil {
return
}
// Mark first so concurrent duplicates still collapse; release on enqueue failure
// so a full queue does not permanently suppress the item.
if !dedup.markIfNew(key) {
return
}
if !writer.TryEnqueue(item) {
dedup.unmark(key)
}
}
func mustNewObservabilityWriter[T any](
name string,
flush batchwriter.FlushFunc[T],
dedup *dedupSet,
keyFn func(T) string,
) *batchwriter.Writer[T] {
cfg := batchwriter.Config{
Name: name,
QueueSize: observabilityQueueSize,
MaxBatchSize: observabilityMaxBatchSize,
MinBatchSize: observabilityMinBatchSize,
FlushInterval: observabilityFlushEvery,
MaxFlushWait: observabilityMaxFlushWait,
}
writer, err := batchwriter.New(
cfg,
flush,
withObservabilityDropHandler[T](name),
batchwriter.WithFlushErrorHandler[T](func(ctx context.Context, items []T, err error) {
logger.ErrorF(ctx, "[OpenFlare] flush %s failed (batch=%d): %v", name, len(items), err)
if dedup == nil || keyFn == nil {
return
}
for _, item := range items {
dedup.unmark(keyFn(item))
}
}),
)
if err != nil {
panic(fmt.Sprintf("openflare chwriter %s: %v", name, err))
}
return writer
}
func mustNewNodeAccessLogWriter() *batchwriter.Writer[analyticsmodel.NodeAccessLog] {
cfg := batchwriter.Config{
Name: "node_access_logs",
QueueSize: nodeAccessLogQueueSize,
MaxBatchSize: nodeAccessLogMaxBatchSize,
MinBatchSize: nodeAccessLogMinBatchSize,
FlushInterval: nodeAccessLogFlushEvery,
MaxFlushWait: nodeAccessLogMaxFlushWait,
}
writer, err := batchwriter.New[analyticsmodel.NodeAccessLog](
cfg,
withFlushRetries(flushNodeAccessLogs),
batchwriter.WithDropHandler[analyticsmodel.NodeAccessLog](func(item analyticsmodel.NodeAccessLog) {
logger.WarnF(context.Background(), "[OpenFlare] node access log queue full, dropping log for node %s path %s", item.NodeID, item.Path)
}),
batchwriter.WithFlushErrorHandler[analyticsmodel.NodeAccessLog](func(ctx context.Context, items []analyticsmodel.NodeAccessLog, err error) {
logger.ErrorF(ctx, "[OpenFlare] flush node access logs failed (batch=%d): %v", len(items), err)
}),
)
if err != nil {
panic(fmt.Sprintf("openflare chwriter node_access_logs: %v", err))
}
return writer
}
func withObservabilityDropHandler[T any](name string) batchwriter.Option[T] {
return batchwriter.WithDropHandler(func(_ T) {
logger.WarnF(context.Background(), "[OpenFlare] %s queue full, dropping observability item", name)
})
}
// withFlushRetries wraps a flush function with a short retry to ride out brief CH blips.
func withFlushRetries[T any](flush batchwriter.FlushFunc[T]) batchwriter.FlushFunc[T] {
return func(ctx context.Context, items []T) error {
var err error
for attempt := 1; attempt <= flushAttempts; attempt++ {
err = flush(ctx, items)
if err == nil {
return nil
}
if attempt == flushAttempts {
break
}
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(flushRetryBackoff * time.Duration(attempt)):
}
}
return err
}
}
func wireModelInsertHooks() {
logstore.SetObservabilityHooks(logstore.ObservabilityHooks{
QueueMetricSnapshot: QueueMetricSnapshot,
QueueEdgeHealth: QueueEdgeHealth,
QueueNodeObsFrps: QueueFrpsObservation,
QueueNodeObsFrpc: QueueFrpcObservation,
})
logstore.SetAccessLogHooks(logstore.AccessLogHooks{
QueueNodeAccessLogs: QueueNodeAccessLogs,
})
}
// 以下 flush 函数作为 batchwriter 的落库目标:激活库由 logstore 在 flush 时决定。
func flushNodeMetricSnapshots(ctx context.Context, rows []analyticsmodel.NodeMetricSnapshot) error {
s, err := logstore.Active(ctx)
if err != nil {
return err
}
return s.Observability.BatchInsertNodeMetricSnapshots(ctx, rows)
}
func flushNodeEdgeHealth(ctx context.Context, rows []analyticsmodel.NodeEdgeHealth) error {
s, err := logstore.Active(ctx)
if err != nil {
return err
}
return s.Observability.BatchInsertNodeEdgeHealth(ctx, rows)
}
func flushNodeObsFrps(ctx context.Context, rows []analyticsmodel.NodeObsFrps) error {
s, err := logstore.Active(ctx)
if err != nil {
return err
}
return s.Observability.BatchInsertNodeObsFrps(ctx, rows)
}
func flushNodeObsFrpc(ctx context.Context, rows []analyticsmodel.NodeObsFrpc) error {
s, err := logstore.Active(ctx)
if err != nil {
return err
}
return s.Observability.BatchInsertNodeObsFrpc(ctx, rows)
}
func flushNodeAccessLogs(ctx context.Context, rows []analyticsmodel.NodeAccessLog) error {
s, err := logstore.Active(ctx)
if err != nil {
return err
}
return s.AccessLogs.BatchInsertNodeAccessLogs(ctx, rows)
}
func metricSnapshotKey(snapshot analyticsmodel.NodeMetricSnapshot) string {
return fmt.Sprintf("%s|%d", snapshot.NodeID, snapshot.CapturedAt.UTC().UnixNano())
}
func edgeHealthKey(row analyticsmodel.NodeEdgeHealth) string {
return fmt.Sprintf("%s|%d", row.NodeID, row.CapturedAt.UTC().UnixNano())
}
func frpsKey(observation analyticsmodel.NodeObsFrps) string {
return fmt.Sprintf("%s|%d", observation.NodeID, observation.CapturedAt.UTC().UnixNano())
}
func frpcKey(observation analyticsmodel.NodeObsFrpc) string {
return fmt.Sprintf("%s|%d", observation.NodeID, observation.CapturedAt.UTC().UnixNano())
}
type batchStopper interface {
Stop(ctx context.Context) error
}
type statsProvider interface {
Stats() batchwriter.Stats
}
func running() bool {
return metricSnapshotWriter != nil && metricSnapshotWriter.Running()
}
@@ -0,0 +1,9 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package observability defines shared error messages for observability operations.
package observability
const (
errInvalidStatusCode = "status_code 必须为 100-599 之间的整数"
)
@@ -0,0 +1,398 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"context"
"encoding/json"
"errors"
"fmt"
"time"
"Wavelet/openflare/plugins/server/domain/observability/chwriter"
"Wavelet/openflare/plugins/server/kernel/model"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
"Wavelet/openflare/plugins/server/kernel/task"
"Wavelet/pkg/logger"
)
const copyBatchSize = 1000
// 迁移目标库名常量(normalizeTarget 归一化后的取值)。
const (
targetPostgres = "postgres"
targetSQLite = "sqlite"
targetClickHouse = "clickhouse"
)
type logDBSwitchPayload struct {
Target string `json:"target"`
}
// LogDBSwitchHandler 切换日志数据库任务处理器。
type LogDBSwitchHandler struct{}
// ValidatePayload 校验并规范化参数。
func (h *LogDBSwitchHandler) ValidatePayload(payload []byte) ([]byte, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf("参数解析失败: %w", err)
}
p.Target = normalizeTarget(p.Target)
if !validTarget(p.Target) {
return nil, fmt.Errorf("目标日志库不合法: %s", p.Target)
}
out, err := json.Marshal(p)
if err != nil {
return nil, err
}
return out, nil
}
func normalizeTarget(v string) string {
switch v {
case targetPostgres, "postgresql":
return targetPostgres
case targetSQLite, "sqlite3":
return targetSQLite
case targetClickHouse, "ch":
return targetClickHouse
}
return v
}
func validTarget(v string) bool {
return v == targetPostgres || v == targetSQLite || v == targetClickHouse
}
// Execute 执行迁移。
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf("参数解析失败: %w", err)
}
p.Target = normalizeTarget(p.Target)
if err := validateSwitch(ctx, p.Target); err != nil {
return nil, err
}
source, err := currentLogDatabase(ctx)
if err != nil {
task.AppendLog(ctx, "读取日志主库失败: %v", err)
return nil, err
}
task.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
// 设置迁移冻结标记(置位后由 ensureWritable 拒绝新写入)。
if err := setMigrationFlag(ctx, "migrating"); err != nil {
return nil, err
}
// 失败也清除(SaveOrUpdateSystemConfig 会失效 RAM 缓存并广播),保持源库可写。
defer func() {
if err := setMigrationFlag(ctx, ""); err != nil {
logger.ErrorF(ctx, "清除日志迁移冻结标记失败: %v", err)
}
}()
// 冻结标记置位后再排空在途批次(chwriter + 用户访问日志 writer),
// 保证排空完成后不再有新批次进入源库。
if err := drainLogWriters(ctx); err != nil {
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
}
src, err := logstore.Active(ctx)
if err != nil {
return nil, err
}
dst, err := buildTargetStore(ctx, p.Target)
if err != nil {
return nil, err
}
// 清空目标库日志表(幂等重试前提)。
if err := clearTargetLogTables(ctx, dst); err != nil {
return nil, err
}
// PG 目标:按源库时间范围预建分区,避免历史数据复制报 "no partition of relation found"。
if err := ensureTargetPartitions(ctx, src, dst, p.Target); err != nil {
return nil, err
}
// 逐表复制(6 张日志表)。
if err := copyAccessLogs(ctx, src, dst); err != nil {
return nil, err
}
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
return nil, err
}
if err := copyObservability(ctx, src, dst); err != nil {
return nil, err
}
// 翻转主库标记。
if err := flipLogDatabase(ctx, p.Target); err != nil {
return nil, err
}
task.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
return &task.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
}
func validateSwitch(ctx context.Context, target string) error {
source, err := currentLogDatabase(ctx)
if err != nil {
return err
}
if source == target {
return errors.New("目标日志库与当前日志库相同,无需迁移")
}
switch target {
case "clickhouse":
if !runtimeconfig.ClickHouseEnabled() {
return errors.New("ClickHouse 未启用,无法迁移到 ClickHouse")
}
case "postgres":
if !runtimeconfig.DatabaseEnabled() {
return errors.New("PostgreSQL 未启用(当前主库为 SQLite),无法迁移到 PostgreSQL")
}
case "sqlite":
if runtimeconfig.DatabaseEnabled() {
return errors.New("当前主库为 PostgreSQL,日志库不能设置为 SQLite")
}
}
return nil
}
func currentLogDatabase(ctx context.Context) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase)
if err != nil {
return "", fmt.Errorf("读取日志主库失败: %w", err)
}
if cfg.Value == "" {
return "", errors.New("日志主库配置为空")
}
return cfg.Value, nil
}
// drainLogWriters 等待 chwriter(节点访问日志 + 可观测 4 表)的在途批次全部落库。
// 见设计 §7.2:先排空再冻结。用户访问日志(w_user_access_logs)记录已禁用,无在途批次。
func drainLogWriters(ctx context.Context) error {
return chwriter.Drain(ctx)
}
// setMigrationFlag 写入迁移冻结标记。用 SaveOrUpdateSystemConfig:行缺失时 upsert,
// 并失效 RAM 缓存 + 广播其他节点,保证 logstore.Migrating/resolveDatabase 立即生效。
func setMigrationFlag(ctx context.Context, v string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDBMigration, v)
}
// flipLogDatabase 翻转日志主库。同上用 SaveOrUpdateSystemConfig,确保各进程缓存失效后指向新库。
func flipLogDatabase(ctx context.Context, target string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, target)
}
// buildTargetStore 构造目标库 Store(不经过 Active 缓存,直接 Build)。
// 迁移期间冻结标记已置位,目标库的清空/复制写入必须放行,故使用 BuildForMigration。
func buildTargetStore(ctx context.Context, database string) (*logstore.Store, error) {
return logstore.BuildForMigration(ctx, database)
}
func clearTargetLogTables(ctx context.Context, dst *logstore.Store) error {
// 依次清空 6 张表:AccessLogs.DeleteAll、UserAccessLogs.DeleteAll、Observability.DeleteAll*
// (SQLite/PG 用 DeleteAll;CH 用 TRUNCATE 语义)。
if _, err := dst.AccessLogs.DeleteAll(ctx); err != nil {
return fmt.Errorf("清空目标访问日志失败: %w", err)
}
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
return fmt.Errorf("清空目标用户访问日志失败: %w", err)
}
for _, fn := range []func(context.Context) (int64, error){
dst.Observability.DeleteAllMetricSnapshots,
dst.Observability.DeleteAllEdgeHealth,
dst.Observability.DeleteAllNodeObservationFrps,
dst.Observability.DeleteAllNodeObservationFrpc,
} {
if _, err := fn(ctx); err != nil {
return err
}
}
return nil
}
// ensureTargetPartitions 目标为 PG 时,按源库时间范围(两表合并)预建分区,
// 否则复制历史数据会报 "no partition of relation found";目标非 PG 为 no-op。
func ensureTargetPartitions(ctx context.Context, src, dst *logstore.Store, target string) error {
if target != targetPostgres {
return nil
}
from, to, err := migrationRange(ctx, src)
if err != nil {
return err
}
if from.IsZero() || to.IsZero() {
task.AppendLog(ctx, "源库无日志数据,跳过分区预建")
return nil
}
if err := dst.AccessLogs.EnsurePartitions(ctx, from, to.AddDate(0, 1, 0)); err != nil {
return fmt.Errorf("预建目标 PG 分区失败: %w", err)
}
task.AppendLog(ctx, "已为目标 PG 预建分区 %s ~ %s", from.Format("2006-01"), to.Format("2006-01"))
return nil
}
// migrationRange 合并源库节点访问日志(logged_at)与用户访问日志(created_at)
// 的最小/最大时间;任一表为空时忽略该表。
func migrationRange(ctx context.Context, src *logstore.Store) (time.Time, time.Time, error) {
fromAccess, toAccess, err := src.AccessLogs.MigrationRange(ctx)
if err != nil {
return time.Time{}, time.Time{}, fmt.Errorf("读取源访问日志时间范围失败: %w", err)
}
fromUser, toUser, err := src.UserAccessLogs.MigrationRange(ctx)
if err != nil {
return time.Time{}, time.Time{}, fmt.Errorf("读取源用户访问日志时间范围失败: %w", err)
}
return minTime(fromAccess, fromUser), maxTime(toAccess, toUser), nil
}
func minTime(a, b time.Time) time.Time {
switch {
case a.IsZero():
return b
case b.IsZero():
return a
case a.Before(b):
return a
default:
return b
}
}
func maxTime(a, b time.Time) time.Time {
switch {
case a.IsZero():
return b
case b.IsZero():
return a
case a.After(b):
return a
default:
return b
}
}
// copyAccessLogs 从 src 复制节点访问日志到 dst。
func copyAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
// 注意:迁移期间 src 已冻结,但复制读取不受冻结影响;每批按 id 升序扫描。
var lastID uint64
for {
rows, err := listNodeAccessLogsByID(ctx, src, lastID, copyBatchSize)
if err != nil {
return err
}
if len(rows) == 0 {
break
}
if err := dst.AccessLogs.BatchInsertNodeAccessLogs(ctx, rows); err != nil {
return fmt.Errorf("写入目标访问日志失败(批 %d): %w", lastID, err)
}
task.AppendLog(ctx, "已复制访问日志 %d 条(截至 id=%d)", len(rows), rows[len(rows)-1].ID)
lastID = rows[len(rows)-1].ID
if len(rows) < copyBatchSize {
break
}
}
return nil
}
func listNodeAccessLogsByID(ctx context.Context, src *logstore.Store, afterID uint64, limit int) ([]analyticsmodel.NodeAccessLog, error) {
return src.AccessLogs.ListForMigration(ctx, afterID, limit)
}
// copyUserAccessLogs 从 src 复制用户访问日志到 dst(按 id 升序分批)。
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
var lastID uint64
for {
rows, err := src.UserAccessLogs.ListForMigration(ctx, lastID, copyBatchSize)
if err != nil {
return err
}
if len(rows) == 0 {
return nil
}
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
return fmt.Errorf("写入目标用户访问日志失败(批 %d): %w", lastID, err)
}
lastID = rows[len(rows)-1].ID
task.AppendLog(ctx, "已复制用户访问日志 %d 条(截至 id=%d)", len(rows), lastID)
if len(rows) < copyBatchSize {
return nil
}
}
}
// copyObservability 复制 4 张可观测表,每张表按 id 升序分批复制,
// 以每批最后一条 id 作为下一批游标(不使用 len 近似)。
func copyObservability(ctx context.Context, src, dst *logstore.Store) error {
if err := copyObsTable(ctx, "metric_snapshots",
src.Observability.ListMetricSnapshotsForMigration,
dst.Observability.BatchInsertNodeMetricSnapshots,
lastMetricSnapshotID); err != nil {
return err
}
if err := copyObsTable(ctx, "edge_health",
src.Observability.ListEdgeHealthForMigration,
dst.Observability.BatchInsertNodeEdgeHealth,
lastEdgeHealthID); err != nil {
return err
}
if err := copyObsTable(ctx, "obs_frps",
src.Observability.ListNodeObsFrpsForMigration,
dst.Observability.BatchInsertNodeObsFrps,
lastObsFrpsID); err != nil {
return err
}
if err := copyObsTable(ctx, "obs_frpc",
src.Observability.ListNodeObsFrpcForMigration,
dst.Observability.BatchInsertNodeObsFrpc,
lastObsFrpcID); err != nil {
return err
}
return nil
}
// copyObsTable 按 id 升序分批复制单张可观测表;idOf 返回批内最后一条 id。
func copyObsTable[T any](ctx context.Context, name string,
list func(context.Context, uint64, int) ([]T, error),
insert func(context.Context, []T) error,
idOf func([]T) uint64,
) error {
var lastID uint64
for {
rows, err := list(ctx, lastID, copyBatchSize)
if err != nil {
return fmt.Errorf("复制 %s 失败: %w", name, err)
}
if len(rows) == 0 {
return nil
}
if err := insert(ctx, rows); err != nil {
return fmt.Errorf("复制 %s 失败: %w", name, err)
}
lastID = idOf(rows)
task.AppendLog(ctx, "已复制 %s %d 条(截至 id=%d)", name, len(rows), lastID)
if len(rows) < copyBatchSize {
return nil
}
}
}
func lastMetricSnapshotID(rows []analyticsmodel.NodeMetricSnapshot) uint64 {
return rows[len(rows)-1].ID
}
func lastEdgeHealthID(rows []analyticsmodel.NodeEdgeHealth) uint64 { return rows[len(rows)-1].ID }
func lastObsFrpsID(rows []analyticsmodel.NodeObsFrps) uint64 { return rows[len(rows)-1].ID }
func lastObsFrpcID(rows []analyticsmodel.NodeObsFrpc) uint64 { return rows[len(rows)-1].ID }
@@ -0,0 +1,381 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"context"
"encoding/json"
"errors"
"fmt"
"sync/atomic"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"Wavelet/openflare/plugins/server/kernel/model"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
db "Wavelet/plugins/infra/database"
)
var logDBSwitchDBSeq int64
// newLogDBSwitchDB 构造内存 sqlite 库(含日志 5 表 + 系统配置表)。
func newLogDBSwitchDB(t *testing.T) *gorm.DB {
t.Helper()
dsn := fmt.Sprintf("file:log-db-switch-%d?mode=memory&cache=shared", atomic.AddInt64(&logDBSwitchDBSeq, 1))
gdb, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
require.NoError(t, err)
require.NoError(t, gdb.AutoMigrate(
&model.SystemConfig{},
&analyticsmodel.NodeAccessLog{},
&analyticsmodel.NodeMetricSnapshot{},
&analyticsmodel.NodeEdgeHealth{},
&analyticsmodel.NodeObsFrps{},
&analyticsmodel.NodeObsFrpc{},
&analyticsmodel.UserAccessLog{},
))
return gdb
}
// TestCopyAccessLogsPreservesIDs sqlite→sqlite 模拟:源 store 3 条,目标空库,
// copyAccessLogs 后 ID 保留、数量一致。
func TestCopyAccessLogsPreservesIDs(t *testing.T) {
t.Cleanup(runtimeconfig.Override(false, false))
logstore.ResetForTest()
defer logstore.ResetForTest()
ctx := context.Background()
srcDB := newLogDBSwitchDB(t)
dstDB := newLogDBSwitchDB(t)
db.SetDB(srcDB)
src, err := logstore.Active(ctx) // 无 reader 时按 seed 规则解析为 sqlite
require.NoError(t, err)
db.SetDB(dstDB)
dst, err := logstore.BuildForMigration(ctx, "sqlite")
require.NoError(t, err)
t.Cleanup(func() { db.SetDB(nil) })
now := time.Now().UTC()
rows := []analyticsmodel.NodeAccessLog{
{ID: 101, NodeID: "n1", LoggedAt: now, RemoteAddr: "1.1.1.1", Host: "a.example.com", Path: "/"},
{ID: 202, NodeID: "n2", LoggedAt: now, RemoteAddr: "2.2.2.2", Host: "b.example.com", Path: "/x"},
{ID: 303, NodeID: "n1", LoggedAt: now, RemoteAddr: "3.3.3.3", Host: "c.example.com", Path: "/y"},
}
require.NoError(t, src.AccessLogs.BatchInsertNodeAccessLogs(ctx, rows))
require.NoError(t, copyAccessLogs(ctx, src, dst))
var got []analyticsmodel.NodeAccessLog
require.NoError(t, dstDB.Order("id ASC").Find(&got).Error)
require.Len(t, got, 3)
for i, wantID := range []uint64{101, 202, 303} {
assert.Equal(t, wantID, got[i].ID, "row %d id preserved", i)
}
assert.Equal(t, "n1", got[0].NodeID)
assert.Equal(t, "n2", got[1].NodeID)
assert.Equal(t, "n1", got[2].NodeID)
assert.Equal(t, "1.1.1.1", got[0].RemoteAddr)
// 源库保持不变。
var srcCount int64
require.NoError(t, srcDB.Model(&analyticsmodel.NodeAccessLog{}).Count(&srcCount).Error)
assert.Equal(t, int64(3), srcCount)
}
// TestCopyUserAccessLogsPreservesIDs sqlite→sqlite 模拟:源库用户访问日志按 id 升序
// 复制到目标库,ID 保留、数量一致,且源库保持不变。
func TestCopyUserAccessLogsPreservesIDs(t *testing.T) {
t.Cleanup(runtimeconfig.Override(false, false))
logstore.ResetForTest()
defer logstore.ResetForTest()
ctx := context.Background()
srcDB := newLogDBSwitchDB(t)
dstDB := newLogDBSwitchDB(t)
db.SetDB(srcDB)
src, err := logstore.Active(ctx)
require.NoError(t, err)
db.SetDB(dstDB)
dst, err := logstore.BuildForMigration(ctx, "sqlite")
require.NoError(t, err)
t.Cleanup(func() { db.SetDB(nil) })
now := time.Now().UTC()
rows := []analyticsmodel.UserAccessLog{
{ID: 11, UserID: 1, Path: "/a", CreatedAt: now},
{ID: 22, UserID: 2, Path: "/b", CreatedAt: now.Add(time.Second)},
{ID: 33, UserID: 1, Path: "/c", CreatedAt: now.Add(2 * time.Second)},
}
require.NoError(t, src.UserAccessLogs.BatchInsert(ctx, rows))
require.NoError(t, copyUserAccessLogs(ctx, src, dst))
var got []analyticsmodel.UserAccessLog
require.NoError(t, dstDB.Order("id ASC").Find(&got).Error)
require.Len(t, got, 3)
for i, wantID := range []uint64{11, 22, 33} {
assert.Equal(t, wantID, got[i].ID, "row %d id preserved", i)
}
var srcCount int64
require.NoError(t, srcDB.Model(&analyticsmodel.UserAccessLog{}).Count(&srcCount).Error)
assert.Equal(t, int64(3), srcCount)
}
// TestClearTargetLogTablesClearsUserAccessLogs 验证清空目标包含用户访问日志表
// (6 张日志表之一),迁移「覆盖目标库已有日志」幂等前提成立。
func TestClearTargetLogTablesClearsUserAccessLogs(t *testing.T) {
t.Cleanup(runtimeconfig.Override(false, false))
logstore.ResetForTest()
defer logstore.ResetForTest()
ctx := context.Background()
dstDB := newLogDBSwitchDB(t)
db.SetDB(dstDB)
t.Cleanup(func() { db.SetDB(nil) })
dst, err := logstore.BuildForMigration(ctx, "sqlite")
require.NoError(t, err)
now := time.Now().UTC()
require.NoError(t, dst.UserAccessLogs.BatchInsert(ctx, []analyticsmodel.UserAccessLog{
{ID: 1, UserID: 1, Path: "/a", CreatedAt: now},
{ID: 2, UserID: 2, Path: "/b", CreatedAt: now},
}))
require.NoError(t, clearTargetLogTables(ctx, dst))
var count int64
require.NoError(t, dstDB.Model(&analyticsmodel.UserAccessLog{}).Count(&count).Error)
assert.Zero(t, count, "用户访问日志应被清空")
}
// TestClearTargetLogTablesDuringMigration 回归:冻结标记置位后,BuildForMigration 构造的
// 目标 store 必须放行用户访问日志清空/写入。skipFreeze 未传播到 UserAccessLogs store 时
// DeleteAll 会误报 ErrMigrating,导致真实切换任务在清空目标库阶段失败。
func TestClearTargetLogTablesDuringMigration(t *testing.T) {
logstore.ResetForTest()
defer logstore.ResetForTest()
gdb := newLogDBSwitchDB(t)
db.SetDB(gdb)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
logstore.SetConfigReader(func(ctx context.Context, key string) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, key)
if err != nil {
return "", err
}
return cfg.Value, nil
})
// 预置目标库已有日志(迁移「覆盖目标库已有日志」幂等前提)。
now := time.Now().UTC()
require.NoError(t, gdb.Create(&analyticsmodel.UserAccessLog{ID: 1, UserID: 1, Path: "/a", CreatedAt: now}).Error)
// 冻结标记置位(与真实任务 Execute 流程一致)。
require.NoError(t, setMigrationFlag(ctx, "migrating"))
t.Cleanup(func() { _ = setMigrationFlag(ctx, "") })
require.True(t, logstore.Migrating(ctx))
dst, err := logstore.BuildForMigration(ctx, "sqlite")
require.NoError(t, err)
require.NoError(t, clearTargetLogTables(ctx, dst), "迁移冻结期间目标库清空必须放行")
var count int64
require.NoError(t, gdb.Model(&analyticsmodel.UserAccessLog{}).Count(&count).Error)
assert.Zero(t, count, "用户访问日志应被清空")
}
// TestValidateSwitch 各非法组合报错。
func TestValidateSwitch(t *testing.T) {
t.Cleanup(func() {
})
gdb := newLogDBSwitchDB(t)
db.SetDB(gdb)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
setLogDB := func(v string) {
require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, v))
}
t.Run("same target rejected", func(t *testing.T) {
setLogDB("sqlite")
t.Cleanup(runtimeconfig.Override(false, false))
err := validateSwitch(ctx, "sqlite")
require.Error(t, err)
assert.Contains(t, err.Error(), "相同")
})
t.Run("clickhouse disabled rejected", func(t *testing.T) {
setLogDB("sqlite")
t.Cleanup(runtimeconfig.Override(false, false))
err := validateSwitch(ctx, "clickhouse")
require.Error(t, err)
assert.Contains(t, err.Error(), "ClickHouse 未启用")
})
t.Run("postgres requires main db enabled", func(t *testing.T) {
setLogDB("sqlite")
t.Cleanup(runtimeconfig.Override(false, false))
err := validateSwitch(ctx, "postgres")
require.Error(t, err)
assert.Contains(t, err.Error(), "PostgreSQL 未启用")
})
t.Run("sqlite rejected when main db is postgres", func(t *testing.T) {
setLogDB("postgres")
t.Cleanup(runtimeconfig.Override(true, false))
err := validateSwitch(ctx, "sqlite")
require.Error(t, err)
assert.Contains(t, err.Error(), "SQLite")
})
t.Run("valid postgres migration", func(t *testing.T) {
setLogDB("sqlite")
t.Cleanup(runtimeconfig.Override(true, false))
require.NoError(t, validateSwitch(ctx, "postgres"))
})
}
// TestLogDBSwitchValidatePayload 参数归一化与非法值拒绝。
func TestLogDBSwitchValidatePayload(t *testing.T) {
h := &LogDBSwitchHandler{}
cases := []struct {
name string
in string
want string
ok bool
}{
{name: "postgresql normalized", in: `{"target":"postgresql"}`, want: "postgres", ok: true},
{name: "sqlite3 normalized", in: `{"target":"sqlite3"}`, want: "sqlite", ok: true},
{name: "ch normalized", in: `{"target":"ch"}`, want: "clickhouse", ok: true},
{name: "postgres passthrough", in: `{"target":"postgres"}`, want: "postgres", ok: true},
{name: "invalid target", in: `{"target":"mysql"}`, ok: false},
{name: "malformed json", in: `not-json`, ok: false},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
out, err := h.ValidatePayload([]byte(c.in))
if !c.ok {
require.Error(t, err)
return
}
require.NoError(t, err)
var p logDBSwitchPayload
require.NoError(t, json.Unmarshal(out, &p))
assert.Equal(t, c.want, p.Target)
})
}
}
// TestExecuteFailureClearsMigrationFlag 迁移失败后 log_db_migration 冻结标记被清除。
// 在 FRESH DB(不预置 log_db_migration 行)上验证:setMigrationFlag 必须 upsert 建行,
// 且失败后经缓存路径(GetSystemConfigByKey)可观察为空。
func TestExecuteFailureClearsMigrationFlag(t *testing.T) {
t.Cleanup(runtimeconfig.Override(true, runtimeconfig.ClickHouseEnabled()))
logstore.ResetForTest()
defer logstore.ResetForTest()
gdb := newLogDBSwitchDB(t)
db.SetDB(gdb)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
// FRESH DB:log_db_migration 行不存在(不预置),log_database 预置为 sqlite。
require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, "sqlite"))
_, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDBMigration)
require.ErrorIs(t, err, gorm.ErrRecordNotFound)
// configReader 对 log_database 报错,使 logstore.Active 在冻结标记置位后失败。
logstore.SetConfigReader(func(_ context.Context, key string) (string, error) {
if key == model.ConfigKeyLogDatabase {
return "", errors.New("reader error")
}
return "", nil
})
_, err = (&LogDBSwitchHandler{}).Execute(ctx, []byte(`{"target":"postgres"}`))
require.Error(t, err)
assert.Contains(t, err.Error(), "reader error")
// 冻结标记必须被 upsert 持久化(行存在)并经缓存路径可观察为空,源库恢复可写。
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDBMigration)
require.NoError(t, err, "setMigrationFlag 应 upsert 创建 log_db_migration 行")
assert.Empty(t, cfg.Value, "失败后冻结标记必须清除,源库保持可写")
assert.False(t, logstore.Migrating(ctx))
}
// TestSetMigrationFlagObservableThroughCache 在 FRESH DB 上验证 setMigrationFlag 写入
// 经缓存路径(logstore.Migrating → repository 读取)实时反映:置位 true、清除 false。
func TestSetMigrationFlagObservableThroughCache(t *testing.T) {
logstore.ResetForTest()
defer logstore.ResetForTest()
gdb := newLogDBSwitchDB(t)
db.SetDB(gdb)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
// 按 bootstrap 同款注入 repository 读取,走 RAM 缓存路径。
logstore.SetConfigReader(func(ctx context.Context, key string) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, key)
if err != nil {
return "", err
}
return cfg.Value, nil
})
// FRESH DB:行缺失 → fail-open false。
assert.False(t, logstore.Migrating(ctx))
require.NoError(t, setMigrationFlag(ctx, "migrating"))
assert.True(t, logstore.Migrating(ctx), "置位后缓存路径必须立即观察到 migrating")
require.NoError(t, setMigrationFlag(ctx, ""))
assert.False(t, logstore.Migrating(ctx), "清除后缓存路径必须立即观察到非 migrating")
}
// TestFlipLogDatabaseRefreshesCachedConfig 验证翻转日志主库后缓存路径立即反映新库
// (logstore.ActiveDatabase / GetSystemConfigByKey),防止各进程继续写旧库(split-brain)。
func TestFlipLogDatabaseRefreshesCachedConfig(t *testing.T) {
logstore.ResetForTest()
defer logstore.ResetForTest()
gdb := newLogDBSwitchDB(t)
db.SetDB(gdb)
t.Cleanup(func() { db.SetDB(nil) })
ctx := context.Background()
logstore.SetConfigReader(func(ctx context.Context, key string) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, key)
if err != nil {
return "", err
}
return cfg.Value, nil
})
require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, "sqlite"))
active, err := logstore.ActiveDatabase(ctx)
require.NoError(t, err)
assert.Equal(t, "sqlite", active) // 预热缓存
require.NoError(t, flipLogDatabase(ctx, "postgres"))
active, err = logstore.ActiveDatabase(ctx)
require.NoError(t, err)
assert.Equal(t, "postgres", active, "翻转后缓存路径必须立即反映新库")
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase)
require.NoError(t, err)
assert.Equal(t, "postgres", cfg.Value)
}
@@ -0,0 +1,275 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"context"
"encoding/json"
"errors"
"sync"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"gorm.io/gorm"
)
const (
defaultObservabilityWindow = 24 * time.Hour
defaultObservabilityLimit = 120
maxObservabilityLimit = 500
defaultTrafficDistributionLimit = 8
nodeObservabilityCacheTTL = 15 * time.Second
)
var nodeObservabilityCache struct {
mu sync.Mutex
views map[string]cachedNodeObservability
}
type cachedNodeObservability struct {
view *NodeView
expiresAt time.Time
}
// NodeQuery filters node observability data.
type NodeQuery struct {
Hours int `json:"hours"`
Limit int `json:"limit"`
}
// NodeAnalytics groups node observability analytics.
type NodeAnalytics struct {
Traffic *TrafficWindowSummary `json:"traffic"`
Distributions TrafficDistributions `json:"distributions"`
Health HealthSummary `json:"health"`
}
// NodeTrends groups node observability trend series.
type NodeTrends struct {
Traffic24h []TrafficTrendPoint `json:"traffic_24h"`
Capacity24h []CapacityTrendPoint `json:"capacity_24h"`
Network24h []NetworkTrendPoint `json:"network_24h"`
DiskIO24h []DiskIOTrendPoint `json:"disk_io_24h"`
}
// RelayDashboardSnapshot summarizes tunnel relay status.
type RelayDashboardSnapshot struct {
TotalProxies int `json:"total_proxies"`
OnlineProxies int `json:"online_proxies"`
OfflineProxies int `json:"offline_proxies"`
Proxies []RelayProxyStat `json:"proxies"`
TotalConnections int `json:"total_connections"`
ClientCounts int `json:"client_counts"`
}
// RelayProxyStat is a single relay proxy entry.
type RelayProxyStat struct {
Name string `json:"name"`
Type string `json:"type"`
Status string `json:"status"`
ClientVersion string `json:"client_version"`
LastStartTime string `json:"last_start_time"`
LastCloseTime string `json:"last_close_time"`
ClientAddr string `json:"client_addr"`
}
// NodeView is the node observability API response.
type NodeView struct {
NodeID string `json:"node_id"`
Profile *model.OpenFlareNodeSystemProfile `json:"profile"`
MetricSnapshots []*NodeMetricSnapshotView `json:"metric_snapshots"`
HealthEvents []*model.OpenFlareHealthEvent `json:"health_events"`
Analytics NodeAnalytics `json:"analytics"`
Trends NodeTrends `json:"trends"`
RelayDashboard *RelayDashboardSnapshot `json:"relay_dashboard,omitempty"`
}
// HealthEventCleanupResult reports health event cleanup outcome.
type HealthEventCleanupResult struct {
NodeID string `json:"node_id"`
DeletedCount int64 `json:"deleted_count"`
}
// GetNodeObservability returns observability details for a node.
func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeView, error) {
now := time.Now()
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
if view, ok := getCachedNodeObservability(node.NodeID); ok {
return view, nil
}
limit := normalizeObservabilityLimit(query.Limit)
since := now.Add(-normalizeObservabilityWindow(query.Hours))
profile, err := repository.GetOpenFlareNodeSystemProfile(ctx, node.NodeID)
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
if errors.Is(err, gorm.ErrRecordNotFound) {
profile = nil
}
snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit)
if err != nil {
return nil, err
}
edgeHealth, err := repository.ListOpenFlareEdgeHealth(ctx, node.NodeID, since, limit)
if err != nil {
return nil, err
}
accessLogRegions, err := repository.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit)
if err != nil {
return nil, err
}
events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit)
if err != nil {
return nil, err
}
distributions := BuildTrafficDistributionsFromAccessLogs(
ctx, since, now, defaultTrafficDistributionLimit, accessLogRegions,
)
trafficSummary := buildTrafficWindowSummaryFromAccessLogs(ctx, node.NodeID, since, now)
view := &NodeView{
NodeID: node.NodeID,
Profile: profile,
MetricSnapshots: BuildMetricSnapshotViews(snapshots, edgeHealth),
HealthEvents: events,
Analytics: NodeAnalytics{
Traffic: trafficSummary,
Distributions: distributions,
Health: buildHealthSummary(latestMetricSnapshot(snapshots), trafficSummary, events),
},
Trends: BuildNodeTrends(ctx, now, node.NodeID, snapshots),
}
if node.NodeType == "tunnel_relay" {
frpsObs, frpsErr := repository.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1)
if frpsErr != nil {
return nil, frpsErr
}
var latestFrps *model.OpenFlareNodeObservationFrps
if len(frpsObs) > 0 {
latestFrps = frpsObs[0]
}
view.RelayDashboard = buildRelayDashboardSnapshot(node, latestFrps)
}
setCachedNodeObservability(node.NodeID, view)
return view, nil
}
func getCachedNodeObservability(nodeID string) (*NodeView, bool) {
nodeObservabilityCache.mu.Lock()
defer nodeObservabilityCache.mu.Unlock()
if nodeObservabilityCache.views == nil {
return nil, false
}
entry, ok := nodeObservabilityCache.views[nodeID]
if !ok || time.Now().After(entry.expiresAt) {
return nil, false
}
return entry.view, true
}
func setCachedNodeObservability(nodeID string, view *NodeView) {
nodeObservabilityCache.mu.Lock()
defer nodeObservabilityCache.mu.Unlock()
if nodeObservabilityCache.views == nil {
nodeObservabilityCache.views = make(map[string]cachedNodeObservability)
}
nodeObservabilityCache.views[nodeID] = cachedNodeObservability{
view: view,
expiresAt: time.Now().Add(nodeObservabilityCacheTTL),
}
}
// CleanupHealthEvents removes all health events for a node.
func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) {
node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil {
return nil, err
}
deletedCount, err := repository.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID)
if err != nil {
return nil, err
}
return &HealthEventCleanupResult{
NodeID: node.NodeID,
DeletedCount: deletedCount,
}, nil
}
func buildRelayDashboardSnapshot(node *model.OpenFlareNode, obs *model.OpenFlareNodeObservationFrps) *RelayDashboardSnapshot {
if node == nil {
return nil
}
totalProxies := 0
totalConnections := 0
clientCounts := 0
proxies := []RelayProxyStat{}
if obs != nil {
totalProxies = obs.FrpsProxyCount
totalConnections = obs.FrpsConnections
clientCounts = obs.FrpsClientCount
if obs.FrpsProxies != "" {
var decoded []RelayProxyStat
if err := json.Unmarshal([]byte(obs.FrpsProxies), &decoded); err == nil {
proxies = decoded
}
}
}
if totalProxies < 0 {
totalProxies = 0
}
onlineProxies := 0
for _, proxy := range proxies {
if proxy.Status == "online" {
onlineProxies++
}
}
if len(proxies) == 0 {
onlineProxies = totalProxies
if node.RelayStatus != "healthy" {
onlineProxies = 0
}
}
return &RelayDashboardSnapshot{
TotalProxies: totalProxies,
OnlineProxies: onlineProxies,
OfflineProxies: totalProxies - onlineProxies,
Proxies: proxies,
TotalConnections: maxInt(totalConnections, 0),
ClientCounts: maxInt(clientCounts, 0),
}
}
func maxInt(a int, b int) int {
if a > b {
return a
}
return b
}
func 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
}
@@ -0,0 +1,130 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestReadQueryStringArrayAcceptsHostsBracketForm(t *testing.T) {
gin.SetMode(gin.TestMode)
cases := []struct {
name string
url string
want []string
}{
{
name: "axios brackets form",
url: "/overview?hours=168&hosts%5B%5D=gist.arctel.de",
want: []string{"gist.arctel.de"},
},
{
name: "repeated hosts keys",
url: "/overview?hosts=a.example&hosts=b.example",
want: []string{"a.example", "b.example"},
},
{
name: "single hosts key",
url: "/overview?hosts=gist.arctel.de",
want: []string{"gist.arctel.de"},
},
{
name: "empty",
url: "/overview?hours=24",
want: nil,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
req, err := http.NewRequest(http.MethodGet, tc.url, nil)
require.NoError(t, err)
c.Request = req
got := readQueryStringArray(c, "hosts")
require.Equal(t, tc.want, got)
})
}
}
func TestReadAccessLogQueryIncludesStatusCode(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
req, err := http.NewRequest(
http.MethodGet,
"/?node_id=n1&remote_addr=1.2.3.4&host=a.example&path=/api&status_code=404&p=2&page_size=50",
nil,
)
require.NoError(t, err)
c.Request = req
got, err := readAccessLogQuery(c)
require.NoError(t, err)
require.Equal(t, "n1", got.NodeID)
require.Equal(t, "1.2.3.4", got.RemoteAddr)
require.Equal(t, "a.example", got.Host)
require.Equal(t, "/api", got.Path)
require.Equal(t, 404, got.StatusCode)
require.Equal(t, 2, got.Page)
require.Equal(t, 50, got.PageSize)
}
func TestReadAccessLogQueryRejectsInvalidStatusCode(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, raw := range []string{"abc", "99", "600", "-1"} {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
req, err := http.NewRequest(http.MethodGet, "/?status_code="+raw, nil)
require.NoError(t, err)
c.Request = req
_, err = readAccessLogQuery(c)
require.Error(t, err, "status_code=%s should be rejected", raw)
}
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
req, err := http.NewRequest(http.MethodGet, "/", nil)
require.NoError(t, err)
c.Request = req
got, err := readAccessLogQuery(c)
require.NoError(t, err)
require.Equal(t, 0, got.StatusCode)
}
func TestResolveAccessLogWindow(t *testing.T) {
since, until, err := resolveAccessLogWindow(
"2026-08-01T00:00:00Z",
"2026-08-02T00:00:00Z",
)
require.NoError(t, err)
require.True(t, until.After(since))
for _, tc := range []struct {
name string
since string
until string
}{
{name: "missing both", since: "", until: ""},
{name: "only since", since: "2026-08-01T00:00:00Z", until: ""},
{name: "only until", since: "", until: "2026-08-02T00:00:00Z"},
{name: "bad since", since: "not-a-time", until: "2026-08-02T00:00:00Z"},
{name: "bad until", since: "2026-08-01T00:00:00Z", until: "not-a-time"},
{name: "reversed", since: "2026-08-02T00:00:00Z", until: "2026-08-01T00:00:00Z"},
} {
t.Run(tc.name, func(t *testing.T) {
_, _, err := resolveAccessLogWindow(tc.since, tc.until)
require.Error(t, err)
})
}
}
@@ -0,0 +1,323 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"errors"
"net/http"
"strconv"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// GetAccessLogOverviewHandler 获取访问日志概览。
// @Summary 获取访问日志概览
// @Description 返回访问日志汇总指标、趋势与 Top 排行,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param host query string false "请求 Host(单域名)"
// @Param hosts query []string false "请求 Host 列表(多域名精确匹配)"
// @Param hours query int false "统计时间范围(小时)"
// @Param bucket_minutes query int false "趋势桶分钟数(1、3、5 或 60,默认 60)"
// @Success 200 {object} response.Any{data=observability.AccessLogOverview} "访问日志概览"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/overview [get]
func GetAccessLogOverviewHandler(c *gin.Context) {
result, err := GetAccessLogOverview(c.Request.Context(), AccessLogOverviewQuery{
NodeID: c.Query("node_id"),
Host: c.Query("host"),
Hosts: readQueryStringArray(c, "hosts"),
Hours: readQueryInt(c, "hours"),
BucketMinutes: readQueryInt(c, "bucket_minutes"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// GetAccessLogsHandler 分页列出访问日志。
// @Summary 列出访问日志
// @Description 分页返回 OpenFlare 访问日志,支持按节点、IP、主机、路径与状态码筛选,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param remote_addr query string false "客户端 IP"
// @Param host query string false "请求 Host"
// @Param path query string false "请求路径"
// @Param status_code query int false "HTTP 状态码(100-599)"
// @Param since query string false "起始时间(RFC3339,需与 until 成对提供)"
// @Param until query string false "结束时间(RFC3339,需与 since 成对提供)"
// @Param p query int false "页码"
// @Param page_size query int false "每页条数"
// @Param sort_by query string false "排序字段"
// @Param sort_order query string false "排序方向"
// @Success 200 {object} response.Any{data=observability.AccessLogList} "访问日志列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs [get]
func GetAccessLogsHandler(c *gin.Context) {
query, err := readAccessLogQuery(c)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
logs, err := ListAccessLogs(c.Request.Context(), query)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(logs))
}
// GetFoldedAccessLogsHandler 分页列出折叠访问日志。
// @Summary 列出折叠访问日志
// @Description 按时间桶聚合访问日志并分页返回,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param remote_addr query string false "客户端 IP"
// @Param host query string false "请求 Host"
// @Param path query string false "请求路径"
// @Param fold_minutes query int false "折叠时间窗口(分钟)"
// @Param p query int false "页码"
// @Param page_size query int false "每页条数"
// @Param sort_by query string false "排序字段"
// @Param sort_order query string false "排序方向"
// @Success 200 {object} response.Any{data=observability.FoldedAccessLogList} "折叠访问日志列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/folds [get]
func GetFoldedAccessLogsHandler(c *gin.Context) {
query, err := readAccessLogQuery(c)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
query.FoldMinutes = readQueryInt(c, "fold_minutes")
logs, err := ListFoldedAccessLogs(c.Request.Context(), query)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(logs))
}
// GetFoldedAccessLogIPsHandler 列出折叠桶内的 IP 汇总。
// @Summary 列出折叠访问日志 IP 汇总
// @Description 在指定时间桶内按 IP 聚合访问统计,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param remote_addr query string false "客户端 IP"
// @Param host query string false "请求 Host"
// @Param path query string false "请求路径"
// @Param bucket_started_at query string false "时间桶起始时间"
// @Param fold_minutes query int false "折叠时间窗口(分钟)"
// @Param p query int false "页码"
// @Param page_size query int false "每页条数"
// @Param sort_by query string false "排序字段"
// @Param sort_order query string false "排序方向"
// @Success 200 {object} response.Any{data=observability.FoldedAccessLogIPList} "折叠 IP 汇总列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/folds/ip-summary [get]
func GetFoldedAccessLogIPsHandler(c *gin.Context) {
result, err := ListFoldedAccessLogIPs(c.Request.Context(), FoldedAccessLogIPQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Path: c.Query("path"),
BucketStartedAt: c.Query("bucket_started_at"),
FoldMinutes: readQueryInt(c, "fold_minutes"),
Page: readQueryInt(c, "p"),
PageSize: readQueryInt(c, "page_size"),
SortBy: c.Query("sort_by"),
SortOrder: c.Query("sort_order"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// GetAccessLogIPSummariesHandler 列出访问日志 IP 汇总。
// @Summary 列出访问日志 IP 汇总
// @Description 按 IP 聚合访问日志统计并分页返回;支持 hours 或 since/until 时间窗,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param remote_addr query string false "客户端 IP"
// @Param host query string false "请求 Host"
// @Param hours query int false "统计时间范围(小时,1-720,默认 168)"
// @Param since query string false "开始时间 RFC3339(与 until 同时提供时优先于 hours)"
// @Param until query string false "结束时间 RFC3339"
// @Param p query int false "页码"
// @Param page_size query int false "每页条数"
// @Param sort_by query string false "排序字段 total_requests|request_length|bytes_sent|success_ratio|last_seen_at|remote_addr"
// @Param sort_order query string false "排序方向"
// @Success 200 {object} response.Any{data=observability.AccessLogIPSummaryList} "IP 汇总列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/ip-summary [get]
func GetAccessLogIPSummariesHandler(c *gin.Context) {
result, err := ListAccessLogIPSummaries(c.Request.Context(), AccessLogIPSummaryQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Hours: readQueryInt(c, "hours"),
Since: c.Query("since"),
Until: c.Query("until"),
Page: readQueryInt(c, "p"),
PageSize: readQueryInt(c, "page_size"),
SortBy: c.Query("sort_by"),
SortOrder: c.Query("sort_order"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// GetAccessLogIPTrendHandler 获取 IP 访问趋势。
// @Summary 获取访问日志 IP 趋势
// @Description 返回指定 IP 在时间范围内的访问趋势数据,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param remote_addr query string false "客户端 IP"
// @Param host query string false "请求 Host"
// @Param hours query int false "统计时间范围(小时)"
// @Param bucket_minutes query int false "时间桶粒度(分钟)"
// @Success 200 {object} response.Any{data=observability.AccessLogIPTrendView} "IP 访问趋势"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/ip-summary/trend [get]
func GetAccessLogIPTrendHandler(c *gin.Context) {
result, err := GetAccessLogIPTrend(c.Request.Context(), AccessLogIPTrendQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Hours: readQueryInt(c, "hours"),
BucketMinutes: readQueryInt(c, "bucket_minutes"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// GetAccessLogIPAnalysisHandler 获取单 IP 访问分析。
// @Summary 获取访问日志 IP 分析
// @Description 返回指定 IP 的汇总指标与 Top 分布,需要管理员权限
// @Tags openflare-observability
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID"
// @Param remote_addr query string false "客户端 IP"
// @Param host query string false "请求 Host"
// @Param hours query int false "统计时间范围(小时)"
// @Success 200 {object} response.Any{data=observability.AccessLogIPAnalysisView} "IP 访问分析"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/ip-summary/analysis [get]
func GetAccessLogIPAnalysisHandler(c *gin.Context) {
result, err := GetAccessLogIPAnalysis(c.Request.Context(), AccessLogIPAnalysisQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Hours: readQueryInt(c, "hours"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// CleanupAccessLogsHandler 清理过期访问日志。
// @Summary 清理访问日志
// @Description 按保留天数清理过期访问日志记录,需要管理员权限
// @Tags openflare-observability
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body observability.AccessLogCleanupInput true "清理参数"
// @Success 200 {object} response.Any{data=observability.AccessLogCleanupResult} "清理结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/access-logs/cleanup [post]
func CleanupAccessLogsHandler(c *gin.Context) {
var input AccessLogCleanupInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := CleanupAccessLogs(c.Request.Context(), input)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
func readAccessLogQuery(c *gin.Context) (AccessLogQuery, error) {
query := AccessLogQuery{
NodeID: c.Query("node_id"),
RemoteAddr: c.Query("remote_addr"),
Host: c.Query("host"),
Path: c.Query("path"),
Since: c.Query("since"),
Until: c.Query("until"),
Page: readQueryInt(c, "p"),
PageSize: readQueryInt(c, "page_size"),
SortBy: c.Query("sort_by"),
SortOrder: c.Query("sort_order"),
}
if raw := c.Query("status_code"); raw != "" {
code, err := strconv.Atoi(raw)
if err != nil || code < 100 || code > 599 {
return AccessLogQuery{}, errors.New(errInvalidStatusCode)
}
query.StatusCode = code
}
return query, nil
}
func readQueryInt(c *gin.Context, key string) int {
value, _ := strconv.Atoi(c.DefaultQuery(key, "0"))
return value
}
// readQueryStringArray reads repeated query values for key, and also accepts
// the Axios/jQuery bracket form key[] / key%5B%5D which Gin does not map to key.
func readQueryStringArray(c *gin.Context, key string) []string {
if values := c.QueryArray(key); len(values) > 0 {
return values
}
if values := c.QueryArray(key + "[]"); len(values) > 0 {
return values
}
return nil
}
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package option provides handler-level error message constants for the openflare option module.
package option
const (
errInvalidParams = "无效的参数"
errOptionInitFailed = "系统选项初始化失败"
errGeoIPProvider = "归属方式仅支持 disabled、mmdb、ip-api、geojs、ipinfo"
errGeoIPIPEmpty = "IP 不能为空"
errGeoIPIPInvalid = "IP 格式无效"
errGeoIPLookupDisabled = "GeoIP 查询已禁用"
)
@@ -0,0 +1,178 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"context"
"errors"
"fmt"
"strings"
"Wavelet/openflare/plugins/server/domain/option/uptimekuma"
"Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/buildinfo"
)
type publicAuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
AuthorizeURL string `json:"authorize_url"`
IconURL string `json:"icon_url"`
}
type statusView struct {
Version string `json:"version"`
StartTime int64 `json:"start_time"`
EmailVerification bool `json:"email_verification"`
ServerAddress string `json:"server_address"`
PasswordRegisterEnabled bool `json:"password_register_enabled"`
CapLoginEnabled bool `json:"cap_login_enabled"`
AuthSources []publicAuthSourceView `json:"auth_sources"`
}
type geoIPLookupRequest struct {
Provider string `json:"provider"`
IP string `json:"ip"`
}
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"`
}
type optionBatchPayload struct {
Options []model.OpenFlareOption `json:"options"`
}
func listOptions(ctx context.Context) ([]model.OpenFlareOption, error) {
// 从 SystemConfig 读取所有业务配置
configs, err := repository.ListAdminSystemConfigs(ctx, "business")
if err != nil {
return nil, err
}
options := make([]model.OpenFlareOption, 0, len(configs))
for _, config := range configs {
// 跳过敏感配置(如密码、令牌)
if config.Visibility == model.ConfigVisibilityHidden && isSecretConfigKey(config.Key) {
continue
}
// 将 snake_case key 转换为 PascalCase 以保持向后兼容
options = append(options, model.OpenFlareOption{
Key: config.Key,
Value: config.Value,
})
}
return options, nil
}
func updateOption(ctx context.Context, option model.OpenFlareOption) error {
return updateOptions(ctx, []model.OpenFlareOption{option})
}
func updateOptionsBatch(ctx context.Context, payload optionBatchPayload) error {
if len(payload.Options) == 0 {
return errors.New(errInvalidParams)
}
return updateOptions(ctx, payload.Options)
}
func updateOptions(ctx context.Context, options []model.OpenFlareOption) error {
if err := validateOptions(ctx, options); err != nil {
return err
}
// 将每个 option 更新到 SystemConfig
for _, opt := range options {
if err := repository.SaveOrUpdateSystemConfig(ctx, opt.Key, opt.Value); err != nil {
return fmt.Errorf("failed to update config %s: %w", opt.Key, err)
}
// 特殊处理:GeoIP 配置变更时刷新运行时
if opt.Key == model.ConfigKeyGeoIPProvider {
if err := geoip.RefreshRuntimeProvider(ctx); err != nil {
return err
}
}
}
return nil
}
func getStatus(ctx context.Context, baseAPIPath string) *statusView {
authSources, err := publicAuthSources(ctx, baseAPIPath)
if err != nil {
authSources = []publicAuthSourceView{}
}
// 从 SystemConfig 读取配置
emailVerification, _ := repository.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled)
serverAddress, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
passwordRegisterEnabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled)
capLoginEnabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyCapLoginEnabled)
return &statusView{
Version: buildinfo.Version,
StartTime: model.StartTime,
EmailVerification: emailVerification,
ServerAddress: serverAddress.Value,
PasswordRegisterEnabled: passwordRegisterEnabled,
CapLoginEnabled: capLoginEnabled,
AuthSources: authSources,
}
}
func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) {
sources, err := repository.GetActiveAuthSources(ctx)
if err != nil {
return nil, err
}
result := make([]publicAuthSourceView, 0, len(sources))
base := strings.TrimRight(baseAPIPath, "/")
for _, source := range sources {
result = append(result, publicAuthSourceView{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
AuthorizeURL: fmt.Sprintf("%s/oauth/%s/authorize", base, source.Name),
IconURL: source.IconURL,
})
}
return result, nil
}
func lookupGeoIP(_ context.Context, provider, rawIP string) (*geoIPLookupView, error) {
view, err := geoip.Lookup(provider, rawIP)
if err != nil {
return nil, err
}
return &geoIPLookupView{
Provider: view.Provider,
IP: view.IP,
ISOCode: view.ISOCode,
Name: view.Name,
Latitude: view.Latitude,
Longitude: view.Longitude,
}, nil
}
func syncUptimeKuma(ctx context.Context) error {
return uptimekuma.SyncToUptimeKuma(ctx)
}
// isSecretConfigKey 判断 SystemConfig 的 key 是否为敏感配置
func isSecretConfigKey(key string) bool {
return strings.Contains(key, "token") ||
strings.Contains(key, "secret") ||
strings.Contains(key, "password")
}
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"context"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupOptionTestDB(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.SystemConfig{}))
db.SetDB(sqliteDB)
// 预填充一些业务配置用于测试
seedConfigs := []model.SystemConfig{
{Key: "geoip_provider", Value: "ipinfo", Type: "business", Visibility: 0},
{Key: "uptime_kuma_password", Value: "secret-pwd", Type: "business", Visibility: 0},
}
for _, cfg := range seedConfigs {
require.NoError(t, sqliteDB.Create(&cfg).Error)
}
return func() {
db.SetDB(nil)
}
}
// setTestConfig 设置测试配置的辅助函数
func setTestConfig(t *testing.T, ctx context.Context, key, value string) {
t.Helper()
require.NoError(t, db.DB(ctx).Model(&model.SystemConfig{}).Where("key = ?", key).Update("value", value).Error)
}
func TestListOptionsFiltersSecretKeys(t *testing.T) {
cleanup := setupOptionTestDB(t)
defer cleanup()
ctx := context.Background()
options, err := listOptions(ctx)
require.NoError(t, err)
keys := make(map[string]string, len(options))
for _, option := range options {
keys[option.Key] = option.Value
}
// geoip_provider 应该出现在列表中
assert.Equal(t, "ipinfo", keys["geoip_provider"])
// 敏感配置(密码)应该被过滤掉
assert.NotContains(t, keys, "uptime_kuma_password")
}
func TestUpdateOptionPersistsToSystemConfig(t *testing.T) {
cleanup := setupOptionTestDB(t)
defer cleanup()
ctx := context.Background()
err := updateOption(ctx, model.OpenFlareOption{
Key: model.ConfigKeyGeoIPProvider,
Value: "mmdb",
})
require.NoError(t, err)
// 验证配置已写入 SystemConfig
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyGeoIPProvider)
require.NoError(t, err)
assert.Equal(t, "mmdb", config.Value)
}
func TestUpdateOpenRestyOptionPersistsToSystemConfig(t *testing.T) {
cleanup := setupOptionTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{
Key: model.ConfigKeyOpenRestyEventsUse,
Value: "epoll",
Type: "business",
Visibility: 0,
}).Error)
err := updateOption(ctx, model.OpenFlareOption{
Key: model.ConfigKeyOpenRestyEventsUse,
Value: "kqueue",
})
require.NoError(t, err)
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyOpenRestyEventsUse)
require.NoError(t, err)
assert.Equal(t, "kqueue", config.Value)
}
func TestLookupGeoIPDisabledProvider(t *testing.T) {
cleanup := setupOptionTestDB(t)
defer cleanup()
ctx := context.Background()
view, err := lookupGeoIP(ctx, "disabled", "8.8.8.8")
require.NoError(t, err)
assert.Equal(t, "disabled", view.Provider)
assert.Equal(t, "8.8.8.8", view.IP)
}
@@ -0,0 +1,284 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"encoding/json"
"fmt"
"regexp"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
openrestyrender "Wavelet/openflare/share/render/openresty"
)
const (
maxOriginErrorPageHTMLBytes = 256 << 10 // 256 KiB
maxSWOfflineDomains = 1000
)
var openRestyOptionValidators = map[string]func(key, value string) error{
model.ConfigKeyOpenRestyDefaultServerReturnStatus: validateOpenRestyDefaultServerReturnStatus,
model.ConfigKeyOpenRestyWorkerProcesses: validateOpenRestyWorkerProcesses,
model.ConfigKeyOpenRestyWorkerConnections: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyWorkerRlimitNofile: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyKeepaliveTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyKeepaliveRequests: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyClientHeaderTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyClientBodyTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestySendTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyProxyConnectTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyProxySendTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyProxyReadTimeout: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyGzipMinLength: validatePositiveIntegerOption,
model.ConfigKeyOpenRestyGzipCompLevel: validateOpenRestyGzipCompLevel,
model.ConfigKeyOpenRestyEventsUse: validateOpenRestyEventsUse,
model.ConfigKeyOpenRestyResolvers: validateOpenRestyResolvers,
model.ConfigKeyOpenRestyEventsMultiAcceptEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyWebsocketEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyHTTP3Enabled: validateBooleanOption,
model.ConfigKeyOpenRestyProxyRequestBufferingEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyProxyBufferingEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyGzipEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyCacheEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyCacheLockEnabled: validateBooleanOption,
model.ConfigKeyOpenRestyProxyBuffers: validateOpenRestyProxyBuffers,
model.ConfigKeyOpenRestyLargeClientHeaderBuffers: validateOpenRestyProxyBuffers,
model.ConfigKeyOpenRestyProxyBufferSize: validateOpenRestySizeValue,
model.ConfigKeyOpenRestyProxyBusyBuffersSize: validateOpenRestySizeValue,
model.ConfigKeyOpenRestyCacheMaxSize: validateOpenRestySizeValue,
model.ConfigKeyOpenRestyClientMaxBodySize: validateOpenRestySizeValue,
model.ConfigKeyOpenRestyCachePath: validateOpenRestyCachePath,
model.ConfigKeyOpenRestyCacheLevels: validateOpenRestyCacheLevels,
model.ConfigKeyOpenRestyCacheInactive: validateOpenRestyDurationToken,
model.ConfigKeyOpenRestyCacheLockTimeout: validateOpenRestyDurationToken,
model.ConfigKeyOpenRestyCacheKeyTemplate: validateOpenRestyCacheKeyTemplate,
model.ConfigKeyOpenRestyCacheUseStale: validateOpenRestyCacheUseStale,
model.ConfigKeyOpenRestyMainConfigTemplate: validateOpenRestyMainConfigTemplate,
model.ConfigKeyOpenRestyDefaultLimitConnPerServer: validateNonNegativeIntegerOption,
model.ConfigKeyOpenRestyDefaultLimitConnPerIP: validateNonNegativeIntegerOption,
model.ConfigKeyOpenRestyDefaultLimitRate: validateOpenRestyDefaultLimitRate,
model.ConfigKeyOpenRestyDefaultLimitReqPerIP: validateOpenRestyDefaultLimitReqPerIP,
model.ConfigKeyOriginErrorPageEnabled: validateBooleanOption,
model.ConfigKeyOriginErrorPageStatusCodes: validateOriginErrorPageStatusCodes,
model.ConfigKeyOriginErrorPageHTML: validateOriginErrorPageHTML,
model.ConfigKeyOriginErrorPageGetOnly: validateBooleanOption,
model.ConfigKeySWOfflineEnabled: validateBooleanOption,
model.ConfigKeySWOfflineHTML: validateSWOfflineHTML,
model.ConfigKeySWOfflineDomains: validateSWOfflineDomains,
}
var openRestyDefaultLimitRatePattern = regexp.MustCompile(`^\d+[kKmM]?$`)
func validateOpenRestyOption(key, value string) error {
// HTML 按原始字节长度校验,避免 TrimSpace 影响上限判断
if key == model.ConfigKeyOriginErrorPageHTML || key == model.ConfigKeySWOfflineHTML {
return validateOriginErrorPageHTML(key, value)
}
trimmed := strings.TrimSpace(value)
if validator, ok := openRestyOptionValidators[key]; ok {
return validator(key, trimmed)
}
return nil
}
func validateOpenRestyDefaultServerReturnStatus(key, trimmed string) error {
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
statusCode, _ := strconv.Atoi(trimmed)
if statusCode < 100 || statusCode > 999 {
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
}
return nil
}
func validateOpenRestyWorkerProcesses(key, trimmed string) error {
if trimmed == "auto" {
return nil
}
return validatePositiveIntegerOption(key, trimmed)
}
func validateOpenRestyGzipCompLevel(key, trimmed string) error {
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
level, _ := strconv.Atoi(trimmed)
if level > maxOpenRestyGzipCompLevel {
return fmt.Errorf("%s 不能大于 %d", key, maxOpenRestyGzipCompLevel)
}
return nil
}
func validateOpenRestyEventsUse(key, trimmed string) error {
if trimmed == "" {
return nil
}
switch trimmed {
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
return nil
default:
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
}
}
func validateOpenRestyResolvers(key, trimmed string) error {
if trimmed == "" {
return nil
}
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
}
return nil
}
func validateOpenRestyProxyBuffers(key, trimmed string) error {
if openRestyProxyBuffersPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
}
func validateOpenRestySizeValue(key, trimmed string) error {
if openRestySizePattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
}
func validateOpenRestyCachePath(key, trimmed string) error {
if strings.ContainsAny(trimmed, "\r\n\t") {
return fmt.Errorf("%s 不能包含换行或制表符", key)
}
return nil
}
func validateOpenRestyCacheLevels(key, trimmed string) error {
if openRestyCacheLevelsPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
}
func validateOpenRestyDurationToken(key, trimmed string) error {
if openRestyDurationTokenPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
}
func validateOpenRestyCacheKeyTemplate(key, trimmed string) error {
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
if strings.ContainsAny(trimmed, "\r\n") {
return fmt.Errorf("%s 不能包含换行", key)
}
return nil
}
func validateOpenRestyCacheUseStale(key, trimmed string) error {
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
allowedTokens := map[string]struct{}{
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
}
for token := range strings.FieldsSeq(trimmed) {
if _, ok := allowedTokens[token]; !ok {
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
}
}
return nil
}
func validateOpenRestyMainConfigTemplate(key, value string) error {
if strings.TrimSpace(value) == "" {
return fmt.Errorf("%s 不能为空", key)
}
return nil
}
func validateOpenRestyDefaultLimitRate(key, trimmed string) error {
if trimmed == "" || trimmed == "0" {
return nil
}
if !openRestyDefaultLimitRatePattern.MatchString(strings.ToLower(trimmed)) {
return fmt.Errorf("%s 格式不合法,请使用 512k、1m 或纯数字,空表示关闭", key)
}
return nil
}
var openRestyDefaultLimitReqPerIPPattern = regexp.MustCompile(`^\d+r/[sm]$`)
func validateOpenRestyDefaultLimitReqPerIP(key, trimmed string) error {
if trimmed == "" || trimmed == "0" {
return nil
}
if !openRestyDefaultLimitReqPerIPPattern.MatchString(strings.ToLower(trimmed)) {
return fmt.Errorf("%s 格式不合法,请输入类似 10r/s、100r/m,或留空关闭", key)
}
return nil
}
func validateOriginErrorPageStatusCodes(key, trimmed string) error {
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
var tags []string
if err := json.Unmarshal([]byte(trimmed), &tags); err != nil {
return fmt.Errorf("%s 必须为 JSON 字符串数组", key)
}
if len(tags) == 0 {
return fmt.Errorf("%s 至少包含一个状态码标签", key)
}
codes, err := openrestyrender.ExpandStatusCodeTags(tags)
if err != nil {
return fmt.Errorf("%s: %w", key, err)
}
if len(codes) == 0 {
return fmt.Errorf("%s 展开后不能为空", key)
}
return nil
}
func validateOriginErrorPageHTML(key, value string) error {
if len(value) > maxOriginErrorPageHTMLBytes {
return fmt.Errorf("%s 长度不能超过 %d 字节(256 KiB)", key, maxOriginErrorPageHTMLBytes)
}
return nil
}
func validateSWOfflineHTML(key, value string) error {
return validateOriginErrorPageHTML(key, value)
}
func validateSWOfflineDomains(key, value string) error {
var domains []string
if err := json.Unmarshal([]byte(value), &domains); err != nil || domains == nil {
return fmt.Errorf("%s 必须为 JSON 字符串数组", key)
}
if len(domains) > maxSWOfflineDomains {
return fmt.Errorf("%s 最多支持 %d 个域名", key, maxSWOfflineDomains)
}
seen := make(map[string]struct{}, len(domains))
for _, raw := range domains {
domain := strings.ToLower(strings.TrimSpace(raw))
if domain == "" {
return fmt.Errorf("%s 包含空域名", key)
}
if raw != domain {
return fmt.Errorf("%s 域名必须为小写且不含首尾空格:%s", key, raw)
}
if _, ok := seen[domain]; ok {
return fmt.Errorf("%s 包含重复域名 %s", key, domain)
}
seen[domain] = struct{}{}
}
return nil
}
@@ -0,0 +1,127 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"fmt"
"strings"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestValidateOriginErrorPageStatusCodes(t *testing.T) {
t.Parallel()
tests := []struct {
name string
value string
wantErr string
}{
{
name: "合法单码与区间",
value: `["522","500-502"]`,
},
{
name: "默认区间",
value: `["500-599"]`,
},
{
name: "非法标签",
value: `["abc"]`,
wantErr: "无效状态码",
},
{
name: "非 JSON 数组",
value: `500-599`,
wantErr: "必须为 JSON 字符串数组",
},
{
name: "空数组",
value: `[]`,
wantErr: "至少包含一个状态码标签",
},
{
name: "越界状态码",
value: `["399"]`,
wantErr: "状态码须在",
},
{
name: "空字符串",
value: "",
wantErr: "不能为空",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
err := validateOpenRestyOption(model.ConfigKeyOriginErrorPageStatusCodes, tt.value)
if tt.wantErr == "" {
require.NoError(t, err)
return
}
require.Error(t, err)
assert.Contains(t, err.Error(), tt.wantErr)
})
}
}
func TestValidateOriginErrorPageHTML(t *testing.T) {
t.Parallel()
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, ""))
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, "<html>ok</html>"))
oversized := strings.Repeat("a", maxOriginErrorPageHTMLBytes+1)
err := validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, oversized)
require.Error(t, err)
assert.Contains(t, err.Error(), "长度不能超过")
// 恰好上限应通过
atLimit := strings.Repeat("b", maxOriginErrorPageHTMLBytes)
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, atLimit))
}
func TestValidateOriginErrorPageEnabled(t *testing.T) {
t.Parallel()
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageEnabled, "true"))
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageEnabled, "false"))
err := validateOpenRestyOption(model.ConfigKeyOriginErrorPageEnabled, "yes")
require.Error(t, err)
assert.Contains(t, err.Error(), "true 或 false")
}
func TestValidateSWOfflineDomains(t *testing.T) {
cases := []struct {
name string
value string
ok bool
}{
{"empty array", `[]`, true},
{"single", `["example.com"]`, true},
{"multiple", `["example.com","api.example.com"]`, true},
{"invalid json", `not-json`, false},
{"null", "null", false},
{"empty element", `[""]`, false},
{"duplicate", `["example.com","example.com"]`, false},
{"whitespace dedup", `[" Example.com ","example.com"]`, false},
{"over limit", fmt.Sprintf(`[%s]`, strings.Repeat(`"a.com",`, maxSWOfflineDomains)+`"a.com"`), false},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
err := validateSWOfflineDomains("sw_offline_domains", tc.value)
if tc.ok && err != nil {
t.Fatalf("want ok, got %v", err)
}
if !tc.ok && err == nil {
t.Fatal("want error, got nil")
}
})
}
}
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// GetStatusHandler 获取公开运行状态。
// @Summary 获取 OpenFlare 公开状态
// @Description 返回版本、认证源与系统公开配置,无需登录
// @Tags openflare-option
// @Produce json
// @Success 200 {object} response.Any{data=option.statusView} "公开状态"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/status [get]
func GetStatusHandler(c *gin.Context) {
view := getStatus(c.Request.Context(), "/api/v1/d")
c.JSON(http.StatusOK, response.OK(view))
}
// ListOptionsHandler 列出全部配置项。
// @Summary 列出 OpenFlare 配置项
// @Description 返回全部非敏感 OpenFlare 配置项,需要管理员权限
// @Tags openflare-option
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.OpenFlareOption} "配置项列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option [get]
func ListOptionsHandler(c *gin.Context) {
options, err := listOptions(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(options))
}
// UpdateOptionHandler 更新单个配置项。
// @Summary 更新 OpenFlare 配置项
// @Description 更新单个 OpenFlare 配置项,需要管理员权限
// @Tags openflare-option
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body model.OpenFlareOption true "配置项"
// @Success 200 {object} response.Any "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/update [post]
func UpdateOptionHandler(c *gin.Context) {
var option model.OpenFlareOption
if !apiutil.BindJSON(c, &option) {
return
}
if apiutil.AbortBadRequestOnError(c, updateOption(c.Request.Context(), option)) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// UpdateOptionsBatchHandler 批量更新配置项。
// @Summary 批量更新 OpenFlare 配置项
// @Description 批量更新多个 OpenFlare 配置项,需要管理员权限
// @Tags openflare-option
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body option.optionBatchPayload true "批量配置项"
// @Success 200 {object} response.Any "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/update-batch [post]
func UpdateOptionsBatchHandler(c *gin.Context) {
var payload optionBatchPayload
if !apiutil.BindJSON(c, &payload) {
return
}
if apiutil.AbortBadRequestOnError(c, updateOptionsBatch(c.Request.Context(), payload)) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// LookupGeoIPHandler 查询 GeoIP 信息。
// @Summary GeoIP 地址查询
// @Description 按提供商与 IP 查询地理位置信息,需要管理员权限
// @Tags openflare-option
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body option.geoIPLookupRequest true "查询参数"
// @Success 200 {object} response.Any{data=option.geoIPLookupView} "GeoIP 查询结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/geoip/lookup [post]
func LookupGeoIPHandler(c *gin.Context) {
var request geoIPLookupRequest
if !apiutil.BindJSON(c, &request) {
return
}
view, err := lookupGeoIP(c.Request.Context(), request.Provider, request.IP)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// SyncUptimeKumaHandler 同步 Uptime Kuma 监控。
// @Summary 同步 Uptime Kuma
// @Description 将 OpenFlare 节点同步到 Uptime Kuma,需要管理员权限
// @Tags openflare-option
// @Accept json
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=string} "同步成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/uptimekuma/sync [post]
func SyncUptimeKumaHandler(c *gin.Context) {
if apiutil.AbortBadRequestOnError(c, syncUptimeKuma(c.Request.Context())) {
return
}
c.JSON(http.StatusOK, response.OK("同步成功"))
}
@@ -0,0 +1,386 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package uptimekuma provides a Socket.IO client and sync implementation for Uptime Kuma.
package uptimekuma
import (
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"maps"
"net/http"
"strconv"
"strings"
"sync"
"time"
)
const emitAckTimeout = 10 * time.Second
// Monitor represents a monitor entry from Uptime Kuma.
type Monitor struct {
ID int `json:"id"`
Name string `json:"name"`
URL string `json:"url"`
Type string `json:"type"`
Interval int `json:"interval"`
MaxRetries int `json:"maxretries"`
RetryInterval int `json:"retryInterval"`
Timeout int `json:"timeout"`
Tags []Tag `json:"tags"`
}
// Tag represents a tag attached to a monitor.
type Tag struct {
ID int `json:"tag_id"`
Name string `json:"name"`
Color string `json:"color"`
}
// TagItem represents a tag returned by getTags.
type TagItem struct {
ID int `json:"id"`
Name string `json:"name"`
Color string `json:"color"`
}
// SocketIOClient is a minimal Engine.IO/Socket.IO polling client for Uptime Kuma.
type SocketIOClient struct {
baseURL string
httpClient *http.Client
sid string
ackMutex sync.Mutex
ackID int
ackChanMap map[int]chan string
doneChan chan struct{}
closeOnce sync.Once
monitorListMutex sync.RWMutex
monitorList map[string]Monitor
monitorListChan chan struct{}
monitorListOnce sync.Once
ctx context.Context
cancel context.CancelFunc
err error
}
// NewSocketIOClient creates a Socket.IO polling client for the given base URL.
func NewSocketIOClient(baseURL string) *SocketIOClient {
ctx, cancel := context.WithCancel(context.Background())
return &SocketIOClient{
baseURL: strings.TrimSuffix(baseURL, "/"),
httpClient: &http.Client{
Timeout: 60 * time.Second,
},
ackChanMap: make(map[int]chan string),
doneChan: make(chan struct{}),
monitorListChan: make(chan struct{}),
monitorList: make(map[string]Monitor),
ctx: ctx,
cancel: cancel,
}
}
// Connect performs the Engine.IO handshake and starts the polling loop.
func (c *SocketIOClient) Connect() error {
slog.Debug("Uptime Kuma client starting handshake", "baseURL", c.baseURL)
u := c.baseURL + "/socket.io/?EIO=4&transport=polling"
reqHandshake, err := http.NewRequestWithContext(c.ctx, http.MethodGet, u, nil)
if err != nil {
return fmt.Errorf("create handshake request failed: %w", err)
}
resp, err := c.httpClient.Do(reqHandshake)
if err != nil {
slog.Error("Uptime Kuma handshake connection failed", "url", u, "error", err)
return fmt.Errorf("handshake request failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
bs, err := io.ReadAll(resp.Body)
if err != nil {
slog.Error("Failed to read Uptime Kuma handshake response body", "error", err)
return fmt.Errorf("read handshake body failed: %w", err)
}
bodyStr := string(bs)
slog.Debug("Received handshake response from Uptime Kuma", "body", bodyStr)
if len(bodyStr) == 0 || bodyStr[0] != '0' {
return fmt.Errorf("invalid handshake response format: %s", bodyStr)
}
var hs struct {
Sid string `json:"sid"`
}
if err := json.Unmarshal([]byte(bodyStr[1:]), &hs); err != nil {
return fmt.Errorf("unmarshal handshake sid failed: %w", err)
}
c.sid = hs.Sid
slog.Debug("Uptime Kuma handshake success", "sid", c.sid)
slog.Debug("Sending namespace connect request to Uptime Kuma", "sid", c.sid)
connectURL := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
req, err := http.NewRequestWithContext(c.ctx, http.MethodPost, connectURL, strings.NewReader("40"))
if err != nil {
return fmt.Errorf("create connect request failed: %w", err)
}
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
respConnect, err := c.httpClient.Do(req)
if err != nil {
slog.Error("Uptime Kuma namespace connect request failed", "sid", c.sid, "error", err)
return fmt.Errorf("namespace connect failed: %w", err)
}
_ = respConnect.Body.Close()
slog.Debug("Namespace connected successfully to Uptime Kuma", "sid", c.sid)
go c.pollLoop()
return nil
}
func (c *SocketIOClient) pollLoop() {
slog.Debug("Uptime Kuma polling loop started", "sid", c.sid)
defer c.Close()
for {
select {
case <-c.doneChan:
slog.Debug("Uptime Kuma polling loop stopped (doneChan closed)", "sid", c.sid)
return
default:
}
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
reqPoll, err := http.NewRequestWithContext(c.ctx, http.MethodGet, u, nil)
if err != nil {
slog.Error("Failed to create Uptime Kuma polling request", "sid", c.sid, "error", err)
c.err = err
return
}
resp, err := c.httpClient.Do(reqPoll)
if err != nil {
slog.Error("Uptime Kuma polling request failed", "sid", c.sid, "error", err)
c.err = err
return
}
bs, err := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if err != nil {
slog.Error("Failed to read Uptime Kuma polling body", "sid", c.sid, "error", err)
c.err = err
return
}
bodyStr := string(bs)
if len(bodyStr) == 0 {
continue
}
slog.Debug("Received polling payload from Uptime Kuma", "length", len(bodyStr))
packets := strings.SplitSeq(bodyStr, "\x1e")
for pkt := range packets {
if len(pkt) == 0 {
continue
}
engineIOType := pkt[0]
payload := pkt[1:]
slog.Debug("Parsing engine.io packet", "type", string(engineIOType), "payload_len", len(payload))
switch engineIOType {
case '2':
slog.Debug("Received engine.io ping, responding with pong", "sid", c.sid)
c.sendPong()
case '4':
if len(payload) == 0 {
continue
}
socketIOType := payload[0]
socketIOPayload := payload[1:]
slog.Debug("Parsing socket.io packet", "type", string(socketIOType), "payload", socketIOPayload)
switch socketIOType {
case '2':
c.handleEvent(socketIOPayload)
case '3':
c.handleAck(socketIOPayload)
}
}
}
}
}
func (c *SocketIOClient) sendPong() {
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
req, err := http.NewRequestWithContext(c.ctx, http.MethodPost, u, strings.NewReader("3"))
if err != nil {
return
}
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
resp, err := c.httpClient.Do(req)
if err == nil {
_ = resp.Body.Close()
}
}
func (c *SocketIOClient) handleEvent(payload string) {
var arr []json.RawMessage
if err := json.Unmarshal([]byte(payload), &arr); err != nil || len(arr) < 2 {
return
}
var eventName string
if err := json.Unmarshal(arr[0], &eventName); err != nil {
return
}
if eventName == "monitorList" {
var list map[string]Monitor
if err := json.Unmarshal(arr[1], &list); err == nil {
c.monitorListMutex.Lock()
c.monitorList = list
c.monitorListMutex.Unlock()
c.monitorListOnce.Do(func() {
close(c.monitorListChan)
})
}
}
}
func (c *SocketIOClient) handleAck(payload string) {
idx := strings.IndexByte(payload, '[')
if idx == -1 {
return
}
ackIDStr := payload[:idx]
ackID, err := strconv.Atoi(ackIDStr)
if err != nil {
return
}
c.ackMutex.Lock()
ch, ok := c.ackChanMap[ackID]
if ok {
delete(c.ackChanMap, ackID)
c.ackMutex.Unlock()
select {
case ch <- payload[idx:]:
default:
}
} else {
c.ackMutex.Unlock()
}
}
// Emit sends a Socket.IO event and waits for the corresponding ack.
func (c *SocketIOClient) Emit(event string, args ...any) (string, error) {
c.ackMutex.Lock()
id := c.ackID
c.ackID++
ch := make(chan string, 1)
c.ackChanMap[id] = ch
c.ackMutex.Unlock()
payloadArr := make([]any, 1, 1+len(args))
payloadArr[0] = event
payloadArr = append(payloadArr, args...)
bs, err := json.Marshal(payloadArr)
if err != nil {
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Failed to marshal event payload", "event", event, "error", err)
return "", err
}
body := fmt.Sprintf("42%d%s", id, string(bs))
slog.Debug("Emitting Socket.IO event", "event", event, "ackID", id, "payload_len", len(bs))
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling&sid=%s", c.baseURL, c.sid)
req, err := http.NewRequestWithContext(c.ctx, http.MethodPost, u, strings.NewReader(body))
if err != nil {
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
return "", err
}
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
resp, err := c.httpClient.Do(req)
if err != nil {
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Failed to send Emit request", "event", event, "ackID", id, "error", err)
return "", err
}
_ = resp.Body.Close()
select {
case result := <-ch:
slog.Debug("Received Ack for event", "event", event, "ackID", id, "response", result)
return result, nil
case <-time.After(emitAckTimeout):
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Timeout waiting for event Ack", "event", event, "ackID", id)
return "", fmt.Errorf("timeout waiting for ack for event: %s", event)
case <-c.doneChan:
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
slog.Error("Client closed while waiting for event Ack", "event", event, "ackID", id)
return "", fmt.Errorf("client closed while waiting for event ack: %s", event)
}
}
// Close shuts down the polling loop.
func (c *SocketIOClient) Close() {
c.closeOnce.Do(func() {
c.cancel()
close(c.doneChan)
})
}
// GetMonitorListChan returns a channel closed when the first monitorList event arrives.
func (c *SocketIOClient) GetMonitorListChan() <-chan struct{} {
return c.monitorListChan
}
// GetMonitorList returns a copy of the current monitor list.
func (c *SocketIOClient) GetMonitorList() map[string]Monitor {
c.monitorListMutex.RLock()
defer c.monitorListMutex.RUnlock()
m := make(map[string]Monitor, len(c.monitorList))
maps.Copy(m, c.monitorList)
return m
}
// ParseAckResponse unmarshals an ack payload and validates the ok status when present.
func ParseAckResponse(response string, target any) error {
var arr []json.RawMessage
if err := json.Unmarshal([]byte(response), &arr); err != nil || len(arr) == 0 {
return fmt.Errorf("invalid ack response format: %s", response)
}
var status struct {
Ok bool `json:"ok"`
Msg string `json:"msg"`
}
if err := json.Unmarshal(arr[0], &status); err == nil {
if !status.Ok {
errMsg := status.Msg
if errMsg == "" {
errMsg = "unknown error from Uptime Kuma"
}
return fmt.Errorf("uptime Kuma error response: %s", errMsg)
}
}
if target != nil {
return json.Unmarshal(arr[0], target)
}
return nil
}
@@ -0,0 +1,313 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package uptimekuma
import (
"context"
"errors"
"fmt"
"log/slog"
"strings"
"sync/atomic"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
)
const uptimeKumaTagOpenFlare = "OpenFlare"
var isSyncing atomic.Bool
// kumaConfig 封装 UptimeKuma 配置
type kumaConfig struct {
URL string
Username string
Password string
MonitorScope string
SelectedSites string
Interval int
Retry int
RetryInterval int
Timeout int
}
// loadKumaConfig 从 SystemConfig 加载 UptimeKuma 配置
func loadKumaConfig(ctx context.Context) *kumaConfig {
url, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaURL)
username, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaUsername)
password, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaPassword)
scope, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaMonitorScope)
selected, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaSelectedSites)
interval, _ := repository.GetIntByKey(ctx, model.ConfigKeyUptimeKumaInterval)
if interval <= 0 {
interval = 60
}
retry, _ := repository.GetIntByKey(ctx, model.ConfigKeyUptimeKumaRetry)
retryInterval, _ := repository.GetIntByKey(ctx, model.ConfigKeyUptimeKumaRetryInterval)
if retryInterval <= 0 {
retryInterval = 60
}
timeout, _ := repository.GetIntByKey(ctx, model.ConfigKeyUptimeKumaTimeout)
if timeout <= 0 {
timeout = 48
}
if scope.Value == "" {
scope.Value = "all"
}
return &kumaConfig{
URL: strings.TrimSpace(url.Value),
Username: strings.TrimSpace(username.Value),
Password: strings.TrimSpace(password.Value),
MonitorScope: scope.Value,
SelectedSites: selected.Value,
Interval: interval,
Retry: retry,
RetryInterval: retryInterval,
Timeout: timeout,
}
}
// SyncToUptimeKuma synchronizes enabled proxy routes to Uptime Kuma monitors.
func SyncToUptimeKuma(ctx context.Context) error {
// 检查是否启用
enabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyUptimeKumaEnabled)
if !enabled {
return errors.New("uptime Kuma integration is disabled")
}
if !isSyncing.CompareAndSwap(false, true) {
return errors.New("sync task is already in progress, please try again later")
}
defer isSyncing.Store(false)
// 加载配置
config := loadKumaConfig(ctx)
// 验证配置
if err := validateKumaConfig(config); err != nil {
return err
}
slog.Info("Starting Uptime Kuma sync process",
"url", config.URL,
"username", config.Username,
"scope", config.MonitorScope,
)
allRoutes, err := repository.ListProxyRoutes(ctx)
if err != nil {
return fmt.Errorf("failed to list local proxy routes: %w", err)
}
expectedRoutes := filterExpectedRoutes(allRoutes, config)
client, err := connectAndLoginUptimeKuma(config.URL, config.Username, config.Password)
if err != nil {
return err
}
defer client.Close()
openFlareTagID, err := ensureOpenFlareTag(client)
if err != nil {
return err
}
existingOpenFlareMonitors := filterOpenFlareMonitors(client.GetMonitorList(), openFlareTagID)
expectedSitesMap := syncRouteMonitors(ctx, client, expectedRoutes, existingOpenFlareMonitors, openFlareTagID, config)
removeStaleMonitors(client, existingOpenFlareMonitors, expectedSitesMap)
return nil
}
func filterExpectedRoutes(allRoutes []*model.ProxyRoute, config *kumaConfig) []*model.ProxyRoute {
scope := config.MonitorScope
if scope == "selected" {
selectedList := strings.Split(config.SelectedSites, ",")
selectedMap := make(map[string]bool)
for _, name := range selectedList {
trimmedName := strings.TrimSpace(name)
if trimmedName != "" {
selectedMap[trimmedName] = true
}
}
var expectedRoutes []*model.ProxyRoute
for _, route := range allRoutes {
if route.Enabled && selectedMap[route.SiteName] {
expectedRoutes = append(expectedRoutes, route)
}
}
return expectedRoutes
}
var expectedRoutes []*model.ProxyRoute
for _, route := range allRoutes {
if route.Enabled {
expectedRoutes = append(expectedRoutes, route)
}
}
return expectedRoutes
}
func ensureOpenFlareTag(client *SocketIOClient) (int, error) {
slog.Debug("Fetching tags from Uptime Kuma")
tagsAck, err := client.Emit("getTags")
if err != nil {
slog.Error("Failed to request tags from Uptime Kuma", "error", err)
return 0, fmt.Errorf("failed to fetch tags: %w", err)
}
var tagsResult struct {
Ok bool `json:"ok"`
Tags []TagItem `json:"tags"`
}
if err := ParseAckResponse(tagsAck, &tagsResult); err != nil {
slog.Error("Failed to parse tags response from Uptime Kuma", "error", err)
return 0, fmt.Errorf("parse tags response failed: %w", err)
}
for _, tag := range tagsResult.Tags {
if tag.Name == uptimeKumaTagOpenFlare {
slog.Debug("Found existing OpenFlare tag", "tag_id", tag.ID)
return tag.ID, nil
}
}
slog.Debug("OpenFlare tag not found, creating new tag")
addTagAck, err := client.Emit("addTag", map[string]string{
"name": uptimeKumaTagOpenFlare,
"color": "#4f46e5",
})
if err != nil {
slog.Error("Failed to create OpenFlare tag in Uptime Kuma", "error", err)
return 0, fmt.Errorf("failed to create tag: %w", err)
}
var tagResult struct {
Ok bool `json:"ok"`
Tag struct {
ID int `json:"id"`
} `json:"tag"`
}
if err := ParseAckResponse(addTagAck, &tagResult); err != nil || tagResult.Tag.ID == 0 {
slog.Error("Failed to parse addTag response from Uptime Kuma", "error", err)
return 0, fmt.Errorf("parse addTag response failed: %w", err)
}
slog.Debug("Successfully created OpenFlare tag", "tag_id", tagResult.Tag.ID)
return tagResult.Tag.ID, nil
}
func filterOpenFlareMonitors(monitors map[string]Monitor, openFlareTagID int) map[string]Monitor {
existingOpenFlareMonitors := make(map[string]Monitor)
for _, monitor := range monitors {
hasOpenFlareTag := false
for _, tag := range monitor.Tags {
if tag.Name == uptimeKumaTagOpenFlare || tag.ID == openFlareTagID {
hasOpenFlareTag = true
break
}
}
if hasOpenFlareTag {
existingOpenFlareMonitors[monitor.Name] = monitor
}
}
return existingOpenFlareMonitors
}
func routeMonitorURL(ctx context.Context, route *model.ProxyRoute) (string, error) {
if route == nil {
return "", errors.New("proxy route is nil")
}
domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return "", err
}
if len(domains) == 0 {
return "", fmt.Errorf("route %s has no zone domains", route.SiteName)
}
domain := domains[0].Domain
if route.EnableHTTPS {
return "https://" + domain, nil
}
return "http://" + domain, nil
}
func monitorPayload(id int, name, targetURL string, config *kumaConfig) map[string]any {
payload := map[string]any{
"type": "http",
"name": name,
"url": targetURL,
"interval": config.Interval,
"maxretries": config.Retry,
"retryInterval": config.RetryInterval,
"timeout": config.Timeout,
"active": true,
"resendInterval": 0,
"expiryNotification": false,
"ignoreTls": false,
"accepted_statuscodes": []string{"200-299"},
"dns_resolve_type": "A",
"conditions": []any{},
}
if id > 0 {
payload["id"] = id
}
return payload
}
func monitorNeedsUpdate(existing Monitor, targetURL string, config *kumaConfig) bool {
return existing.URL != targetURL ||
existing.Interval != config.Interval ||
existing.MaxRetries != config.Retry ||
existing.RetryInterval != config.RetryInterval ||
existing.Timeout != config.Timeout
}
func createMonitor(client *SocketIOClient, siteName, targetURL string, openFlareTagID int, config *kumaConfig) error {
slog.Info("Creating monitor in Uptime Kuma", "name", siteName, "url", targetURL)
addAck, err := client.Emit("add", monitorPayload(0, siteName, targetURL, config))
if err != nil {
return err
}
var addResult struct {
Ok bool `json:"ok"`
MonitorID int `json:"monitorID"`
}
if err := ParseAckResponse(addAck, &addResult); err != nil || addResult.MonitorID == 0 {
return fmt.Errorf("parse add monitor result failed: %w", err)
}
slog.Debug("Adding OpenFlare tag to the new monitor",
"name", siteName,
"monitor_id", addResult.MonitorID,
"tag_id", openFlareTagID,
)
tagAck, err := client.Emit("addMonitorTag", openFlareTagID, addResult.MonitorID, "")
if err != nil {
return err
}
if err := ParseAckResponse(tagAck, nil); err != nil {
return fmt.Errorf("parse add tag result failed: %w", err)
}
slog.Debug("OpenFlare tag successfully added to monitor", "name", siteName, "monitor_id", addResult.MonitorID)
return nil
}
func updateMonitor(client *SocketIOClient, monitorID int, siteName, targetURL string, config *kumaConfig) error {
slog.Info("Updating monitor in Uptime Kuma due to settings mismatch", "name", siteName)
editAck, err := client.Emit("editMonitor", monitorPayload(monitorID, siteName, targetURL, config))
if err != nil {
return err
}
if err := ParseAckResponse(editAck, nil); err != nil {
return fmt.Errorf("parse edit monitor result failed: %w", err)
}
slog.Info("Successfully updated monitor in Uptime Kuma", "name", siteName)
return nil
}
@@ -0,0 +1,115 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package uptimekuma
import (
"context"
"errors"
"fmt"
"log/slog"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
)
const monitorListWaitTimeout = 5 * time.Second
// validateKumaConfig 验证 kumaConfig 配置完整性
func validateKumaConfig(config *kumaConfig) error {
if strings.TrimSpace(config.URL) == "" {
return errors.New("uptime Kuma URL is not configured")
}
if strings.TrimSpace(config.Username) == "" {
return errors.New("uptime Kuma username is not configured")
}
if strings.TrimSpace(config.Password) == "" {
return errors.New("uptime Kuma password is not configured")
}
return nil
}
func connectAndLoginUptimeKuma(kumaURL, kumaUsername, kumaPassword string) (*SocketIOClient, error) {
slog.Debug("Connecting to Uptime Kuma socket endpoint", "url", kumaURL)
client := NewSocketIOClient(kumaURL)
if err := client.Connect(); err != nil {
slog.Error("Failed to connect to Uptime Kuma endpoint", "url", kumaURL, "error", err)
return nil, fmt.Errorf("failed to connect to Uptime Kuma: %w", err)
}
slog.Debug("Sending login request to Uptime Kuma", "username", kumaUsername)
loginAck, err := client.Emit("login", map[string]string{
"username": kumaUsername,
"password": kumaPassword,
})
if err != nil {
client.Close()
slog.Error("Failed to send login request to Uptime Kuma", "username", kumaUsername, "error", err)
return nil, fmt.Errorf("login request failed: %w", err)
}
var loginResult struct {
Ok bool `json:"ok"`
}
if err := ParseAckResponse(loginAck, &loginResult); err != nil || !loginResult.Ok {
client.Close()
slog.Error("Uptime Kuma login verification failed", "username", kumaUsername, "error", err)
return nil, fmt.Errorf("login failed: %w", err)
}
slog.Debug("Successfully logged into Uptime Kuma", "username", kumaUsername)
slog.Debug("Waiting for monitor list push from Uptime Kuma")
select {
case <-client.GetMonitorListChan():
slog.Debug("Received monitor list from Uptime Kuma")
case <-time.After(monitorListWaitTimeout):
client.Close()
slog.Error("Timeout waiting for Uptime Kuma monitorList push event")
return nil, errors.New("timeout waiting for monitorList event from Uptime Kuma")
}
return client, nil
}
func syncRouteMonitors(ctx context.Context, client *SocketIOClient, expectedRoutes []*model.ProxyRoute, existingMonitors map[string]Monitor, openFlareTagID int, config *kumaConfig) map[string]bool {
expectedSitesMap := make(map[string]bool, len(expectedRoutes))
for _, route := range expectedRoutes {
expectedSitesMap[route.SiteName] = true
targetURL, urlErr := routeMonitorURL(ctx, route)
if urlErr != nil {
slog.Error("Failed to resolve monitor URL", "name", route.SiteName, "error", urlErr)
continue
}
existing, exists := existingMonitors[route.SiteName]
if !exists {
if err := createMonitor(client, route.SiteName, targetURL, openFlareTagID, config); err != nil {
slog.Error("Failed to add monitor to Uptime Kuma", "name", route.SiteName, "error", err)
}
continue
}
if monitorNeedsUpdate(existing, targetURL, config) {
if err := updateMonitor(client, existing.ID, route.SiteName, targetURL, config); err != nil {
slog.Error("Failed to edit monitor in Uptime Kuma", "name", route.SiteName, "error", err)
}
}
}
return expectedSitesMap
}
func removeStaleMonitors(client *SocketIOClient, existingMonitors map[string]Monitor, expectedSitesMap map[string]bool) {
for name, monitor := range existingMonitors {
if expectedSitesMap[name] {
continue
}
slog.Info("Deleting monitor in Uptime Kuma", "name", name, "monitorID", monitor.ID)
deleteAck, err := client.Emit("deleteMonitor", monitor.ID)
if err != nil {
slog.Error("Failed to delete monitor in Uptime Kuma", "name", name, "monitorID", monitor.ID, "error", err)
continue
}
if err := ParseAckResponse(deleteAck, nil); err != nil {
slog.Error("Failed to parse delete monitor result", "name", name, "monitorID", monitor.ID, "error", err)
}
}
}
@@ -0,0 +1,366 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package uptimekuma
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
type mockKumaServer struct {
mu sync.Mutex
postsReceived []string
pendingPackets chan string
monitorList string
}
func newMockKumaServer(monitorList string) *mockKumaServer {
return &mockKumaServer{
pendingPackets: make(chan string, 100),
monitorList: monitorList,
}
}
func (s *mockKumaServer) ServeHTTP(w http.ResponseWriter, r *http.Request) {
s.mu.Lock()
defer s.mu.Unlock()
transport := r.URL.Query().Get("transport")
sid := r.URL.Query().Get("sid")
if r.Method == http.MethodGet {
if transport == "polling" && sid == "" {
w.Header().Set("Content-Type", "text/plain;charset=UTF-8")
_, _ = w.Write([]byte(`0{"sid":"mock-sid"}`))
return
}
if transport == "polling" && sid == "mock-sid" {
w.Header().Set("Content-Type", "text/plain;charset=UTF-8")
select {
case pkt := <-s.pendingPackets:
_, _ = w.Write([]byte(pkt))
case <-time.After(100 * time.Millisecond):
_, _ = w.Write([]byte(""))
}
return
}
} else if r.Method == http.MethodPost {
bodyBytes, _ := io.ReadAll(r.Body)
bodyStr := string(bodyBytes)
s.postsReceived = append(s.postsReceived, bodyStr)
w.Header().Set("Content-Type", "text/plain;charset=UTF-8")
w.WriteHeader(http.StatusOK)
if bodyStr == "40" {
s.pendingPackets <- fmt.Sprintf(`42["monitorList",%s]`, s.monitorList)
return
}
if strings.HasPrefix(bodyStr, "42") {
payload := bodyStr[2:]
digitsEnd := 0
for digitsEnd < len(payload) && payload[digitsEnd] >= '0' && payload[digitsEnd] <= '9' {
digitsEnd++
}
if digitsEnd == 0 {
return
}
ackIDStr := payload[:digitsEnd]
jsonArrayStr := payload[digitsEnd:]
var arr []json.RawMessage
if err := json.Unmarshal([]byte(jsonArrayStr), &arr); err != nil || len(arr) == 0 {
return
}
var eventName string
_ = json.Unmarshal(arr[0], &eventName)
switch eventName {
case "login", "loginByToken":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
case "getTags":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true,\"tags\":[{\"id\":10,\"name\":\"OpenFlare\",\"color\":\"#4f46e5\"}]}]", ackIDStr)
case "addTag":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true,\"tag\":{\"id\":10}}]", ackIDStr)
case "add":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true,\"monitorID\":100}]", ackIDStr)
case "addMonitorTag":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
case "editMonitor":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
case "deleteMonitor":
s.pendingPackets <- fmt.Sprintf("43%s[{\"ok\":true}]", ackIDStr)
}
}
}
}
func setupSyncTestDB(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.ProxyRoute{}, &model.Zone{}, &model.ZoneDomain{}, &model.SystemConfig{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func createRouteZoneDomain(t *testing.T, ctx context.Context, route *model.ProxyRoute, domain string) {
t.Helper()
zone := &model.Zone{Domain: fmt.Sprintf("zone-%d.example", route.ID)}
require.NoError(t, db.DB(ctx).Create(zone).Error)
require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{
ZoneID: zone.ID,
ProxyRouteID: &route.ID,
Domain: domain,
}).Error)
}
func backupUptimeKumaConfig(ctx context.Context) func() {
// 备份所有 UptimeKuma 相关配置
configs := []string{
model.ConfigKeyUptimeKumaEnabled,
model.ConfigKeyUptimeKumaURL,
model.ConfigKeyUptimeKumaUsername,
model.ConfigKeyUptimeKumaPassword,
model.ConfigKeyUptimeKumaMonitorScope,
model.ConfigKeyUptimeKumaSelectedSites,
model.ConfigKeyUptimeKumaInterval,
model.ConfigKeyUptimeKumaRetry,
model.ConfigKeyUptimeKumaRetryInterval,
model.ConfigKeyUptimeKumaTimeout,
}
oldValues := make(map[string]string)
for _, key := range configs {
config, _ := repository.GetSystemConfigByKey(ctx, key)
oldValues[key] = config.Value
}
return func() {
// 恢复所有配置
for key, value := range oldValues {
_ = db.DB(ctx).Model(&model.SystemConfig{}).Where("key = ?", key).Update("value", value).Error
}
}
}
// setTestConfig 设置测试配置的辅助函数(不存在则创建)
func setTestConfig(ctx context.Context, key, value string) {
_ = repository.SaveOrUpdateSystemConfig(ctx, key, value)
}
func TestSyncToUptimeKumaDisabled(t *testing.T) {
cleanup := setupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
restore := backupUptimeKumaConfig(ctx)
defer restore()
setTestConfig(ctx, model.ConfigKeyUptimeKumaEnabled, "false")
err := SyncToUptimeKuma(ctx)
require.Error(t, err)
assert.Contains(t, err.Error(), "disabled")
}
func TestSyncToUptimeKumaSuccess(t *testing.T) {
cleanup := setupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
restore := backupUptimeKumaConfig(ctx)
defer restore()
require.NoError(t, db.DB(ctx).Where("1 = 1").Delete(&model.ProxyRoute{}).Error)
routeA := &model.ProxyRoute{
SiteName: "site-a",
OriginURL: "http://10.0.0.1",
Enabled: true,
EnableHTTPS: false,
}
routeB := &model.ProxyRoute{
SiteName: "site-b",
OriginURL: "https://10.0.0.2",
Enabled: true,
EnableHTTPS: true,
}
routeC := &model.ProxyRoute{
SiteName: "site-c",
OriginURL: "http://10.0.0.3",
Enabled: false,
EnableHTTPS: false,
}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeA))
require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeB))
require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeC))
createRouteZoneDomain(t, ctx, routeA, "site-a.com")
createRouteZoneDomain(t, ctx, routeB, "site-b.com")
createRouteZoneDomain(t, ctx, routeC, "site-c.com")
monitorListJSON := `{
"99": {
"id": 99,
"name": "site-old",
"url": "http://site-old.com",
"interval": 60,
"tags": [{"tag_id": 10, "name": "OpenFlare"}]
},
"98": {
"id": 98,
"name": "site-a",
"url": "http://site-a.com",
"interval": 30,
"tags": [{"tag_id": 10, "name": "OpenFlare"}]
}
}`
mockSrv := newMockKumaServer(monitorListJSON)
server := httptest.NewServer(mockSrv)
defer server.Close()
// 设置测试配置
setTestConfig(ctx, model.ConfigKeyUptimeKumaEnabled, "true")
setTestConfig(ctx, model.ConfigKeyUptimeKumaURL, server.URL)
setTestConfig(ctx, model.ConfigKeyUptimeKumaUsername, "admin")
setTestConfig(ctx, model.ConfigKeyUptimeKumaPassword, "password")
setTestConfig(ctx, model.ConfigKeyUptimeKumaMonitorScope, "all")
setTestConfig(ctx, model.ConfigKeyUptimeKumaInterval, "60")
setTestConfig(ctx, model.ConfigKeyUptimeKumaRetry, "0")
setTestConfig(ctx, model.ConfigKeyUptimeKumaRetryInterval, "60")
setTestConfig(ctx, model.ConfigKeyUptimeKumaTimeout, "48")
require.NoError(t, SyncToUptimeKuma(ctx))
mockSrv.mu.Lock()
posts := mockSrv.postsReceived
mockSrv.mu.Unlock()
hasLogin := false
hasGetTags := false
hasAddSiteB := false
hasTagSiteB := false
hasEditSiteA := false
hasDeleteOld := false
for _, body := range posts {
if strings.Contains(body, `"login"`) && strings.Contains(body, `"admin"`) && strings.Contains(body, `"password"`) {
hasLogin = true
}
if strings.Contains(body, `"getTags"`) {
hasGetTags = true
}
if strings.Contains(body, `"add"`) && strings.Contains(body, `"site-b"`) && strings.Contains(body, `"https://site-b.com"`) {
hasAddSiteB = true
}
if strings.Contains(body, `"addMonitorTag"`) && strings.Contains(body, `10`) && strings.Contains(body, `100`) {
hasTagSiteB = true
}
if strings.Contains(body, `"editMonitor"`) && strings.Contains(body, `98`) && strings.Contains(body, `"site-a"`) && strings.Contains(body, `"interval":60`) {
hasEditSiteA = true
}
if strings.Contains(body, `"deleteMonitor"`) && strings.Contains(body, `99`) {
hasDeleteOld = true
}
}
assert.True(t, hasLogin, "expected login event to be called")
assert.True(t, hasGetTags, "expected getTags event to be called")
assert.True(t, hasAddSiteB, "expected site-b to be added")
assert.True(t, hasTagSiteB, "expected site-b to be tagged")
assert.True(t, hasEditSiteA, "expected site-a to be edited/updated")
assert.True(t, hasDeleteOld, "expected site-old to be deleted")
}
func TestSyncToUptimeKumaSelectedScope(t *testing.T) {
cleanup := setupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
restore := backupUptimeKumaConfig(ctx)
defer restore()
require.NoError(t, db.DB(ctx).Where("1 = 1").Delete(&model.ProxyRoute{}).Error)
routeA := &model.ProxyRoute{
SiteName: "site-a",
OriginURL: "http://10.0.0.1",
Enabled: true,
EnableHTTPS: false,
}
routeB := &model.ProxyRoute{
SiteName: "site-b",
OriginURL: "http://10.0.0.2",
Enabled: true,
EnableHTTPS: false,
}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeA))
require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeB))
createRouteZoneDomain(t, ctx, routeA, "site-a.com")
createRouteZoneDomain(t, ctx, routeB, "site-b.com")
mockSrv := newMockKumaServer(`{}`)
server := httptest.NewServer(mockSrv)
defer server.Close()
// 设置测试配置
setTestConfig(ctx, model.ConfigKeyUptimeKumaEnabled, "true")
setTestConfig(ctx, model.ConfigKeyUptimeKumaURL, server.URL)
setTestConfig(ctx, model.ConfigKeyUptimeKumaUsername, "admin")
setTestConfig(ctx, model.ConfigKeyUptimeKumaPassword, "password")
setTestConfig(ctx, model.ConfigKeyUptimeKumaMonitorScope, "selected")
setTestConfig(ctx, model.ConfigKeyUptimeKumaSelectedSites, "site-a")
require.NoError(t, SyncToUptimeKuma(ctx))
mockSrv.mu.Lock()
posts := mockSrv.postsReceived
mockSrv.mu.Unlock()
hasLogin := false
hasAddSiteA := false
hasAddSiteB := false
for _, body := range posts {
if strings.Contains(body, `"login"`) && strings.Contains(body, `"admin"`) && strings.Contains(body, `"password"`) {
hasLogin = true
}
if strings.Contains(body, `"add"`) && strings.Contains(body, `"site-a"`) {
hasAddSiteA = true
}
if strings.Contains(body, `"add"`) && strings.Contains(body, `"site-b"`) {
hasAddSiteB = true
}
}
assert.True(t, hasLogin, "expected login event to be called")
assert.True(t, hasAddSiteA, "expected site-a to be added")
assert.False(t, hasAddSiteB, "expected site-b NOT to be added (not in selected scope)")
}
@@ -0,0 +1,241 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"context"
"errors"
"fmt"
"regexp"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
)
const maxOpenRestyGzipCompLevel = 9
var (
openRestySizePattern = regexp.MustCompile(`^\d+[kKmMgG]?$`)
openRestyProxyBuffersPattern = regexp.MustCompile(`^\d+\s+\d+[kKmMgG]?$`)
openRestyCacheLevelsPattern = regexp.MustCompile(`^\d{1,2}(?::\d{1,2}){0,2}$`)
openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`)
)
const optionValueTrue = "true"
// protectedConfigKeyMessage 命中受保护 key 时返回给管理员的业务错误文案。
const protectedConfigKeyMessage = "该配置项由系统任务管理,禁止手动修改"
// protectedConfigKeys 仅允许内部(迁移任务/bootstrap)写入的 key。
var protectedConfigKeys = map[string]bool{
model.ConfigKeyLogDatabase: true,
model.ConfigKeyLogDBMigration: true,
}
func isProtectedConfigKey(key string) bool { return protectedConfigKeys[key] }
func buildOptionValidationState(ctx context.Context, options []model.OpenFlareOption) map[string]string {
// 从 SystemConfig 读取所有业务配置构建状态
configs, err := repository.ListAdminSystemConfigs(ctx, "business")
state := make(map[string]string, len(configs)+len(options))
if err == nil {
for _, config := range configs {
state[config.Key] = config.Value
}
}
// 应用待验证的新值
for _, option := range options {
state[option.Key] = option.Value
}
return state
}
func validateOptionWithState(ctx context.Context, option model.OpenFlareOption, state map[string]string) error {
if err := validateOpenRestyOption(option.Key, option.Value); err != nil {
return err
}
if err := validateGeoIPOption(option.Key, option.Value); err != nil {
return err
}
if err := validateLogRetentionOption(option.Key, option.Value); err != nil {
return err
}
if err := validateAgentOption(option.Key, option.Value); err != nil {
return err
}
if err := validatePagesOption(option.Key, option.Value); err != nil {
return err
}
return validateUptimeKumaOption(ctx, option.Key, option.Value, state)
}
func validatePositiveIntegerOption(key, value string) error {
intValue, err := strconv.Atoi(value)
if err != nil || intValue <= 0 {
return fmt.Errorf("%s 必须为大于 0 的整数", key)
}
return nil
}
func validateNonNegativeIntegerOption(key, value string) error {
intValue, err := strconv.Atoi(value)
if err != nil || intValue < 0 {
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
}
return nil
}
func validateBooleanOption(key, value string) error {
switch value {
case optionValueTrue, "false":
return nil
default:
return fmt.Errorf("%s 必须为 true 或 false", key)
}
}
func validateGeoIPOption(key, value string) error {
if key != model.ConfigKeyGeoIPProvider {
return nil
}
if geoip.IsValidProvider(value) {
return nil
}
return fmt.Errorf("%s 仅支持 disabled、mmdb、ip-api、geojs、ipinfo", key)
}
func validateLogRetentionOption(key, value string) error {
switch key {
case model.ConfigKeyLogRetentionDaysPostgres, model.ConfigKeyLogRetentionDaysSQLite, model.ConfigKeyLogRetentionDaysClickHouse:
intValue, err := strconv.Atoi(value)
if err != nil || intValue < 1 {
return fmt.Errorf("%s 必须为大于等于 1 的整数天", key)
}
}
return nil
}
func validateAgentOption(key, value string) error {
if key == model.ConfigKeyAgentWebsocketUpgradeEnabled {
return validateBooleanOption(key, strings.TrimSpace(value))
}
return nil
}
func validatePagesOption(key, value string) error {
trimmed := strings.TrimSpace(value)
switch key {
case model.ConfigKeyPagesMaxPackageSizeMB:
intValue, err := strconv.Atoi(trimmed)
if err != nil || intValue < 1 || intValue > 2048 {
return fmt.Errorf("%s 必须为 1~2048 的整数(MiB)", key)
}
case model.ConfigKeyPagesMaxHistoryCount:
intValue, err := strconv.Atoi(trimmed)
if err != nil || intValue < 0 {
return fmt.Errorf("%s 必须为大于等于 0 的整数(0 表示不限制)", key)
}
}
return nil
}
func validateUptimeKumaOption(ctx context.Context, key, value string, state map[string]string) error {
trimmed := strings.TrimSpace(value)
switch key {
case model.ConfigKeyUptimeKumaEnabled:
return validateUptimeKumaEnabled(ctx, key, trimmed, state)
case model.ConfigKeyUptimeKumaUsername:
return validateUptimeKumaUsername(trimmed, state)
case model.ConfigKeyUptimeKumaURL:
return validateUptimeKumaURL(trimmed)
case model.ConfigKeyUptimeKumaMonitorScope:
return validateUptimeKumaMonitorScope(trimmed)
case model.ConfigKeyUptimeKumaSyncInterval, model.ConfigKeyUptimeKumaInterval, model.ConfigKeyUptimeKumaRetryInterval, model.ConfigKeyUptimeKumaTimeout:
return validatePositiveIntegerOption(key, trimmed)
case model.ConfigKeyUptimeKumaRetry:
return validateUptimeKumaRetry(key, trimmed)
}
return nil
}
func validateUptimeKumaEnabled(ctx context.Context, key, trimmed string, state map[string]string) error {
if err := validateBooleanOption(key, trimmed); err != nil {
return err
}
if trimmed != optionValueTrue {
return nil
}
url := strings.TrimSpace(state[model.ConfigKeyUptimeKumaURL])
username := strings.TrimSpace(state[model.ConfigKeyUptimeKumaUsername])
password := strings.TrimSpace(state[model.ConfigKeyUptimeKumaPassword])
if url == "" {
return errors.New("启用 Uptime Kuma 时地址不能为空")
}
if username == "" {
return errors.New("启用 Uptime Kuma 时用户名不能为空")
}
// 如果待验证的密码为空,且当前配置中也没有密码,则报错
if password == "" {
existingPwd, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaPassword)
if strings.TrimSpace(existingPwd.Value) == "" {
return errors.New("启用 Uptime Kuma 时密码不能为空")
}
}
return nil
}
func validateUptimeKumaUsername(trimmed string, state map[string]string) error {
if trimmed == "" && state[model.ConfigKeyUptimeKumaEnabled] == optionValueTrue {
return errors.New("启用 Uptime Kuma 时用户名不能为空")
}
return nil
}
func validateUptimeKumaURL(trimmed string) error {
if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
return errors.New("uptime Kuma 地址必须以 http:// 或 https:// 开头")
}
return nil
}
func validateUptimeKumaMonitorScope(trimmed string) error {
if trimmed != "all" && trimmed != "selected" {
return errors.New("监控范围必须为全部站点 (all) 或选择站点 (selected)")
}
return nil
}
func validateUptimeKumaRetry(key, trimmed string) error {
intValue, err := strconv.Atoi(trimmed)
if err != nil || intValue < 0 {
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
}
return nil
}
func validateOptions(ctx context.Context, options []model.OpenFlareOption) error {
if len(options) == 0 {
return errors.New(errInvalidParams)
}
state := buildOptionValidationState(ctx, options)
for _, option := range options {
if strings.TrimSpace(option.Key) == "" {
return errors.New(errInvalidParams)
}
if isProtectedConfigKey(option.Key) {
return errors.New(protectedConfigKeyMessage)
}
if err := validateOptionWithState(ctx, option, state); err != nil {
return err
}
}
return nil
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"errors"
"net/url"
"strings"
"Wavelet/openflare/share/pagesarchive"
)
// downloadPagesPackageFromURL is the deprecated one-shot URL adapter. It uses
// the same bounded downloader as persisted sources and allows insecure TLS for
// operator-managed internal artifact services.
func downloadPagesPackageFromURL(
ctx context.Context,
rawURL string,
maxPackageBytes int64,
) (tempPath string, checksum string, size int64, format pagesarchive.Format, fileName string, err error) {
if _, err := parseAndValidatePagesDownloadURL(rawURL); err != nil {
return "", "", 0, "", "", err
}
candidate, err := FetchRemoteSource(ctx, RemoteSourceRequest{
URL: strings.TrimSpace(rawURL),
AllowInsecure: true,
MaxPackageBytes: maxPackageBytes,
})
if err != nil {
if strings.Contains(err.Error(), errPagesSourceRemoteURLInvalid) {
return "", "", 0, "", "", errors.New(errPagesPackageURLInvalid)
}
return "", "", 0, "", "", err
}
// Ownership transfers to the existing one-shot caller, which removes the
// temporary file after the candidate deployment has been created.
return candidate.TempPath, candidate.Checksum, candidate.PackageSize, candidate.Format, candidate.SafeLabel, nil
}
func parseAndValidatePagesDownloadURL(raw string) (*url.URL, error) {
value := strings.TrimSpace(raw)
if value == "" {
return nil, errors.New(errPagesPackageURLRequired)
}
parsed, err := url.Parse(value)
if err != nil || parsed.User != nil || parsed.Fragment != "" || parsed.Opaque != "" {
return nil, errors.New(errPagesPackageURLInvalid)
}
scheme := strings.ToLower(strings.TrimSpace(parsed.Scheme))
if (scheme != remoteSourceSchemeHTTP && scheme != remoteSourceSchemeHTTPS) || strings.TrimSpace(parsed.Hostname()) == "" {
return nil, errors.New(errPagesPackageURLInvalid)
}
return parsed, nil
}

Some files were not shown because too many files have changed in this diff Show More