mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 14:56:38 +08:00
feat(cloudflare): add DNS pointing integration
Implement Cloudflare connection management, pointing groups and members, asynchronous A-record reconciliation, node IP triggers, admin APIs, management pages, migrations, tests, and documentation.
This commit is contained in:
@@ -9,11 +9,13 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
cf "github.com/Rain-kl/Wavelet/internal/apps/openflare/cloudflare"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
|
||||
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
)
|
||||
|
||||
// RegisterWithAccessToken registers an agent on a reserved node token.
|
||||
@@ -118,6 +120,11 @@ func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload N
|
||||
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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
"github.com/Rain-kl/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 := 0; attempt < maxRequestAttempts; attempt++ {
|
||||
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,453 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cloudflare
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/credential"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/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
|
||||
}
|
||||
items := make([]AvailableDomain, 0, len(domains))
|
||||
for _, domain := range domains {
|
||||
items = append(items, AvailableDomain{ID: domain.ID, ZoneID: domain.ZoneID, Domain: domain.Domain})
|
||||
}
|
||||
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 {
|
||||
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,154 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cloudflare
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
)
|
||||
|
||||
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 {
|
||||
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,218 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cloudflare
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/credential"
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"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,383 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cloudflare
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/shared/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,43 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cloudflare
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||
"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())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cloudflare
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
)
|
||||
|
||||
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.
|
||||
var SyncMemberMeta = task.TaskMeta{Type: TaskTypeSyncMember, AsynqTask: SyncMemberTask, Name: "Cloudflare 域名同步", Description: "同步单个域名的 Cloudflare A 记录", MaxRetry: 3, Queue: task.QueueDefault, Retryable: true, InternalOnly: true}
|
||||
|
||||
// SyncGroupMeta describes group reconciliation.
|
||||
var SyncGroupMeta = task.TaskMeta{Type: TaskTypeSyncGroup, AsynqTask: SyncGroupTask, Name: "Cloudflare 分组同步", Description: "同步指向分组内全部域名", MaxRetry: 2, Queue: task.QueueDefault, Retryable: true, InternalOnly: true}
|
||||
|
||||
// SyncByNodeMeta describes node-triggered reconciliation.
|
||||
var SyncByNodeMeta = task.TaskMeta{Type: TaskTypeSyncByNode, AsynqTask: SyncByNodeTask, Name: "Cloudflare 节点同步", Description: "同步当前指向指定节点的全部域名", 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) {
|
||||
var input SyncMemberPayload
|
||||
if err := decodePayload(payload, &input); err != nil || input.MemberID == 0 {
|
||||
return nil, errors.New("无效的 Cloudflare 成员同步参数")
|
||||
}
|
||||
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)
|
||||
task.AppendLog(ctx, "正在同步 Cloudflare 成员 ID=%d", input.MemberID)
|
||||
if err = ReconcileMember(ctx, input.MemberID); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", errSyncFailed, err)
|
||||
}
|
||||
return &task.TaskResult{Message: "Cloudflare 域名同步成功"}, 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) {
|
||||
var input SyncGroupPayload
|
||||
if err := decodePayload(payload, &input); err != nil || input.GroupID == 0 {
|
||||
return nil, errors.New("无效的 Cloudflare 分组同步参数")
|
||||
}
|
||||
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())
|
||||
}
|
||||
members, err := repository.ListCFPointingMembersByGroupID(ctx, input.GroupID)
|
||||
return executeBatchSync(ctx, members, err, "分组")
|
||||
}
|
||||
|
||||
// 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())
|
||||
}
|
||||
members, err := repository.ListCFPointingMembersByActiveNodeID(ctx, input.NodeID)
|
||||
return executeBatchSync(ctx, members, err, "节点")
|
||||
}
|
||||
|
||||
func executeBatchSync(ctx context.Context, members []model.CFPointingMember, listErr error, scope string) (*task.TaskResult, error) {
|
||||
if listErr != nil {
|
||||
return nil, listErr
|
||||
}
|
||||
for _, member := range members {
|
||||
if err := ReconcileMember(ctx, member.ID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return &task.TaskResult{Message: fmt.Sprintf("Cloudflare %s同步完成,共 %d 个域名", scope, len(members))}, 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,106 @@
|
||||
// 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"`
|
||||
}
|
||||
|
||||
// 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,60 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package credential seals and opens OpenFlare integration credentials.
|
||||
package credential
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
// Prefix identifies values encrypted with the current credential format.
|
||||
const Prefix = "enc:v1:"
|
||||
|
||||
func encryptionKey() string {
|
||||
if config.Config == nil || strings.TrimSpace(config.Config.App.SessionSecret) == "" {
|
||||
return ""
|
||||
}
|
||||
sum := sha256.Sum256([]byte(config.Config.App.SessionSecret))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// Seal trims and encrypts plaintext when a session secret is configured.
|
||||
// Plaintext storage is preserved for installations without a session secret.
|
||||
func Seal(plaintext string) (string, error) {
|
||||
plaintext = strings.TrimSpace(plaintext)
|
||||
if plaintext == "" {
|
||||
return "", nil
|
||||
}
|
||||
key := encryptionKey()
|
||||
if key == "" {
|
||||
return plaintext, nil
|
||||
}
|
||||
encrypted, err := util.Encrypt(key, plaintext)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return Prefix + encrypted, nil
|
||||
}
|
||||
|
||||
// Open decrypts a sealed value and accepts legacy plaintext values.
|
||||
func Open(stored string) (string, error) {
|
||||
stored = strings.TrimSpace(stored)
|
||||
if stored == "" {
|
||||
return "", nil
|
||||
}
|
||||
if !strings.HasPrefix(stored, Prefix) {
|
||||
return stored, nil
|
||||
}
|
||||
key := encryptionKey()
|
||||
if key == "" {
|
||||
return "", errors.New("cannot decrypt sensitive field without session secret")
|
||||
}
|
||||
return util.Decrypt(key, strings.TrimPrefix(stored, Prefix))
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package credential
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
)
|
||||
|
||||
func TestSealAndOpenSensitiveValue(t *testing.T) {
|
||||
previous := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = "cloudflare-pointing-test-secret"
|
||||
t.Cleanup(func() { config.Config.App.SessionSecret = previous })
|
||||
|
||||
sealed, err := Seal(`{"api_token":"secret-token"}`)
|
||||
if err != nil {
|
||||
t.Fatalf("Seal() error = %v", err)
|
||||
}
|
||||
if !strings.HasPrefix(sealed, Prefix) {
|
||||
t.Fatalf("Seal() = %q, want prefix %q", sealed, Prefix)
|
||||
}
|
||||
if strings.Contains(sealed, "secret-token") {
|
||||
t.Fatalf("Seal() = %q, want token redacted", sealed)
|
||||
}
|
||||
|
||||
opened, err := Open(sealed)
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
if want := `{"api_token":"secret-token"}`; opened != want {
|
||||
t.Errorf("Open(Seal(value)) = %q, want %q", opened, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSealWithoutSessionSecretKeepsPlaintextCompatibility(t *testing.T) {
|
||||
previous := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = ""
|
||||
t.Cleanup(func() { config.Config.App.SessionSecret = previous })
|
||||
|
||||
sealed, err := Seal(" legacy-value ")
|
||||
if err != nil {
|
||||
t.Fatalf("Seal() error = %v", err)
|
||||
}
|
||||
if sealed != "legacy-value" {
|
||||
t.Errorf("Seal() = %q, want %q", sealed, "legacy-value")
|
||||
}
|
||||
|
||||
opened, err := Open(sealed)
|
||||
if err != nil {
|
||||
t.Fatalf("Open(plaintext) error = %v", err)
|
||||
}
|
||||
if opened != "legacy-value" {
|
||||
t.Errorf("Open(plaintext) = %q, want %q", opened, "legacy-value")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenEncryptedValueRequiresSessionSecret(t *testing.T) {
|
||||
previous := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = "cloudflare-pointing-test-secret"
|
||||
sealed, err := Seal("secret")
|
||||
if err != nil {
|
||||
t.Fatalf("Seal() error = %v", err)
|
||||
}
|
||||
|
||||
config.Config.App.SessionSecret = ""
|
||||
t.Cleanup(func() { config.Config.App.SessionSecret = previous })
|
||||
if _, err := Open(sealed); err == nil {
|
||||
t.Fatal("Open(encrypted) error = nil, want missing session secret error")
|
||||
}
|
||||
}
|
||||
@@ -10,10 +10,12 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
cf "github.com/Rain-kl/Wavelet/internal/apps/openflare/cloudflare"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/observability"
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -232,6 +234,7 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) {
|
||||
return nil, err
|
||||
}
|
||||
ipManualOverride := resolveNodeIPManualOverride(input, node, ip)
|
||||
previousIP := node.IP
|
||||
node.Name = name
|
||||
node.IP = ip
|
||||
node.IPManualOverride = ipManualOverride
|
||||
@@ -255,6 +258,11 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) {
|
||||
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
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cf "github.com/Rain-kl/Wavelet/internal/apps/openflare/cloudflare"
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
@@ -130,6 +131,27 @@ func TestUpdateNode(t *testing.T) {
|
||||
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()
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/credential"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
@@ -128,7 +129,7 @@ func TestCreateCertificateEncryptsPrivateKey(t *testing.T) {
|
||||
stored, err := repository.GetTLSCertificateByID(ctx, certificate.ID)
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, keyPEM, stored.KeyPEM)
|
||||
assert.Contains(t, stored.KeyPEM, sensitiveValuePrefix)
|
||||
assert.Contains(t, stored.KeyPEM, credential.Prefix)
|
||||
|
||||
content, err := GetCertificateContent(ctx, certificate.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -3,40 +3,10 @@
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
const sensitiveValuePrefix = "enc:v1:"
|
||||
|
||||
func sensitiveEncryptionKey() string {
|
||||
if config.Config == nil || strings.TrimSpace(config.Config.App.SessionSecret) == "" {
|
||||
return ""
|
||||
}
|
||||
sum := sha256.Sum256([]byte(config.Config.App.SessionSecret))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
import "github.com/Rain-kl/Wavelet/internal/apps/openflare/credential"
|
||||
|
||||
func sealSensitive(plaintext string) (string, error) {
|
||||
plaintext = strings.TrimSpace(plaintext)
|
||||
if plaintext == "" {
|
||||
return "", nil
|
||||
}
|
||||
key := sensitiveEncryptionKey()
|
||||
if key == "" {
|
||||
return plaintext, nil
|
||||
}
|
||||
encrypted, err := util.Encrypt(key, plaintext)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return sensitiveValuePrefix + encrypted, nil
|
||||
return credential.Seal(plaintext)
|
||||
}
|
||||
|
||||
// OpenKeyPEM decrypts a stored certificate private key for runtime distribution.
|
||||
@@ -45,16 +15,5 @@ func OpenKeyPEM(stored string) (string, error) {
|
||||
}
|
||||
|
||||
func openSensitive(stored string) (string, error) {
|
||||
stored = strings.TrimSpace(stored)
|
||||
if stored == "" {
|
||||
return "", nil
|
||||
}
|
||||
if !strings.HasPrefix(stored, sensitiveValuePrefix) {
|
||||
return stored, nil
|
||||
}
|
||||
key := sensitiveEncryptionKey()
|
||||
if key == "" {
|
||||
return "", errors.New("cannot decrypt sensitive field without session secret")
|
||||
}
|
||||
return util.Decrypt(key, strings.TrimPrefix(stored, sensitiveValuePrefix))
|
||||
return credential.Open(stored)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user