mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 23:16:37 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -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, ®istration)
|
||||
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
Reference in New Issue
Block a user