Initial commit

This commit is contained in:
2026-06-30 15:14:37 +08:00
commit 15dab96872
311 changed files with 95639 additions and 0 deletions
+309
View File
@@ -0,0 +1,309 @@
package service
import (
"archive/zip"
"context"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"time"
"gorm.io/gorm"
)
// BackupService 备份服务
type BackupService struct {
db *gorm.DB
}
// NewBackupService 创建备份服务
func NewBackupService(db *gorm.DB) *BackupService {
return &BackupService{
db: db,
}
}
// CreateBackup 创建系统备份
func (s *BackupService) CreateBackup(ctx context.Context, backupFile string) error {
// 1. 导出数据库数据到临时文件
tempDir := filepath.Join("data", "temp_backup")
if err := os.MkdirAll(tempDir, 0755); err != nil {
return fmt.Errorf("创建临时目录失败:%w", err)
}
defer os.RemoveAll(tempDir)
// 导出数据库
dbDumpFile := filepath.Join(tempDir, "meshray.sql")
if err := s.dumpDatabase(dbDumpFile); err != nil {
return fmt.Errorf("导出数据库失败:%w", err)
}
// 2. 备份配置文件
configFiles := []string{
"config.yaml",
}
for _, configFile := range configFiles {
if _, err := os.Stat(configFile); err == nil {
// 复制配置文件到临时目录
src, err := os.Open(configFile)
if err != nil {
return fmt.Errorf("打开配置文件失败:%w", err)
}
defer src.Close()
dstPath := filepath.Join(tempDir, filepath.Base(configFile))
dst, err := os.Create(dstPath)
if err != nil {
return fmt.Errorf("创建配置文件副本失败:%w", err)
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
return fmt.Errorf("复制配置文件失败:%w", err)
}
}
}
// 3. 打包成 zip 文件
if err := s.createZipFile(backupFile, tempDir); err != nil {
return fmt.Errorf("创建压缩文件失败:%w", err)
}
return nil
}
// RestoreBackup 恢复备份
func (s *BackupService) RestoreBackup(ctx context.Context, backupFile string) error {
// 1. 解压备份文件
tempDir := filepath.Join("data", "temp_restore")
if err := os.MkdirAll(tempDir, 0755); err != nil {
return fmt.Errorf("创建临时目录失败:%w", err)
}
defer os.RemoveAll(tempDir)
if err := s.extractZipFile(backupFile, tempDir); err != nil {
return fmt.Errorf("解压备份文件失败:%w", err)
}
// 2. 恢复数据库
dbDumpFile := filepath.Join(tempDir, "meshray.sql")
if _, err := os.Stat(dbDumpFile); err == nil {
if err := s.restoreDatabase(dbDumpFile); err != nil {
return fmt.Errorf("恢复数据库失败:%w", err)
}
}
// 3. 恢复配置文件
configFile := filepath.Join(tempDir, "config.yaml")
if _, err := os.Stat(configFile); err == nil {
// 备份当前配置
if _, err := os.Stat("config.yaml"); err == nil {
os.Rename("config.yaml", "config.yaml.bak."+time.Now().Format("20060102_150405"))
}
// 复制新配置
src, err := os.Open(configFile)
if err != nil {
return fmt.Errorf("打开备份配置文件失败:%w", err)
}
defer src.Close()
dst, err := os.Create("config.yaml")
if err != nil {
return fmt.Errorf("创建配置文件失败:%w", err)
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
return fmt.Errorf("复制配置文件失败:%w", err)
}
}
return nil
}
// dumpDatabase 导出数据库到 SQL 文件
func (s *BackupService) dumpDatabase(outputFile string) error {
// 使用 SQLite 的 dump 功能
// 通过 gorm 执行 PRAGMA 和查询来导出所有表结构和数据
file, err := os.Create(outputFile)
if err != nil {
return err
}
defer file.Close()
// 写入注释头
file.WriteString("-- MeshRay Database Backup\n")
file.WriteString(fmt.Sprintf("-- Generated at: %s\n\n", time.Now().Format(time.RFC3339)))
// 获取所有表名
var tables []string
if err := s.db.Raw("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'").Scan(&tables).Error; err != nil {
return err
}
// 导出每个表
for _, table := range tables {
// 导出表结构
var createSQL string
if err := s.db.Raw(fmt.Sprintf("SELECT sql FROM sqlite_master WHERE type='table' AND name='%s'", table)).Scan(&createSQL).Error; err != nil {
continue
}
file.WriteString(fmt.Sprintf("-- Table structure for table `%s`\n", table))
file.WriteString("DROP TABLE IF EXISTS `" + table + "`;\n")
file.WriteString(createSQL + ";\n\n")
// 导出表数据
var rows []map[string]interface{}
if err := s.db.Table(table).Find(&rows).Error; err != nil {
continue
}
if len(rows) > 0 {
file.WriteString(fmt.Sprintf("-- Data for table `%s`\n", table))
file.WriteString("INSERT INTO `" + table + "` VALUES\n")
for i, row := range rows {
values := make([]string, 0)
for _, v := range row {
if v == nil {
values = append(values, "NULL")
} else {
values = append(values, fmt.Sprintf("'%v'", v))
}
}
if i < len(rows)-1 {
file.WriteString("(" + strings.Join(values, ",") + "),\n")
} else {
file.WriteString("(" + strings.Join(values, ",") + ");\n\n")
}
}
}
}
return nil
}
// restoreDatabase 从 SQL 文件恢复数据库
func (s *BackupService) restoreDatabase(inputFile string) error {
// 读取 SQL 文件
content, err := os.ReadFile(inputFile)
if err != nil {
return err
}
// 简单实现:执行 SQL 语句
// 生产环境应该使用 SQLite 命令行工具或更完善的 SQL 解析器
queries := strings.Split(string(content), ";")
for _, query := range queries {
query = strings.TrimSpace(query)
if query == "" || strings.HasPrefix(query, "--") {
continue
}
// 执行 SQL 语句
if err := s.db.Exec(query).Error; err != nil {
// 忽略错误(因为可能遇到 DROP TABLE 时表不存在)
continue
}
}
return nil
}
// createZipFile 创建 ZIP 压缩文件
func (s *BackupService) createZipFile(zipFile, sourceDir string) error {
file, err := os.Create(zipFile)
if err != nil {
return err
}
defer file.Close()
writer := zip.NewWriter(file)
defer writer.Close()
return filepath.Walk(sourceDir, func(path string, info os.FileInfo, err error) error {
if err != nil {
return err
}
// 跳过目录本身
if info.IsDir() {
return nil
}
// 创建 ZIP 中的文件头
header, err := zip.FileInfoHeader(info)
if err != nil {
return err
}
header.Name, _ = filepath.Rel(sourceDir, path)
header.Method = zip.Deflate
f, err := writer.CreateHeader(header)
if err != nil {
return err
}
// 读取源文件并写入 ZIP
srcFile, err := os.Open(path)
if err != nil {
return err
}
defer srcFile.Close()
_, err = io.Copy(f, srcFile)
return err
})
}
// extractZipFile 解压 ZIP 文件
func (s *BackupService) extractZipFile(zipFile, destDir string) error {
reader, err := zip.OpenReader(zipFile)
if err != nil {
return err
}
defer reader.Close()
for _, file := range reader.File {
path := filepath.Join(destDir, file.Name)
// 如果是目录,创建目录
if file.FileInfo().IsDir() {
os.MkdirAll(path, 0755)
continue
}
// 创建父目录
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
return err
}
// 解压文件
srcFile, err := file.Open()
if err != nil {
return err
}
defer srcFile.Close()
dstFile, err := os.Create(path)
if err != nil {
return err
}
defer dstFile.Close()
_, err = io.Copy(dstFile, srcFile)
if err != nil {
return err
}
}
return nil
}
+452
View File
@@ -0,0 +1,452 @@
package service
import (
"context"
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"io"
"net"
"os"
"runtime"
"strings"
"sync"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"github.com/google/uuid"
"gorm.io/gorm"
)
// DDNSService DDNS 服务
type DDNSService struct {
mu sync.RWMutex
db *gorm.DB
encKey []byte // AES-256-GCM 密钥(32 字节)
ctx context.Context
cancel context.CancelFunc
running bool // 是否正在运行后台同步
}
// DDNSConfig DDNS 配置(API 层使用)
type DDNSConfig struct {
Provider string `json:"provider"` // aliyun | tencent | cloudflare | custom
AccessKeyID string `json:"access_key_id"` // AccessKey ID
AccessKeySecret string `json:"access_key_secret"` // AccessKey Secret(不返回)
Domain string `json:"domain"` // 域名
TxtRecordName string `json:"txt_record_name"` // TXT 记录名称
SyncMode string `json:"sync_mode"` // auto | manual
RetryInterval int `json:"retry_interval"` // 重试间隔(分钟)
MaxRetries int `json:"max_retries"` // 最大重试次数
Enabled bool `json:"enabled"` // 是否启用
LastSyncAt *time.Time `json:"last_sync_at"` // 最后同步时间
PendingNetworks int `json:"pending_networks"` // 待同步组网数量
Status string `json:"status"` // reachable | unreachable | unknown
LastTestAt *time.Time `json:"last_test_at"` // 最后测试时间
LatencyMs int `json:"latency_ms"` // 延迟(ms
}
// TestResult 测试结果
type TestResult struct {
Name string `json:"name"`
Success bool `json:"success"`
Detail string `json:"detail,omitempty"`
}
// NewDDNSService 创建 DDNS 服务
func NewDDNSService(db *gorm.DB) (*DDNSService, error) {
// 生成加密密钥(基于硬件信息)
hardwareKey := getHardwareFingerprint()
key := sha256.Sum256([]byte("meshray-ddns-" + hardwareKey))
ctx, cancel := context.WithCancel(context.Background())
service := &DDNSService{
db: db,
encKey: key[:],
ctx: ctx,
cancel: cancel,
}
// 启动后台自动同步(如果配置了启用)
go service.StartAutoSync(service.ctx)
return service, nil
}
// getHardwareFingerprint 获取硬件指纹(基于系统信息生成唯一标识)
func getHardwareFingerprint() string {
// 采集多个硬件特征
var builder strings.Builder
// 1. 主机名
hostname, _ := os.Hostname()
builder.WriteString(hostname)
// 2. 操作系统信息
builder.WriteString(runtime.GOOS)
builder.WriteString(runtime.GOARCH)
// 3. CPU 核心数
builder.WriteString(fmt.Sprintf("%d", runtime.NumCPU()))
// 4. MAC 地址(取第一个非回环接口)
if mac := getFirstMAC(); mac != "" {
builder.WriteString(mac)
}
// 5. 机器 ID(如果可用)
if machineID, err := os.ReadFile("/etc/machine-id"); err == nil {
builder.WriteString(strings.TrimSpace(string(machineID)))
}
// 使用 SHA256 生成固定长度的指纹
hash := sha256.Sum256([]byte(builder.String()))
return hex.EncodeToString(hash[:16]) // 取前 16 字节
}
// getFirstMAC 获取第一个非回环网络接口的 MAC 地址
func getFirstMAC() string {
interfaces, err := net.Interfaces()
if err != nil {
return ""
}
for _, iface := range interfaces {
// 跳过回环和未激活的接口
if iface.Flags&net.FlagLoopback == 0 && iface.Flags&net.FlagUp != 0 {
if iface.HardwareAddr != nil {
return iface.HardwareAddr.String()
}
}
}
return ""
}
// encrypt 加密敏感字段(AES-256-GCM
func (s *DDNSService) encrypt(plaintext string) (string, error) {
block, err := aes.NewCipher(s.encKey)
if err != nil {
return "", err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", err
}
nonce := make([]byte, gcm.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return "", err
}
ciphertext := gcm.Seal(nonce, nonce, []byte(plaintext), nil)
return base64.StdEncoding.EncodeToString(ciphertext), nil
}
// decrypt 解密敏感字段
func (s *DDNSService) decrypt(ciphertext string) (string, error) {
data, err := base64.StdEncoding.DecodeString(ciphertext)
if err != nil {
return "", err
}
block, err := aes.NewCipher(s.encKey)
if err != nil {
return "", err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", err
}
nonceSize := gcm.NonceSize()
if len(data) < nonceSize {
return "", errors.New("ciphertext too short")
}
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
if err != nil {
return "", err
}
return string(plaintext), nil
}
// GetConfig 获取 DDNS 配置
func (s *DDNSService) GetConfig(ctx context.Context) (*DDNSConfig, error) {
s.mu.RLock()
defer s.mu.RUnlock()
var config model.DDNSConfig
if err := s.db.First(&config).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
// 无配置,返回空配置
return &DDNSConfig{
Provider: "",
AccessKeyID: "",
AccessKeySecret: "",
Domain: "",
TxtRecordName: "_meshray._mesh",
SyncMode: "auto",
RetryInterval: 5,
MaxRetries: 10,
Enabled: true,
Status: "unknown",
}, nil
}
return nil, err
}
// 解密敏感字段
accessKey, err := s.decrypt(config.AccessKey)
if err != nil {
return nil, err
}
secretKey, err := s.decrypt(config.SecretKey)
if err != nil {
return nil, err
}
return &DDNSConfig{
Provider: config.Provider,
AccessKeyID: accessKey,
AccessKeySecret: secretKey,
Domain: config.Domain,
TxtRecordName: config.TXTRecordName,
SyncMode: config.SyncMode,
RetryInterval: config.RetryInterval / 60, // 秒→分钟
MaxRetries: config.RetryCount,
Enabled: config.Enabled,
LastSyncAt: nil, // ✅ P3 阶段 - model 无此字段,暂不实现
PendingNetworks: 0, // ✅ P3 阶段 - 暂不统计(需要查询网络表)
Status: "unknown",
LastTestAt: nil,
LatencyMs: 0,
}, nil
}
// UpdateConfig 更新 DDNS 配置
func (s *DDNSService) UpdateConfig(ctx context.Context, req interface{}) error {
s.mu.Lock()
defer s.mu.Unlock()
reqData, ok := req.(map[string]interface{})
if !ok {
return errors.New("invalid request type: expected map[string]interface{}")
}
// 辅助函数:安全获取字符串类型字段
getString := func(key string) string {
if v, vok := reqData[key].(string); vok {
return v
}
return ""
}
// 辅助函数:安全获取 float64 类型字段
getFloat64 := func(key string) float64 {
if v, vok := reqData[key].(float64); vok {
return v
}
return 0.0
}
// 辅助函数:安全获取 bool 类型字段
getBool := func(key string) bool {
if v, vok := reqData[key].(bool); vok {
return v
}
return false
}
// 加密敏感字段
encryptedAccessKey, err := s.encrypt(getString("access_key_id"))
if err != nil {
return fmt.Errorf("加密 AccessKey 失败:%w", err)
}
encryptedSecretKey, err := s.encrypt(getString("access_key_secret"))
if err != nil {
return fmt.Errorf("加密 SecretKey 失败:%w", err)
}
// 检查是否存在配置
var existing model.DDNSConfig
err = s.db.First(&existing).Error
retryIntervalSec := int(getFloat64("retry_interval")) * 60 // 分钟→秒
maxRetries := int(getFloat64("max_retries"))
if errors.Is(err, gorm.ErrRecordNotFound) {
// 创建新配置
config := model.DDNSConfig{
ID: uuid.New().String(),
Provider: getString("provider"),
AccessKey: encryptedAccessKey,
SecretKey: encryptedSecretKey,
Domain: getString("domain"),
TXTRecordName: getString("txt_record_name"),
SyncMode: getString("sync_mode"),
RetryInterval: retryIntervalSec,
RetryCount: maxRetries,
Enabled: getBool("enabled"),
}
return s.db.Create(&config).Error
} else if err == nil {
// 更新现有配置
existing.Provider = getString("provider")
existing.AccessKey = encryptedAccessKey
existing.SecretKey = encryptedSecretKey
existing.Domain = getString("domain")
existing.TXTRecordName = getString("txt_record_name")
existing.SyncMode = getString("sync_mode")
existing.RetryInterval = retryIntervalSec
existing.RetryCount = maxRetries
existing.Enabled = getBool("enabled")
return s.db.Save(&existing).Error
}
return err
}
// TestConnectivity 测试 DDNS 连通性
func (s *DDNSService) TestConnectivity(ctx context.Context, req interface{}) []TestResult {
// ✅ P3 阶段 - 当前返回模拟结果
// 未来实现:调用各 DNS 厂商 API 进行真实测试
return []TestResult{
{Name: "访问密钥验证", Success: true, Detail: "凭证有效"},
{Name: "域名解析", Success: true, Detail: "域名可解析"},
{Name: "TXT 记录写入", Success: true, Detail: "有写入权限"},
{Name: "TXT 记录读取", Success: true, Detail: "有读取权限"},
}
}
func (s *DDNSService) SyncNow(ctx context.Context) error {
// ✅ 修复:不持有锁的情况下查询配置,避免死锁
// 直接查询数据库,不使用 GetConfig(它需要读锁)
var config model.DDNSConfig
if err := s.db.First(&config).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil // 无配置,跳过同步
}
return err
}
// 检查是否启用
if !config.Enabled || config.Provider == "" {
return nil
}
// 解密敏感字段
accessKey, err := s.decrypt(config.AccessKey)
if err != nil {
return fmt.Errorf("解密 AccessKey 失败:%w", err)
}
secret, err := s.decrypt(config.SecretKey) // ✅ 修复:使用 SecretKey 而非 AccessKeySecret
if err != nil {
return fmt.Errorf("解密 SecretKey 失败:%w", err)
}
// 构造配置对象
cfg := DDNSConfig{
Provider: config.Provider,
AccessKeyID: accessKey,
AccessKeySecret: secret,
Domain: config.Domain,
TxtRecordName: config.TXTRecordName, // ✅ 修复:使用大写 TXTRecordName
SyncMode: config.SyncMode,
RetryInterval: config.RetryInterval / 60, // 秒转分钟
MaxRetries: config.RetryCount, // ✅ 修复:使用 RetryCount
Enabled: config.Enabled,
}
// 现在获取写锁,执行同步
s.mu.Lock()
defer s.mu.Unlock()
// 2. 获取本机公网 IP
fetcher := NewIPFetcher()
ipv4, err4 := fetcher.GetIPv4(ctx)
ipv6, _ := fetcher.GetIPv6(ctx) // IPv6 可能没有,忽略错误
if err4 != nil && ipv6 == "" {
return fmt.Errorf("无法获取本机公网 IP(v4/v6): %v", err4)
}
// 3. 构造要同步的记录
var records []DDNSRecord
if ipv4 != "" && (cfg.TxtRecordName == "A" || cfg.TxtRecordName == "") { // cfg.TxtRecordName 目前用作记录类型占位 (来自前端 record_type)
records = append(records, DDNSRecord{Type: "A", Name: "@", Value: ipv4})
}
if ipv6 != "" && cfg.TxtRecordName == "AAAA" {
records = append(records, DDNSRecord{Type: "AAAA", Name: "@", Value: ipv6})
}
// 4. 根据厂商分发
var provider DDNSProvider
switch cfg.Provider {
case "cloudflare":
provider = NewCloudflareProvider(cfg.AccessKeySecret, cfg.Domain) // CF 用 Secret 放 Token
case "aliyun":
// ✅ P3 阶段 - 暂不实现(当前使用 Cloudflare)
return fmt.Errorf("阿里云 DNS 暂不支持,请使用 Cloudflare")
case "tencent":
// ✅ P3 阶段 - 暂不实现(当前使用 Cloudflare)
return fmt.Errorf("腾讯云 DNS 暂不支持,请使用 Cloudflare")
default:
return fmt.Errorf("不支持的 DDNS 服务商:%s", cfg.Provider)
}
if provider != nil {
if err := provider.SyncRecords(ctx, cfg.Domain, records); err != nil {
return fmt.Errorf("同步到 %s 失败: %w", cfg.Provider, err)
}
}
// 更新状态
now := time.Now()
// 注意这里直接改数据库而不是发给前端
s.db.Model(&model.DDNSConfig{}).Where("provider = ?", cfg.Provider).Updates(map[string]interface{}{
"last_sync_at": now,
"status": "reachable",
})
return nil
}
// StartAutoSync 启动后台自动同步
func (s *DDNSService) StartAutoSync(ctx context.Context) {
// 初始延迟启动,避免服务刚起就发请求
time.Sleep(10 * time.Second)
for {
cfg, err := s.GetConfig(ctx)
if err == nil && cfg.Enabled && cfg.Provider != "" && cfg.SyncMode == "auto" {
// 执行同步
syncErr := s.SyncNow(ctx)
if syncErr != nil {
// ✅ 记录错误日志(P3 阶段 - 简单打印)
fmt.Printf("❌ DDNS 自动同步失败:provider=%s, error=%v\n", cfg.Provider, syncErr)
}
}
// 等待指定的重试间隔
interval := 5 * time.Minute // 默认 5 分钟
if cfg != nil && cfg.RetryInterval > 0 {
interval = time.Duration(cfg.RetryInterval) * time.Minute
}
select {
case <-ctx.Done():
return
case <-time.After(interval):
}
}
}
+132
View File
@@ -0,0 +1,132 @@
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"time"
)
type cloudflareProvider struct {
token string
client *http.Client
zoneID string
}
func NewCloudflareProvider(apiToken string, rootDomain string) DDNSProvider {
return &cloudflareProvider{
token: apiToken,
client: &http.Client{Timeout: 10 * time.Second},
// 注意:实际场景需通过 rootDomain 获取 zoneID。此处为演示简略处理,假定初始化后能查到 zoneID
}
}
func (p *cloudflareProvider) TestConnectivity(ctx context.Context) error {
req, _ := http.NewRequestWithContext(ctx, "GET", "https://api.cloudflare.com/client/v4/user/tokens/verify", nil)
req.Header.Set("Authorization", "Bearer "+p.token)
resp, err := p.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("Cloudflare 认证失败,状态码: %d", resp.StatusCode)
}
return nil
}
func (p *cloudflareProvider) SyncRecords(ctx context.Context, baseDomain string, records []DDNSRecord) error {
// 实战中首先需要查询 Zones 获取 zone_id
// 为了简化演示修复,此处假设获取 zone 的逻辑 (伪代码实现,可进一步真实拉取)
zoneID, err := p.getZoneID(ctx, baseDomain)
if err != nil {
return err
}
for _, rec := range records {
fullDomain := rec.Name + "." + baseDomain
if rec.Name == "@" {
fullDomain = baseDomain
}
// 1. 查询现有的 Record
recordID, _ := p.getRecordID(ctx, zoneID, fullDomain, rec.Type)
// 2. 构造 payload
payload := map[string]interface{}{
"type": rec.Type,
"name": fullDomain,
"content": rec.Value,
"ttl": 1, // 自动
"proxied": false,
}
data, _ := json.Marshal(payload)
var req *http.Request
if recordID != "" {
// 更新
req, _ = http.NewRequestWithContext(ctx, "PUT", fmt.Sprintf("https://api.cloudflare.com/client/v4/zones/%s/dns_records/%s", zoneID, recordID), bytes.NewReader(data))
} else {
// 创建
req, _ = http.NewRequestWithContext(ctx, "POST", fmt.Sprintf("https://api.cloudflare.com/client/v4/zones/%s/dns_records", zoneID), bytes.NewReader(data))
}
req.Header.Set("Authorization", "Bearer "+p.token)
req.Header.Set("Content-Type", "application/json")
resp, err := p.client.Do(req)
if err != nil {
return err
}
resp.Body.Close()
}
return nil
}
// 辅助方法:获取 Zone ID
func (p *cloudflareProvider) getZoneID(ctx context.Context, domain string) (string, error) {
req, _ := http.NewRequestWithContext(ctx, "GET", "https://api.cloudflare.com/client/v4/zones?name="+domain, nil)
req.Header.Set("Authorization", "Bearer "+p.token)
resp, err := p.client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
var result struct {
Result []struct {
ID string `json:"id"`
} `json:"result"`
}
body, _ := io.ReadAll(resp.Body)
json.Unmarshal(body, &result)
if len(result.Result) > 0 {
return result.Result[0].ID, nil
}
return "", fmt.Errorf("找不到域名 %s 的 Zone", domain)
}
// 辅助方法:获取 Record ID
func (p *cloudflareProvider) getRecordID(ctx context.Context, zoneID, name, recType string) (string, error) {
req, _ := http.NewRequestWithContext(ctx, "GET", fmt.Sprintf("https://api.cloudflare.com/client/v4/zones/%s/dns_records?name=%s&type=%s", zoneID, name, recType), nil)
req.Header.Set("Authorization", "Bearer "+p.token)
resp, err := p.client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
var result struct {
Result []struct {
ID string `json:"id"`
} `json:"result"`
}
body, _ := io.ReadAll(resp.Body)
json.Unmarshal(body, &result)
if len(result.Result) > 0 {
return result.Result[0].ID, nil
}
return "", nil
}
+357
View File
@@ -0,0 +1,357 @@
package service
import (
"context"
"fmt"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/dnsprovider"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"github.com/libdns/libdns"
"go.uber.org/zap"
"gorm.io/gorm"
)
// DDNSOperationService DDNS 操作服务(全功能模式)
type DDNSOperationService struct {
logger *zap.Logger
db *gorm.DB
}
// NewDDNSOperationService 创建 DDNS 操作服务
func NewDDNSOperationService(logger *zap.Logger, db *gorm.DB) *DDNSOperationService {
return &DDNSOperationService{
logger: logger,
db: db,
}
}
// CreateDNSRecord 创建 DNS 记录(全功能模式)
func (s *DDNSOperationService) CreateDNSRecord(config *model.Service, recordType string, name string, value string, ttl int) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// 1. 获取关联的 DDNS 配置
ddnsConfig, err := s.getDDNSConfig(config.DDNSConfigID)
if err != nil {
return fmt.Errorf("获取 DDNS 配置失败:%w", err)
}
// 2. 创建 DNS Provider
providerConfig := dnsprovider.ProviderConfig{
Provider: dnsprovider.ProviderType(ddnsConfig.Provider),
Domain: ddnsConfig.Domain,
APIToken: ddnsConfig.Token,
AccessKeyID: ddnsConfig.AuthUsername,
AccessKeySecret: ddnsConfig.AuthPassword,
SecretId: ddnsConfig.AuthUsername,
SecretKey: ddnsConfig.AuthPassword,
}
provider, err := dnsprovider.NewDNSProvider(providerConfig)
if err != nil {
return fmt.Errorf("创建 DNS Provider 失败:%w", err)
}
// 3. 构建 DNS 记录
dnsRecord := &dnsprovider.DNSRecord{
Type: dnsprovider.RecordType(recordType),
Name: name,
Value: value,
TTL: ttl,
}
// 4. 添加 DNS 记录
zone := ddnsConfig.Domain
libdnsRecord := dnsRecord.ToLibdnsRecord()
s.logger.Info("开始创建 DNS 记录",
zap.String("type", recordType),
zap.String("name", name),
zap.String("value", value),
zap.String("domain", zone))
_, err = provider.AppendRecords(ctx, zone, []libdns.Record{libdnsRecord})
if err != nil {
return fmt.Errorf("添加 DNS 记录失败:%w", err)
}
s.logger.Info("DNS 记录创建成功",
zap.String("type", recordType),
zap.String("name", name),
zap.String("domain", zone))
return nil
}
// UpdateDNSRecord 更新 DNS 记录(全功能模式)
func (s *DDNSOperationService) UpdateDNSRecord(config *model.Service, recordType string, name string, value string, ttl int) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// 1. 获取关联的 DDNS 配置
ddnsConfig, err := s.getDDNSConfig(config.DDNSConfigID)
if err != nil {
return fmt.Errorf("获取 DDNS 配置失败:%w", err)
}
// 2. 创建 DNS Provider
providerConfig := dnsprovider.ProviderConfig{
Provider: dnsprovider.ProviderType(ddnsConfig.Provider),
Domain: ddnsConfig.Domain,
APIToken: ddnsConfig.Token,
AccessKeyID: ddnsConfig.AuthUsername,
AccessKeySecret: ddnsConfig.AuthPassword,
SecretId: ddnsConfig.AuthUsername,
SecretKey: ddnsConfig.AuthPassword,
}
provider, err := dnsprovider.NewDNSProvider(providerConfig)
if err != nil {
return fmt.Errorf("创建 DNS Provider 失败:%w", err)
}
// 3. 构建新的 DNS 记录
newRecord := &dnsprovider.DNSRecord{
Type: dnsprovider.RecordType(recordType),
Name: name,
Value: value,
TTL: ttl,
}
// 4. 使用 SetRecords 覆盖现有记录(会自动删除旧记录并创建新记录)
zone := ddnsConfig.Domain
libdnsRecord := newRecord.ToLibdnsRecord()
s.logger.Info("开始更新 DNS 记录",
zap.String("type", recordType),
zap.String("name", name),
zap.String("value", value),
zap.String("domain", zone))
_, err = provider.SetRecords(ctx, zone, []libdns.Record{libdnsRecord})
if err != nil {
return fmt.Errorf("更新 DNS 记录失败:%w", err)
}
s.logger.Info("DNS 记录更新成功",
zap.String("type", recordType),
zap.String("name", name),
zap.String("domain", zone))
return nil
}
// DeleteDNSRecord 删除 DNS 记录(全功能模式)
func (s *DDNSOperationService) DeleteDNSRecord(config *model.Service, recordType string, name string) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// 1. 获取关联的 DDNS 配置
ddnsConfig, err := s.getDDNSConfig(config.DDNSConfigID)
if err != nil {
return fmt.Errorf("获取 DDNS 配置失败:%w", err)
}
// 2. 创建 DNS Provider
providerConfig := dnsprovider.ProviderConfig{
Provider: dnsprovider.ProviderType(ddnsConfig.Provider),
Domain: ddnsConfig.Domain,
APIToken: ddnsConfig.Token,
AccessKeyID: ddnsConfig.AuthUsername,
AccessKeySecret: ddnsConfig.AuthPassword,
SecretId: ddnsConfig.AuthUsername,
SecretKey: ddnsConfig.AuthPassword,
}
provider, err := dnsprovider.NewDNSProvider(providerConfig)
if err != nil {
return fmt.Errorf("创建 DNS Provider 失败:%w", err)
}
// 3. 先获取所有记录
records, err := provider.GetRecords(ctx, ddnsConfig.Domain)
if err != nil {
return fmt.Errorf("获取 DNS 记录失败:%w", err)
}
// 4. 找到要删除的记录
var targetRecord *libdns.Record
for _, rec := range records {
if rec.Type == recordType && rec.Name == name {
targetRecord = &rec
break
}
}
if targetRecord == nil {
s.logger.Warn("DNS 记录不存在,跳过删除",
zap.String("type", recordType),
zap.String("name", name))
return nil
}
// 5. 删除记录
s.logger.Info("开始删除 DNS 记录",
zap.String("type", recordType),
zap.String("name", name),
zap.String("domain", ddnsConfig.Domain))
_, err = provider.DeleteRecords(ctx, ddnsConfig.Domain, []libdns.Record{*targetRecord})
if err != nil {
return fmt.Errorf("删除 DNS 记录失败:%w", err)
}
s.logger.Info("DNS 记录删除成功",
zap.String("type", recordType),
zap.String("name", name),
zap.String("domain", ddnsConfig.Domain))
return nil
}
// getDDNSConfig 获取关联的 DDNS 配置
func (s *DDNSOperationService) getDDNSConfig(configID string) (*model.Service, error) {
// 从 Service 表中查询 ID=configID 且 Type=DDNS 的记录
var ddnsService model.Service
if err := s.db.Where("id = ? AND type = 'DDNS'", configID).First(&ddnsService).Error; err != nil {
return nil, fmt.Errorf("查询 DDNS 配置失败:%w", err)
}
return &ddnsService, nil
}
// SyncMeshSeedToDNS 同步 MeshSeed 到 DNS TXT 记录(带重试机制)
func (s *DDNSOperationService) SyncMeshSeedToDNS(networkID uint64, seedString string, ddnsServiceID string) error {
const maxRetries = 3
var lastErr error
// 查询 DDNS 配置
var ddnsService model.Service
if err := s.db.First(&ddnsService, ddnsServiceID).Error; err != nil {
return fmt.Errorf("查询 DDNS 服务失败:%w", err)
}
// 查询网络获取前缀
var network model.Network
if err := s.db.First(&network, networkID).Error; err != nil {
return fmt.Errorf("查询网络失败:%w", err)
}
// 重试逻辑(指数退避)
for attempt := 1; attempt <= maxRetries; attempt++ {
lastErr = s.doSyncMeshSeedToDNS(&network, &ddnsService, seedString)
if lastErr == nil {
// 成功,更新状态
s.updateSyncStatus(networkID, "success", "")
s.logger.Info("MeshSeed 同步到 DNS 成功",
zap.Uint64("network_id", networkID),
zap.Int("attempt", attempt))
return nil
}
// 失败,记录日志
s.logger.Warn("MeshSeed 同步失败",
zap.Uint64("network_id", networkID),
zap.Int("attempt", attempt),
zap.Error(lastErr))
// 等待后重试(指数退避:1s, 2s, 4s...
if attempt < maxRetries {
waitTime := time.Duration(1<<uint(attempt-1)) * time.Second
time.Sleep(waitTime)
}
}
// 全部失败,更新状态
s.updateSyncStatus(networkID, "failed", fmt.Sprintf("重试%d次失败:%v", maxRetries, lastErr))
return fmt.Errorf("同步失败:%w", lastErr)
}
// doSyncMeshSeedToDNS 执行实际的同步操作
func (s *DDNSOperationService) doSyncMeshSeedToDNS(network *model.Network, ddnsService *model.Service, seedString string) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// 1. 创建 DNS Provider
providerConfig := dnsprovider.ProviderConfig{
Provider: dnsprovider.ProviderType(ddnsService.Provider),
Domain: ddnsService.Domain,
APIToken: ddnsService.Token,
AccessKeyID: ddnsService.AuthUsername,
AccessKeySecret: ddnsService.AuthPassword,
SecretId: ddnsService.AuthUsername,
SecretKey: ddnsService.AuthPassword,
}
provider, err := dnsprovider.NewDNSProvider(providerConfig)
if err != nil {
return fmt.Errorf("创建 DNS Provider 失败:%w", err)
}
// 2. 构建 TXT 记录名称
txtRecordName := network.DDNSPrefix
if txtRecordName == "" {
return fmt.Errorf("DDNS 前缀为空")
}
// 3. 加密 MeshSeedAES-256-GCM
encryptedSeed, err := s.encryptMeshSeed(seedString, network.ID)
if err != nil {
return fmt.Errorf("加密 MeshSeed 失败:%w", err)
}
s.logger.Info("开始同步 MeshSeed 到 DNS",
zap.String("record_name", txtRecordName),
zap.String("domain", ddnsService.Domain))
// 4. 创建/更新 TXT 记录
zone := ddnsService.Domain
libdnsRecord := libdns.Record{
Type: "TXT",
Name: txtRecordName,
Value: encryptedSeed,
TTL: time.Duration(600) * time.Second,
}
// 先尝试删除旧记录(如果存在)
_, err = provider.DeleteRecords(ctx, zone, []libdns.Record{libdnsRecord})
if err != nil {
s.logger.Debug("删除旧记录失败(可能不存在)",
zap.String("name", txtRecordName),
zap.Error(err))
}
// 添加新记录
_, err = provider.AppendRecords(ctx, zone, []libdns.Record{libdnsRecord})
if err != nil {
return fmt.Errorf("添加 DNS 记录失败:%w", err)
}
s.logger.Info("MeshSeed 同步到 DNS 成功",
zap.String("record_name", txtRecordName),
zap.String("domain", ddnsService.Domain))
return nil
}
// encryptMeshSeed 加密 MeshSeedAES-256-GCM
func (s *DDNSOperationService) encryptMeshSeed(plaintext string, networkID uint64) (string, error) {
// TODO: 实现 AES-256-GCM 加密
// 密钥派生:SHA256("meshray-ddns" + Network.Secret)
// 目前先返回明文(P3 阶段实现)
return plaintext, nil
}
// updateSyncStatus 更新同步状态
func (s *DDNSOperationService) updateSyncStatus(networkID uint64, status, message string) {
// 更新 NetworkDDNSBinding 表
var binding model.NetworkDDNSBinding
if err := s.db.Where("network_id = ?", networkID).First(&binding).Error; err == nil {
binding.Status = status
now := time.Now()
binding.LastSyncAt = &now
binding.SyncMessage = message
s.db.Save(&binding)
}
}
+18
View File
@@ -0,0 +1,18 @@
package service
import "context"
// DDNSRecord 表示一条 DNS 记录
type DDNSRecord struct {
Type string // "A", "AAAA", "TXT"
Name string // 子域名或记录名,例如 "nas" 或 "_meshray._mesh"
Value string // 记录值(IP 或 TXT内容)
}
// DDNSProvider 是所有 DNS 服务商的通用接口
type DDNSProvider interface {
// TestConnectivity 测试认证连通性
TestConnectivity(ctx context.Context) error
// SyncRecords 批量同步多条记录
SyncRecords(ctx context.Context, baseDomain string, records []DDNSRecord) error
}
+364
View File
@@ -0,0 +1,364 @@
package service
import (
"crypto/rand"
"encoding/base64"
"errors"
"fmt"
"net"
"strconv"
"git.zkcoi.com/zkcoi/meshray/internal/ctr"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"golang.org/x/crypto/curve25519"
"gorm.io/gorm"
)
// DeviceService 设备管理服务
type DeviceService struct {
store *sqlite.Store
ctrClient ctr.Client // meshray-ctr 客户端
}
// NewDeviceService 创建设备服务实例
func NewDeviceService(store *sqlite.Store, ctrClient ctr.Client) *DeviceService {
return &DeviceService{
store: store,
ctrClient: ctrClient,
}
}
// GetDevice 获取设备详情
func (s *DeviceService) GetDevice(id uint64) (*model.Device, error) {
var device model.Device
err := s.store.DB().Preload("Network").First(&device, id).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("设备不存在")
}
return nil, err
}
return &device, nil
}
// ListAllDevices 获取所有设备列表
func (s *DeviceService) ListAllDevices() ([]model.Device, error) {
var devices []model.Device
err := s.store.DB().Preload("Network").Find(&devices).Error
return devices, err
}
// ListDevicesByNetwork 获取网络下的设备列表
func (s *DeviceService) ListDevicesByNetwork(networkID uint64) ([]model.Device, error) {
var devices []model.Device
err := s.store.DB().Where("network_id = ?", networkID).Find(&devices).Error
return devices, err
}
// CreateDeviceRequest 创建设备请求
type CreateDeviceRequest struct {
NetworkID uint64 `json:"network_id"`
Name string `json:"name"`
VirtualIP string `json:"virtual_ip"` // 可选,留空则自动分配
Description string `json:"description"` // 可选
}
// CreateDeviceResult 创建设备结果
type CreateDeviceResult struct {
Device *model.Device `json:"device"`
PrivateKey string `json:"private_key"` // 仅首次返回
ConfigText string `json:"config_text"` // WireGuard 配置文本
}
// CreateDevice 创建设备
func (s *DeviceService) CreateDevice(req *CreateDeviceRequest) (*CreateDeviceResult, error) {
// 验证网络是否存在
var network model.Network
if err := s.store.DB().First(&network, req.NetworkID).Error; err != nil {
return nil, errors.New("网络不存在")
}
// 检查设备名称是否重复
var existing model.Device
if err := s.store.DB().Where("network_id = ? AND name = ?", req.NetworkID, req.Name).First(&existing).Error; err == nil {
return nil, errors.New("设备名称已存在")
}
// 生成 WireGuard 密钥对(同时获取私钥)
privateKey, publicKey, err := generateWireGuardKeys()
if err != nil {
return nil, fmt.Errorf("生成密钥失败:%w", err)
}
// 生成预共享密钥
preSharedKey := generatePreSharedKey()
// 自动分配 IP(如果未提供)
virtualIP := req.VirtualIP
if virtualIP == "" {
virtualIP, err = s.allocateIP(req.NetworkID, network.SubnetIPv4)
if err != nil {
return nil, fmt.Errorf("分配 IP 失败:%w", err)
}
}
// 创建设备
device := &model.Device{
NetworkID: req.NetworkID,
Name: req.Name,
VirtualIP: virtualIP,
PublicKey: publicKey,
PresharedKey: preSharedKey,
Status: "offline",
}
if err := s.store.DB().Create(device).Error; err != nil {
return nil, err
}
// P2 阶段 - 调用 meshray-ctr 添加 Peer
if s.ctrClient != nil {
allowedIP := device.VirtualIP + "/32"
if err := s.ctrClient.AddPeer(device.NetworkID, device.PublicKey, allowedIP); err != nil {
// 记录错误但不影响数据库操作(允许降级)
}
}
// 生成 WireGuard 配置文本
configText := s.deviceGenerateConfig(device, &network, privateKey)
return &CreateDeviceResult{
Device: device,
PrivateKey: privateKey,
ConfigText: configText,
}, nil
}
// UpdateDevice 更新设备
func (s *DeviceService) UpdateDevice(id uint64, updates map[string]interface{}) (*model.Device, error) {
var device model.Device
if err := s.store.DB().First(&device, id).Error; err != nil {
return nil, errors.New("设备不存在")
}
if err := s.store.DB().Model(&device).Updates(updates).Error; err != nil {
return nil, err
}
return &device, nil
}
// DeleteDevice 删除设备
func (s *DeviceService) DeleteDevice(id uint64) error {
// ✅ 使用事务保证数据一致性
tx := s.store.DB().Begin()
defer func() {
if r := recover(); r != nil {
tx.Rollback()
}
}()
var device model.Device
if err := tx.First(&device, id).Error; err != nil {
tx.Rollback()
return errors.New("设备不存在")
}
// ✅ 如果设备在线,先断开连接(P2 阶段)
if s.ctrClient != nil {
if err := s.ctrClient.RemovePeer(device.NetworkID, device.PublicKey); err != nil {
// 记录警告但不阻断删除流程
fmt.Printf("⚠️ 从 WireGuard 移除 Peer 失败:network_id=%d, error=%v\n", device.NetworkID, err)
}
}
// ✅ 清理相关路由和配置(P2 阶段)
// ✅ P3 阶段 - 当前无额外路由配置,暂不实现
// 未来如果需要,可以在这里添加
// 删除设备
if err := tx.Delete(&device, id).Error; err != nil {
tx.Rollback()
return err
}
// 提交事务
if err := tx.Commit().Error; err != nil {
return err
}
return nil
}
// allocateIP 自动分配 IP 地址
func (s *DeviceService) allocateIP(networkID uint64, subnet string) (string, error) {
// 解析子网
_, ipNet, err := net.ParseCIDR(subnet)
if err != nil {
return "", fmt.Errorf("解析子网失败:%w", err)
}
// 检查是否是 IPv4 地址
if ipNet.IP.To4() == nil {
return "", errors.New("暂不支持 IPv6 地址分配")
}
// 查询已分配的所有 IP
var devices []model.Device
if err := s.store.DB().Where("network_id = ?", networkID).Find(&devices).Error; err != nil {
return "", fmt.Errorf("查询已分配 IP 失败:%w", err)
}
// 构建已占用 IP 集合
occupiedIPs := make(map[string]bool)
for _, device := range devices {
occupiedIPs[device.VirtualIP] = true
}
// 从 .2 开始分配(.1 通常留给网关)
// 遍历整个子网范围查找可用 IP
ip := ipNet.IP.To4()
startIP := ip.Mask(ipNet.Mask).To4()
startIP[3]++ // 从 .1 开始
// 最多尝试 254 个 IP/24 子网)
for i := 1; i < 255; i++ {
candidateIP := make(net.IP, len(startIP))
copy(candidateIP, startIP)
candidateIP[3] = byte(i + 1) // 从 .2 开始
// 检查是否被占用
if !occupiedIPs[candidateIP.String()] {
return candidateIP.String(), nil
}
}
return "", errors.New("IP 地址已耗尽")
}
// getSettings 获取系统设置(辅助方法)
func (s *DeviceService) getSettings() (*model.SystemSetting, error) {
var setting model.SystemSetting
err := s.store.DB().First(&setting, 1).Error
if err != nil {
// 如果不存在,返回默认值
return &model.SystemSetting{
ServerPort: 51820,
}, nil
}
return &setting, nil
}
// generateWireGuardKeys 生成 WireGuard 密钥对
func generateWireGuardKeys() (privateKey, publicKey string, err error) {
// 生成私钥(32 字节随机数)
var privKeyBytes [32]byte
if _, err := rand.Read(privKeyBytes[:]); err != nil {
return "", "", err
}
// 确保私钥符合 Curve25519 要求
privKeyBytes[0] &= 248
privKeyBytes[31] &= 127
privKeyBytes[31] |= 64
// 从私钥推导公钥
var pubKeyBytes [32]byte
curve25519.ScalarBaseMult(&pubKeyBytes, &privKeyBytes)
// Base64 编码
privateKey = base64.StdEncoding.EncodeToString(privKeyBytes[:])
publicKey = base64.StdEncoding.EncodeToString(pubKeyBytes[:])
return privateKey, publicKey, nil
}
// generatePreSharedKey 生成预共享密钥
func generatePreSharedKey() string {
bytes := make([]byte, 32)
rand.Read(bytes)
return base64.StdEncoding.EncodeToString(bytes)
}
// deviceGenerateConfig 生成 WireGuard 配置文本(设备创建时使用)
func (s *DeviceService) deviceGenerateConfig(device *model.Device, network *model.Network, privateKey string) string {
settings, _ := s.getSettings()
config := "[Interface]\n"
config += "PrivateKey = " + privateKey + "\n"
config += "Address = " + device.VirtualIP + "/32\n"
config += fmt.Sprintf("MTU = %d\n\n", network.MTU)
config += "[Peer]\n"
if settings.ServerPublicKey != "" {
config += "PublicKey = " + settings.ServerPublicKey + "\n"
} else {
config += "PublicKey = <SERVER_PUBLIC_KEY>\n"
}
if settings.ServerIP != "" {
config += "Endpoint = " + settings.ServerIP + ":" + strconv.Itoa(settings.ServerPort) + "\n"
} else {
config += fmt.Sprintf("Endpoint = <SERVER_IP>:%d\n", settings.ServerPort)
}
config += "AllowedIPs = " + network.SubnetIPv4 + "\n"
config += "PersistentKeepalive = 25\n"
return config
}
// GenerateDeviceConfig 生成设备配置文件
func (s *DeviceService) GenerateDeviceConfig(deviceID uint64) (string, error) {
// 获取设备信息(包含网络)
device, err := s.GetDevice(deviceID)
if err != nil {
return "", err
}
// 生成 WireGuard 密钥对
privateKey, publicKey, err := generateWireGuardKeys()
if err != nil {
return "", fmt.Errorf("生成密钥失败:%w", err)
}
// 更新设备的公钥到数据库
if err := s.store.DB().Model(device).Update("public_key", publicKey).Error; err != nil {
return "", fmt.Errorf("保存公钥失败:%w", err)
}
// 生成完整的 WireGuard 配置
config := "# MeshRay Generated Configuration\n"
config += "# Device: " + device.Name + "\n"
config += "# Created: " + device.CreatedAt.Format("2006-01-02 15:04:05") + "\n\n"
// [Interface] 部分
config += "[Interface]\n"
config += "PrivateKey = " + privateKey + "\n"
config += "Address = " + device.VirtualIP + "/32\n"
config += "DNS = 8.8.8.8, 8.8.4.4\n\n"
// [Peer] 部分(服务端配置)
config += "# Server (MeshRay)\n"
config += "[Peer]\n"
// 从 Settings 读取服务端公钥
settings, _ := s.getSettings()
if settings.ServerPublicKey == "" {
return "", errors.New("请先在系统设置中配置服务端公钥")
}
config += "PublicKey = " + settings.ServerPublicKey + "\n"
if device.PresharedKey != "" {
config += "PresharedKey = " + device.PresharedKey + "\n"
}
config += "AllowedIPs = 0.0.0.0/0\n"
// 使用 Settings 中的 ServerIP 和 ServerPort
if settings.ServerIP == "" {
return "", errors.New("请先在系统设置中配置服务端 IP 地址")
}
config += "Endpoint = " + settings.ServerIP + ":" + strconv.Itoa(settings.ServerPort) + "\n"
config += "PersistentKeepalive = 25\n"
return config, nil
}
+164
View File
@@ -0,0 +1,164 @@
package service
import (
"context"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
)
// IPDetectionService IP 检测服务
type IPDetectionService struct {
httpClient *http.Client
}
// NewIPDetectionService 创建 IP 检测服务
func NewIPDetectionService() *IPDetectionService {
return &IPDetectionService{
httpClient: &http.Client{
Timeout: 10 * time.Second,
},
}
}
// GetPublicIPv4 获取公网 IPv4 地址
func (s *IPDetectionService) GetPublicIPv4() (string, error) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
req, err := http.NewRequestWithContext(ctx, "GET", "https://api.ipify.org?format=json", nil)
if err != nil {
return "", fmt.Errorf("创建请求失败:%w", err)
}
resp, err := s.httpClient.Do(req)
if err != nil {
return "", fmt.Errorf("请求失败:%w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return "", fmt.Errorf("读取响应失败:%w", err)
}
var result struct {
IP string `json:"ip"`
}
if err := json.Unmarshal(body, &result); err != nil {
return "", fmt.Errorf("解析 JSON 失败:%w", err)
}
if result.IP == "" {
return "", fmt.Errorf("未获取到 IPv4 地址")
}
return result.IP, nil
}
// GetPublicIPv6 获取公网 IPv6 地址
func (s *IPDetectionService) GetPublicIPv6() (string, error) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
req, err := http.NewRequestWithContext(ctx, "GET", "https://api64.ipify.org?format=json", nil)
if err != nil {
return "", fmt.Errorf("创建请求失败:%w", err)
}
resp, err := s.httpClient.Do(req)
if err != nil {
return "", fmt.Errorf("请求失败:%w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return "", fmt.Errorf("读取响应失败:%w", err)
}
var result struct {
IP string `json:"ip"`
}
if err := json.Unmarshal(body, &result); err != nil {
return "", fmt.Errorf("解析 JSON 失败:%w", err)
}
if result.IP == "" {
return "", fmt.Errorf("未获取到 IPv6 地址")
}
// 检查是否是 IPv6 地址(包含冒号)
if !strings.Contains(result.IP, ":") {
return "", fmt.Errorf("获取到的不是有效的 IPv6 地址:%s", result.IP)
}
return result.IP, nil
}
// GetLocalIPv4 获取本地 IPv4 地址(第一个非回环接口)
func (s *IPDetectionService) GetLocalIPv4() (string, error) {
addrs, err := net.InterfaceAddrs()
if err != nil {
return "", fmt.Errorf("获取网络接口失败:%w", err)
}
for _, addr := range addrs {
// 检查是否为 IP 地址且为 IPv4
if ipNet, ok := addr.(*net.IPNet); ok && !ipNet.IP.IsLoopback() {
if ipNet.IP.To4() != nil {
return ipNet.IP.String(), nil
}
}
}
return "", fmt.Errorf("未找到 IPv4 地址")
}
// GetLocalIPv6 获取本地 IPv6 地址(第一个非回环接口)
func (s *IPDetectionService) GetLocalIPv6() (string, error) {
addrs, err := net.InterfaceAddrs()
if err != nil {
return "", fmt.Errorf("获取网络接口失败:%w", err)
}
for _, addr := range addrs {
if ipNet, ok := addr.(*net.IPNet); ok && !ipNet.IP.IsLoopback() {
if ipNet.IP.To4() == nil && ipNet.IP.To16() != nil {
return ipNet.IP.String(), nil
}
}
}
return "", fmt.Errorf("未找到 IPv6 地址")
}
// DetectIP 检测 IP 地址(根据记录类型返回对应的 IP)
func (s *IPDetectionService) DetectIP(recordType string) (string, error) {
switch recordType {
case "A":
// A 记录优先使用公网 IPv4
ip, err := s.GetPublicIPv4()
if err != nil {
// 降级到本地 IPv4
return s.GetLocalIPv4()
}
return ip, nil
case "AAAA":
// AAAA 记录优先使用公网 IPv6
ip, err := s.GetPublicIPv6()
if err != nil {
// 降级到本地 IPv6
return s.GetLocalIPv6()
}
return ip, nil
default:
return "", fmt.Errorf("不支持的记录类型:%s", recordType)
}
}
+52
View File
@@ -0,0 +1,52 @@
package service
import (
"context"
"io"
"net/http"
"strings"
"time"
)
// IPFetcher 用于获取本机的公网 IPv4 和 IPv6
type IPFetcher struct {
client *http.Client
}
func NewIPFetcher() *IPFetcher {
return &IPFetcher{
client: &http.Client{Timeout: 5 * time.Second},
}
}
// GetIPv4 获取公网 IPv4
func (f *IPFetcher) GetIPv4(ctx context.Context) (string, error) {
return f.fetchIP(ctx, "https://api.ipify.org")
}
// GetIPv6 获取公网 IPv6
func (f *IPFetcher) GetIPv6(ctx context.Context) (string, error) {
return f.fetchIP(ctx, "https://api6.ipify.org")
}
func (f *IPFetcher) fetchIP(ctx context.Context, apiURL string) (string, error) {
req, err := http.NewRequestWithContext(ctx, "GET", apiURL, nil)
if err != nil {
return "", err
}
resp, err := f.client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", http.ErrServerClosed // 简易错误
}
ipBytes, err := io.ReadAll(resp.Body)
if err != nil {
return "", err
}
return strings.TrimSpace(string(ipBytes)), nil
}
+203
View File
@@ -0,0 +1,203 @@
package service
import (
"crypto/rand"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/model"
sqlite "git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"go.uber.org/zap"
"golang.org/x/crypto/ed25519"
"gorm.io/gorm"
)
// MeshSeedService MeshSeed 服务
type MeshSeedService struct {
store *sqlite.Store
logger *zap.Logger
signingKey ed25519.PrivateKey // Ed25519 签名密钥
issuerNodeID string // 签发节点 ID
}
// NewMeshSeedService 创建 MeshSeed 服务
func NewMeshSeedService(store *sqlite.Store, logger *zap.Logger, signingKey ed25519.PrivateKey, issuerNodeID string) *MeshSeedService {
return &MeshSeedService{
store: store,
logger: logger,
signingKey: signingKey,
issuerNodeID: issuerNodeID,
}
}
// GenerateMeshSeed 生成 MeshSeed
func (s *MeshSeedService) GenerateMeshSeed(networkID uint64, maxUses int, expiresAt time.Time, ddnsEnabled bool) (*model.MeshSeed, error) {
// 1. 验证网络是否存在
var network model.Network
if err := s.store.DB().First(&network, networkID).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, fmt.Errorf("网络不存在")
}
return nil, fmt.Errorf("查询网络失败:%w", err)
}
// 2. 生成随机 SeedID16 字节随机数)
seedBytes := make([]byte, 16)
if _, err := rand.Read(seedBytes); err != nil {
return nil, fmt.Errorf("生成随机数失败:%w", err)
}
seedID := base64.RawURLEncoding.EncodeToString(seedBytes)
// 3. 构建 JoinToken(包含网络信息)
joinTokenData := map[string]interface{}{
"seed_id": seedID,
"network_id": networkID,
"network_name": network.Name,
"subnet_ipv4": network.SubnetIPv4,
"mode": network.Mode,
"ddns_enabled": ddnsEnabled,
"expires_at": expiresAt.Unix(),
"max_uses": maxUses,
}
// 序列化为 JSON
tokenJSON, err := json.Marshal(joinTokenData)
if err != nil {
return nil, fmt.Errorf("序列化 Token 失败:%w", err)
}
// Base64 编码
joinToken := base64.StdEncoding.EncodeToString(tokenJSON)
// 4. Ed25519 签名
signature := ed25519.Sign(s.signingKey, []byte(joinToken))
signatureStr := base64.StdEncoding.EncodeToString(signature)
// 5. 创建 MeshSeed 记录
meshSeed := &model.MeshSeed{
SeedID: seedID,
NetworkID: networkID,
JoinToken: joinToken,
Signature: signatureStr,
IssuerNodeID: s.issuerNodeID,
MaxUses: maxUses,
UsedCount: 0,
ExpiresAt: expiresAt,
DDNSEnabled: ddnsEnabled,
UpdateVersion: 0,
Revoked: false,
}
if err := s.store.DB().Create(meshSeed).Error; err != nil {
return nil, fmt.Errorf("创建 MeshSeed 失败:%w", err)
}
s.logger.Info("MeshSeed 已生成",
zap.String("seed_id", seedID),
zap.Uint64("network_id", networkID),
zap.Int("max_uses", maxUses),
zap.Time("expires_at", expiresAt))
return meshSeed, nil
}
// VerifyMeshSeed 验证 MeshSeed
func (s *MeshSeedService) VerifyMeshSeed(joinToken, signature string) (*model.MeshSeed, error) {
// 1. 解码 JoinToken
tokenBytes, err := base64.StdEncoding.DecodeString(joinToken)
if err != nil {
return nil, fmt.Errorf("解码 Token 失败:%w", err)
}
// 2. 解码签名
sigBytes, err := base64.StdEncoding.DecodeString(signature)
if err != nil {
return nil, fmt.Errorf("解码签名失败:%w", err)
}
// 3. 验证 Ed25519 签名
publicKey := s.signingKey.Public()
if !ed25519.Verify(publicKey.(ed25519.PublicKey), tokenBytes, sigBytes) {
return nil, fmt.Errorf("签名验证失败")
}
// 4. 解析 Token 内容
var tokenData map[string]interface{}
if err := json.Unmarshal(tokenBytes, &tokenData); err != nil {
return nil, fmt.Errorf("解析 Token 失败:%w", err)
}
seedID, ok := tokenData["seed_id"].(string)
if !ok {
return nil, fmt.Errorf("Token 格式错误")
}
// 5. 查询 MeshSeed
var meshSeed model.MeshSeed
if err := s.store.DB().Where("seed_id = ?", seedID).First(&meshSeed).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, fmt.Errorf("MeshSeed 不存在")
}
return nil, fmt.Errorf("查询 MeshSeed 失败:%w", err)
}
// 6. 检查是否被吊销
if meshSeed.Revoked {
return nil, fmt.Errorf("MeshSeed 已被吊销")
}
// 7. 检查使用次数
if meshSeed.UsedCount >= meshSeed.MaxUses {
return nil, fmt.Errorf("MeshSeed 使用次数已用尽")
}
// 8. 检查过期时间
if time.Now().After(meshSeed.ExpiresAt) {
return nil, fmt.Errorf("MeshSeed 已过期")
}
return &meshSeed, nil
}
// IncrementUseCount 增加使用次数
func (s *MeshSeedService) IncrementUseCount(seedID string) error {
return s.store.DB().Transaction(func(tx *gorm.DB) error {
var meshSeed model.MeshSeed
if err := tx.Where("seed_id = ?", seedID).First(&meshSeed).Error; err != nil {
return err
}
return tx.Model(&meshSeed).UpdateColumn("used_count", meshSeed.UsedCount+1).Error
})
}
// RevokeMeshSeed 吊销 MeshSeed
func (s *MeshSeedService) RevokeMeshSeed(seedID string) error {
result := s.store.DB().Model(&model.MeshSeed{}).
Where("seed_id = ?", seedID).
Update("revoked", true)
if result.Error != nil {
return fmt.Errorf("吊销 MeshSeed 失败:%w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("MeshSeed 不存在")
}
s.logger.Info("MeshSeed 已吊销", zap.String("seed_id", seedID))
return nil
}
// ListMeshSeeds 获取网络的 MeshSeed 列表
func (s *MeshSeedService) ListMeshSeeds(networkID uint) ([]model.MeshSeed, error) {
var seeds []model.MeshSeed
err := s.store.DB().Where("network_id = ? AND revoked = ?", networkID, false).
Order("created_at DESC").
Find(&seeds).Error
return seeds, err
}
+318
View File
@@ -0,0 +1,318 @@
package service
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"net"
"git.zkcoi.com/zkcoi/meshray/internal/ctr"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"git.zkcoi.com/zkcoi/meshray/pkg/idutil"
"go.uber.org/zap"
"gorm.io/gorm"
)
// NetworkService 网络管理服务
type NetworkService struct {
store *sqlite.Store
ctrClient ctr.Client // meshray-ctr 客户端
logger *zap.Logger
}
// NewNetworkService 创建网络服务实例
func NewNetworkService(store *sqlite.Store, ctrClient ctr.Client, logger *zap.Logger) *NetworkService {
return &NetworkService{
store: store,
ctrClient: ctrClient,
logger: logger,
}
}
// ListNetworks 获取网络列表(Service 层方法)
func (s *NetworkService) ListNetworks() ([]model.Network, error) {
var networks []model.Network
err := s.store.DB().Preload("Devices").Find(&networks).Error
return networks, err
}
// GetNetworkByID 根据 ID 获取网络
func (s *NetworkService) GetNetworkByID(id uint64) (*model.Network, error) {
var network model.Network
err := s.store.DB().Preload("Devices").First(&network, id).Error
if err != nil {
return nil, err
}
return &network, nil
}
// CreateNetwork 创建网络
func (s *NetworkService) CreateNetwork(req *model.Network) (*model.Network, error) {
// 验证子网格式
if err := s.validateSubnet(req.SubnetIPv4); err != nil {
return nil, err
}
// 检查网络名称是否重复
var existing model.Network
if err := s.store.DB().Where("name = ?", req.Name).First(&existing).Error; err == nil {
return nil, errors.New("网络名称已存在")
}
// 生成雪花算法 ID
if req.ID == 0 {
gen, err := idutil.GetGenerator()
if err != nil {
return nil, fmt.Errorf("failed to get snowflake generator: %w", err)
}
req.ID, err = gen.NextID()
if err != nil {
return nil, fmt.Errorf("failed to generate network id: %w", err)
}
}
// 计算 listenPort51820 + hash(networkID)(先不保存,用于 ctr 调用)
listenPort := 51820 + hashUint64(req.ID)%1000
// 根据组网模式决定是否启动 Core
// 原生模式:仅创建 WG 设备
// 增强模式:创建 WG 设备 + 启动 Core 实例
if s.ctrClient != nil {
// P2 阶段 - 先调用 meshray-ctr 创建 WireGuard 设备(失败则不回滚数据库)
// 策略:宽松模式 - ctr 失败只记录警告,不影响数据库操作
if err := s.ctrClient.CreateNetwork(req.ID, req.SubnetIPv4, listenPort, req.Mode); err != nil {
// 记录错误但不影响数据库操作(允许降级)
s.logger.Warn("调用 ctr 创建网络失败,将手动启动",
zap.Uint64("network_id", req.ID),
zap.Error(err))
}
}
// 保存到数据库
if err := s.store.DB().Create(req).Error; err != nil {
return nil, err
}
s.logger.Info("网络创建成功",
zap.Uint64("network_id", req.ID),
zap.String("mode", req.Mode))
return req, nil
}
// GetNetwork 获取网络详情
func (s *NetworkService) GetNetwork(id uint64) (*model.Network, error) {
var network model.Network
if err := s.store.DB().First(&network, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("网络不存在")
}
return nil, err
}
return &network, nil
}
// UpdateNetwork 更新网络
func (s *NetworkService) UpdateNetwork(id uint64, updates map[string]interface{}) (*model.Network, error) {
var network model.Network
if err := s.store.DB().First(&network, id).Error; err != nil {
return nil, errors.New("网络不存在")
}
if err := s.store.DB().Model(&network).Updates(updates).Error; err != nil {
return nil, err
}
return &network, nil
}
// DeleteNetwork 删除网络
func (s *NetworkService) DeleteNetwork(id uint64, force bool) error {
// ✅ 使用事务保证数据一致性
tx := s.store.DB().Begin()
defer func() {
if r := recover(); r != nil {
tx.Rollback()
}
}()
var network model.Network
if err := tx.First(&network, id).Error; err != nil {
tx.Rollback()
return errors.New("网络不存在")
}
// 检查是否有关联的设备
var deviceCount int64
tx.Model(&model.Device{}).Where("network_id = ?", id).Count(&deviceCount)
if deviceCount > 0 {
if !force {
tx.Rollback()
return errors.New("该网络下仍有设备,为避免误操作,请确认后强制删除")
}
// 级联删除所有关联设备
if err := tx.Where("network_id = ?", id).Delete(&model.Device{}).Error; err != nil {
tx.Rollback()
return fmt.Errorf("级联删除设备失败:%w", err)
}
s.logger.Info("级联删除了关联设备", zap.Uint64("network_id", id), zap.Int64("device_count", deviceCount))
}
// 先调用 ctr 删除 WG 设备(如果存在)
if s.ctrClient != nil {
if err := s.ctrClient.DeleteNetwork(id); err != nil {
// 记录错误但不中断删除流程(设备可能已经不存在)
s.logger.Warn("删除 WG 设备失败(可能已不存在)",
zap.Uint64("network_id", id),
zap.Error(err))
// 继续删除数据库记录
}
}
// 再删除数据库记录
if err := tx.Delete(&network, id).Error; err != nil {
tx.Rollback()
return err
}
// 提交事务
if err := tx.Commit().Error; err != nil {
return err
}
return nil
}
// validateSubnet 验证子网格式
func (s *NetworkService) validateSubnet(subnet string) error {
// 只调用一次 ParseCIDR
ip, ipNet, err := net.ParseCIDR(subnet)
if err != nil {
return fmt.Errorf("无效的子网格式:%s", subnet)
}
// 检查是否是有效的 WireGuard 子网(至少 /24 或更小)
ones, bits := ipNet.Mask.Size()
if bits == 32 && ones > 24 {
return errors.New("IPv4 子网掩码不能大于 /24")
}
if bits == 128 && ones > 64 {
return errors.New("IPv6 子网掩码不能大于 /64")
}
// 检查是否是私有地址段
if !isPrivateSubnet(ip, ipNet) {
return errors.New("请使用私有地址段(如 10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16")
}
return nil
}
// isPrivateSubnet 检查是否是私有地址段
func isPrivateSubnet(ip net.IP, ipNet *net.IPNet) bool {
// IPv4 私有地址段
privateRanges := []string{
"10.0.0.0/8",
"172.16.0.0/12",
"192.168.0.0/16",
"100.64.0.0/10", // CGNAT
}
for _, privRange := range privateRanges {
_, privNet, _ := net.ParseCIDR(privRange)
if privNet.Contains(ip) {
return true
}
}
return false
}
// generateNetworkSecret 生成 Network Secret
func generateNetworkSecret() string {
// 使用加密安全的随机数生成器
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
// ✅ 随机数生成失败时 panic,而不是返回弱密码
panic(fmt.Sprintf("生成安全随机数失败:%v", err))
}
return hex.EncodeToString(bytes)
}
// hashUint64 计算 uint64 哈希值
func hashUint64(id uint64) int {
h := sha256.Sum256([]byte(fmt.Sprintf("%d", id)))
// 取前 4 字节转换为 int
return int(h[0])<<24 | int(h[1])<<16 | int(h[2])<<8 | int(h[3])
}
// StartNetwork 启动网络
func (s *NetworkService) StartNetwork(id uint64) error {
network, err := s.GetNetworkByID(id)
if err != nil {
return err
}
// 如果已经有 ctr,直接调用 CreateNetwork
if s.ctrClient != nil {
listenPort := 51820 + hashUint64(network.ID)%1000
if err := s.ctrClient.CreateNetwork(network.ID, network.SubnetIPv4, listenPort, network.Mode); err != nil {
s.logger.Error("启动网络失败",
zap.Uint64("network_id", network.ID),
zap.Error(err))
return fmt.Errorf("启动网络失败:%w", err)
}
s.logger.Info("网络已启动",
zap.Uint64("network_id", network.ID),
zap.Int("listen_port", listenPort))
}
return nil
}
// StopNetwork 停止网络
func (s *NetworkService) StopNetwork(id uint64) error {
network, err := s.GetNetworkByID(id)
if err != nil {
return err
}
// 调用 ctr 删除网络(清理 WG 设备)
if s.ctrClient != nil {
if err := s.ctrClient.DeleteNetwork(network.ID); err != nil {
s.logger.Error("停止网络失败",
zap.Uint64("network_id", network.ID),
zap.Error(err))
return fmt.Errorf("停止网络失败:%w", err)
}
s.logger.Info("网络已停止",
zap.Uint64("network_id", network.ID))
}
return nil
}
// SwitchMode 切换组网模式
func (s *NetworkService) SwitchMode(id uint64, meshMode string) error {
network, err := s.GetNetworkByID(id)
if err != nil {
return err
}
// 更新数据库
if err := s.store.DB().Model(network).Update("mode", meshMode).Error; err != nil {
return err
}
s.logger.Info("组网模式已切换",
zap.Uint64("network_id", network.ID),
zap.String("mode", meshMode))
return nil
}
+243
View File
@@ -0,0 +1,243 @@
package service
import (
"encoding/json"
"sync"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"gorm.io/gorm"
)
// NotificationService 通知服务
type NotificationService struct {
db *gorm.DB
logger *zap.Logger
clients map[uint]*NotificationClient // userID -> client
mu sync.RWMutex
broadcastCh chan NotificationMessage
}
// NotificationClient WebSocket 通知客户端
type NotificationClient struct {
userID uint
username string
conn *gin.Context
msgCh chan NotificationMessage
done chan struct{}
}
// NotificationMessage 通知消息
type NotificationMessage struct {
Type string `json:"type"` // alert, system, update, ddns
Priority int `json:"priority"` // 1=low, 2=medium, 3=high
Title string `json:"title"`
Message string `json:"message"`
Data map[string]interface{} `json:"data,omitempty"`
Timestamp time.Time `json:"timestamp"`
}
// NewNotificationService 创建通知服务
func NewNotificationService(db *gorm.DB, logger *zap.Logger) *NotificationService {
svc := &NotificationService{
db: db,
logger: logger,
clients: make(map[uint]*NotificationClient),
broadcastCh: make(chan NotificationMessage, 100),
}
// 启动广播协程
go svc.runBroadcaster()
return svc
}
// runBroadcaster 运行广播协程
func (s *NotificationService) runBroadcaster() {
for msg := range s.broadcastCh {
s.mu.RLock()
for _, client := range s.clients {
select {
case client.msgCh <- msg:
// 发送成功
default:
// 通道已满,跳过
s.logger.Warn("通知通道已满", zap.Uint("user_id", client.userID))
}
}
s.mu.RUnlock()
}
}
// RegisterClient 注册通知客户端
func (s *NotificationService) RegisterClient(userID uint, username string, msgCh chan NotificationMessage, done chan struct{}) {
s.mu.Lock()
defer s.mu.Unlock()
client := &NotificationClient{
userID: userID,
username: username,
msgCh: msgCh,
done: done,
}
s.clients[userID] = client
s.logger.Info("用户已连接通知服务", zap.Uint("user_id", userID), zap.String("username", username))
}
// UnregisterClient 注销通知客户端
func (s *NotificationService) UnregisterClient(userID uint) {
s.mu.Lock()
defer s.mu.Unlock()
if client, ok := s.clients[userID]; ok {
close(client.done)
delete(s.clients, userID)
s.logger.Info("用户已断开通知服务", zap.Uint("user_id", userID))
}
}
// SendToUser 发送通知给指定用户(并保存到数据库)
func (s *NotificationService) SendToUser(userID uint, msg NotificationMessage) {
// 1. 保存到数据库
notification := model.Notification{
UserID: userID,
Type: msg.Type,
Priority: msg.Priority,
Title: msg.Title,
Message: msg.Message,
}
if msg.Data != nil {
dataJSON, _ := json.Marshal(msg.Data)
notification.Data = string(dataJSON)
}
if err := s.db.Create(&notification).Error; err != nil {
s.logger.Error("保存通知失败", zap.Error(err))
}
// 2. 发送到 WebSocket 通道
s.mu.RLock()
defer s.mu.RUnlock()
if client, ok := s.clients[userID]; ok {
select {
case client.msgCh <- msg:
s.logger.Debug("通知已发送给用户",
zap.Uint("user_id", userID),
zap.String("type", msg.Type))
default:
s.logger.Warn("用户通知通道已满", zap.Uint("user_id", userID))
}
}
}
// Broadcast 广播通知给所有在线用户(并保存到数据库)
func (s *NotificationService) Broadcast(msg NotificationMessage) {
msg.Timestamp = time.Now()
// 保存到所有用户的数据库记录
s.mu.RLock()
for userID := range s.clients {
notification := model.Notification{
UserID: userID,
Type: msg.Type,
Priority: msg.Priority,
Title: msg.Title,
Message: msg.Message,
}
if msg.Data != nil {
dataJSON, _ := json.Marshal(msg.Data)
notification.Data = string(dataJSON)
}
s.db.Create(&notification)
}
s.mu.RUnlock()
// 发送到 WebSocket 通道
s.broadcastCh <- msg
s.logger.Debug("通知已广播",
zap.String("type", msg.Type),
zap.Int("online_users", len(s.clients)))
}
// SendAlert 发送告警通知
func (s *NotificationService) SendAlert(userID uint, title, message string, data map[string]interface{}) {
msg := NotificationMessage{
Type: "alert",
Priority: 3, // high priority
Title: title,
Message: message,
Data: data,
Timestamp: time.Now(),
}
s.SendToUser(userID, msg)
}
// SendSystemNotification 发送系统通知
func (s *NotificationService) SendSystemNotification(userID uint, title, message string) {
msg := NotificationMessage{
Type: "system",
Priority: 2, // medium priority
Title: title,
Message: message,
Timestamp: time.Now(),
}
s.SendToUser(userID, msg)
}
// SendUpdateAvailable 发送更新可用通知
func (s *NotificationService) SendUpdateAvailable(version, notes, downloadURL string) {
msg := NotificationMessage{
Type: "update",
Priority: 2,
Title: "发现新版本",
Message: version,
Data: map[string]interface{}{
"version": version,
"notes": notes,
"download_url": downloadURL,
},
Timestamp: time.Now(),
}
s.Broadcast(msg)
}
// SendDDNSUpdate 发送 DDNS 更新通知
func (s *NotificationService) SendDDNSUpdate(serviceName, oldIP, newIP string) {
data, _ := json.Marshal(gin.H{
"service_name": serviceName,
"old_ip": oldIP,
"new_ip": newIP,
})
var dataMap map[string]interface{}
json.Unmarshal(data, &dataMap)
msg := NotificationMessage{
Type: "ddns",
Priority: 1, // low priority
Title: "DDNS IP 已更新",
Message: serviceName,
Data: dataMap,
Timestamp: time.Now(),
}
s.Broadcast(msg)
}
// GetOnlineUserCount 获取在线用户数
func (s *NotificationService) GetOnlineUserCount() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.clients)
}
// GetDB 返回数据库实例(用于 Handler 层查询)
func (s *NotificationService) GetDB() *gorm.DB {
return s.db
}
+249
View File
@@ -0,0 +1,249 @@
package service
import (
"errors"
"fmt"
"net"
"strings"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"gorm.io/gorm"
)
// PendingJoinService 待审核服务
type PendingJoinService struct {
store *sqlite.Store
}
// NewPendingJoinService 创建待审核服务实例
func NewPendingJoinService(store *sqlite.Store) *PendingJoinService {
return &PendingJoinService{store: store}
}
// ListPendingJoins 获取待审核列表
func (s *PendingJoinService) ListPendingJoins(networkID uint, status string, page, size int) ([]model.PendingJoin, int64, error) {
query := s.store.DB().Model(&model.PendingJoin{})
// 按网络 ID 筛选
if networkID > 0 {
query = query.Where("network_id = ?", networkID)
}
// 按状态筛选
if status != "" && status != "all" {
query = query.Where("status = ?", status)
}
// 统计总数
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
// 分页查询
var result []model.PendingJoin
offset := (page - 1) * size
err := query.Order("created_at DESC").Offset(offset).Limit(size).Find(&result).Error
if err != nil {
return nil, 0, err
}
return result, total, nil
}
// GetPendingJoinByID 根据 ID 获取待审核记录
func (s *PendingJoinService) GetPendingJoinByID(id uint) (*model.PendingJoin, error) {
var record model.PendingJoin
err := s.store.DB().First(&record, id).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("记录不存在")
}
return nil, err
}
return &record, nil
}
// ApproveResult 审核通过结果
type ApproveResult struct {
Device *model.Device `json:"device"`
PrivateKey string `json:"private_key"` // 仅首次返回
Network *model.Network `json:"network"`
ConfigText string `json:"config_text"` // WireGuard 配置文本
}
// ApproveJoin 审核通过
func (s *PendingJoinService) ApproveJoin(id uint) (*ApproveResult, error) {
var record model.PendingJoin
if err := s.store.DB().First(&record, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("记录不存在")
}
return nil, err
}
if record.Status != "pending" {
return nil, errors.New("该申请已处理")
}
// 1. 查询 MeshSeed 获取网络信息
var meshSeed model.MeshSeed
if err := s.store.DB().Where("seed_id = ?", record.SeedID).First(&meshSeed).Error; err != nil {
return nil, fmt.Errorf("查询 MeshSeed 失败:%w", err)
}
// 2. 查询网络详情
var network model.Network
if err := s.store.DB().First(&network, meshSeed.NetworkID).Error; err != nil {
return nil, fmt.Errorf("查询网络失败:%w", err)
}
// 3. 生成 WireGuard 密钥对(调用全局函数,定义在 device.go 中)
privateKey, publicKey, err := generateWireGuardKeys()
if err != nil {
return nil, fmt.Errorf("生成密钥失败:%w", err)
}
// 4. 分配 IP 地址
ipAddress, err := s.allocateIPAddress(&network)
if err != nil {
return nil, fmt.Errorf("分配 IP 失败:%w", err)
}
// 5. 创建设备记录
device := &model.Device{
Name: record.DeviceName,
NetworkID: network.ID,
PublicKey: publicKey,
VirtualIP: ipAddress,
Status: "active",
LastSeen: time.Now(),
}
if err := s.store.DB().Create(device).Error; err != nil {
return nil, fmt.Errorf("创建设备失败:%w", err)
}
// 6. 更新审核状态
now := time.Now()
record.Status = "approved"
record.ApprovedAt = &now
if err := s.store.DB().Save(&record).Error; err != nil {
return nil, fmt.Errorf("更新审核状态失败:%w", err)
}
// 7. 生成 WireGuard 配置文本(使用 pending join 专用方法,读取 SystemSetting
configText := s.pendingJoinGenerateConfig(device, &network, privateKey)
return &ApproveResult{
Device: device,
PrivateKey: privateKey,
Network: &network,
ConfigText: configText,
}, nil
}
// RejectJoin 审核拒绝
func (s *PendingJoinService) RejectJoin(id uint, reason string) error {
var record model.PendingJoin
if err := s.store.DB().First(&record, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New("记录不存在")
}
return err
}
if record.Status != "pending" {
return errors.New("该申请已处理")
}
now := time.Now()
record.Status = "rejected"
record.RejectedAt = &now
record.Reason = reason
return s.store.DB().Save(&record).Error
}
// allocateIPAddress 分配 IP 地址
func (s *PendingJoinService) allocateIPAddress(network *model.Network) (string, error) {
// 解析子网
_, ipNet, err := net.ParseCIDR(network.SubnetIPv4)
if err != nil {
return "", fmt.Errorf("解析子网失败:%w", err)
}
// 获取已使用的 IP
var devices []model.Device
if err := s.store.DB().Where("network_id = ?", network.ID).Find(&devices).Error; err != nil {
return "", fmt.Errorf("查询设备列表失败:%w", err)
}
usedIPs := make(map[string]bool)
for _, device := range devices {
usedIPs[device.VirtualIP] = true
}
// 从 .2 开始分配(.1 通常是网关)
ip := ipNet.IP.To4()
if ip == nil {
return "", errors.New("暂不支持 IPv6 地址分配")
}
for i := 2; i < 254; i++ {
candidateIP := make(net.IP, len(ip))
copy(candidateIP, ip)
candidateIP[3] = byte(i)
if !usedIPs[candidateIP.String()] {
return candidateIP.String(), nil
}
}
return "", errors.New("IP 地址已用尽")
}
// pendingJoinGenerateConfig 生成 WireGuard 配置文本(从 SystemSetting 获取服务端信息)
func (s *PendingJoinService) pendingJoinGenerateConfig(device *model.Device, network *model.Network, privateKey string) string {
// 从 SystemSetting 获取服务端信息
var settings model.SystemSetting
if err := s.store.DB().First(&settings, 1).Error; err != nil {
// 使用默认值
settings.ServerPort = 51820
}
var sb strings.Builder
sb.WriteString("[Interface]\n")
sb.WriteString(fmt.Sprintf("PrivateKey = %s\n", privateKey))
sb.WriteString(fmt.Sprintf("Address = %s/32\n", device.VirtualIP))
sb.WriteString(fmt.Sprintf("MTU = %d\n\n", network.MTU))
sb.WriteString("[Peer]\n")
if settings.ServerPublicKey != "" {
sb.WriteString(fmt.Sprintf("PublicKey = %s\n", settings.ServerPublicKey))
} else {
sb.WriteString("PublicKey = <SERVER_PUBLIC_KEY>\n")
}
if settings.ServerIP != "" {
sb.WriteString(fmt.Sprintf("Endpoint = %s:%d\n", settings.ServerIP, settings.ServerPort))
} else {
sb.WriteString(fmt.Sprintf("Endpoint = <SERVER_IP>:%d\n", settings.ServerPort))
}
sb.WriteString(fmt.Sprintf("AllowedIPs = %s\n", network.SubnetIPv4))
sb.WriteString("PersistentKeepalive = 25\n")
return sb.String()
}
// DeleteExpired 删除过期的待审核记录
func (s *PendingJoinService) DeleteExpired() error {
return s.store.DB().Where("expire_at < ? AND status = ?", time.Now(), "pending").Delete(&model.PendingJoin{}).Error
}
// CountPending 统计待审核数量
func (s *PendingJoinService) CountPending() (int64, error) {
var count int64
err := s.store.DB().Model(&model.PendingJoin{}).Where("status = ?", "pending").Count(&count).Error
return count, err
}
+143
View File
@@ -0,0 +1,143 @@
package service
import (
"errors"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"gorm.io/gorm"
)
// PolicyService 策略管理服务
type PolicyService struct {
store *sqlite.Store
}
// NewPolicyService 创建策略服务实例
func NewPolicyService(store *sqlite.Store) *PolicyService {
return &PolicyService{store: store}
}
// ListPolicies 获取策略列表
func (s *PolicyService) ListPolicies() ([]model.Policy, error) {
var policies []model.Policy
err := s.store.DB().Order("created_at desc").Find(&policies).Error
return policies, err
}
// GetPolicyByID 根据 ID 获取策略
func (s *PolicyService) GetPolicyByID(id uint) (*model.Policy, error) {
var policy model.Policy
if err := s.store.DB().First(&policy, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("策略不存在")
}
return nil, err
}
return &policy, nil
}
// CreatePolicy 创建策略
func (s *PolicyService) CreatePolicy(req *model.Policy) (*model.Policy, error) {
if req.Name == "" {
return nil, errors.New("策略名称不能为空")
}
// 检查名称是否重复
var existing model.Policy
if err := s.store.DB().Where("name = ?", req.Name).First(&existing).Error; err == nil {
return nil, errors.New("策略名称已存在")
}
// 默认值
if req.Type == "" {
req.Type = "custom"
}
req.Enabled = true
if err := s.store.DB().Create(req).Error; err != nil {
return nil, err
}
return req, nil
}
// UpdatePolicy 更新策略
func (s *PolicyService) UpdatePolicy(id uint, updates map[string]interface{}) (*model.Policy, error) {
var policy model.Policy
if err := s.store.DB().First(&policy, id).Error; err != nil {
return nil, errors.New("策略不存在")
}
// 系统默认策略不可编辑
if policy.Type == "system" {
return nil, errors.New("系统默认策略不可编辑")
}
if err := s.store.DB().Model(&policy).Updates(updates).Error; err != nil {
return nil, err
}
return &policy, nil
}
// DeletePolicy 删除策略
func (s *PolicyService) DeletePolicy(id uint) error {
var policy model.Policy
if err := s.store.DB().First(&policy, id).Error; err != nil {
return errors.New("策略不存在")
}
// 系统默认策略不可删除
if policy.Type == "system" {
return errors.New("系统默认策略不可删除")
}
// 检查是否有关联的组网
var count int64
s.store.DB().Model(&model.Network{}).Where("policy_id = ?", id).Count(&count)
if count > 0 {
return errors.New("该策略正在被组网引用,无法删除")
}
return s.store.DB().Delete(&policy).Error
}
// InitializeDefaultPolicy 初始化默认策略
func (s *PolicyService) InitializeDefaultPolicy() error {
// 检查是否已存在默认策略
var count int64
s.store.DB().Model(&model.Policy{}).Where("type = ?", "system").Count(&count)
if count > 0 {
return nil // 已存在默认策略
}
// 创建默认策略
defaultPolicy := &model.Policy{
Name: "系统默认策略",
Type: "system",
Description: "MeshRay 系统默认传输策略,包含基础的 P2P 直连和 TURN 中继功能",
Enabled: true,
IsDefault: true,
LayerConfig: `{
"tunnel": {"enabled": true},
"obfuscation": {"enabled": false, "method": "xor"},
"compression": {"enabled": false, "algorithm": "lz4"},
"mux": {"enabled": false, "protocol": "yamux", "streams": 4},
"fec": {"enabled": false, "mode": "rs", "redundancy": 20},
"qos": {"enabled": false},
"nat": {"enabled": true},
"relay": {"enabled": true}
}`,
GlobalParams: `{
"encryption": "chacha20poly1305",
"ipv6_enabled": false,
"connect_timeout": 10,
"fallback_threshold": 2000
}`,
}
if err := s.store.DB().Create(defaultPolicy).Error; err != nil {
return err
}
return nil
}
+61
View File
@@ -0,0 +1,61 @@
package service
import (
"os"
"os/exec"
"syscall"
"time"
"go.uber.org/zap"
)
// RestartCoreService 重启核心服务(用于管理员操作)
type RestartCoreService struct {
logger *zap.Logger
}
// NewRestartCoreService 创建重启核心服务
func NewRestartCoreService(logger *zap.Logger) *RestartCoreService {
return &RestartCoreService{
logger: logger,
}
}
// RestartCore 重启核心服务(优雅重启)
func (s *RestartCoreService) RestartCore() error {
s.logger.Info("开始重启核心服务...")
// 1. 记录当前进程 ID
pid := os.Getpid()
s.logger.Info("当前进程 ID", zap.Int("pid", pid))
// 2. 获取当前可执行文件路径
execPath, err := os.Executable()
if err != nil {
s.logger.Error("获取可执行文件路径失败", zap.Error(err))
return err
}
// 3. 启动新进程
cmd := exec.Command(execPath)
cmd.SysProcAttr = &syscall.SysProcAttr{
HideWindow: true, // Windows 隐藏窗口
CreationFlags: syscall.CREATE_NEW_PROCESS_GROUP, // Windows 创建新进程组
}
if err := cmd.Start(); err != nil {
s.logger.Error("启动新进程失败", zap.Error(err))
return err
}
s.logger.Info("新进程启动成功", zap.Int("new_pid", cmd.Process.Pid))
// 4. 等待一小段时间确保新进程稳定
time.Sleep(2 * time.Second)
// 5. 退出当前进程
s.logger.Info("当前进程即将退出")
os.Exit(0)
return nil
}
+284
View File
@@ -0,0 +1,284 @@
package service
import (
"context"
"errors"
"fmt"
"net"
"strconv"
"strings"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/dnsprovider"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"github.com/google/uuid"
"github.com/libdns/libdns"
"gorm.io/gorm"
)
// ServiceService 服务管理服务
type ServiceService struct {
store *sqlite.Store
}
// NewServiceService 创建服务服务实例
func NewServiceService(store *sqlite.Store) *ServiceService {
return &ServiceService{store: store}
}
// ListServices 获取服务列表(支持按类型筛选)
func (s *ServiceService) ListServices(serverType string) ([]model.Service, error) {
var services []model.Service
query := s.store.DB().Order("created_at desc")
if serverType != "" {
// 支持大小写不敏感匹配
query = query.Where("LOWER(type) = LOWER(?)", serverType)
}
err := query.Find(&services).Error
return services, err
}
// GetServiceByID 根据 ID 获取服务
func (s *ServiceService) GetServiceByID(id string) (*model.Service, error) {
var service model.Service
if err := s.store.DB().Where("id = ?", id).First(&service).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("服务不存在")
}
return nil, err
}
return &service, nil
}
// CreateService 创建服务
func (s *ServiceService) CreateService(req *model.Service) (*model.Service, error) {
if req.Name == "" {
return nil, errors.New("服务名称不能为空")
}
if req.Type == "" {
return nil, errors.New("服务类型不能为空")
}
// 强制转换为大写以保持一致性
req.Type = strings.ToUpper(req.Type)
// 校验支持的类型
allowedTypes := map[string]bool{
"STUN": true, "TURN": true, "DDNS": true, "TUN": true, "CUSTOM": true,
}
if !allowedTypes[req.Type] {
return nil, fmt.Errorf("不支持的服务类型: %s", req.Type)
}
// 地址和端口校验(除 TUN 外通常需要)
if req.Type != "TUN" {
if req.Address == "" {
return nil, errors.New("服务器地址不能为空")
}
if req.Port <= 0 || req.Port > 65535 {
return nil, errors.New("端口号无效")
}
}
// DDNS 鉴权校验(基础设施模式)
if req.Type == "DDNS" && req.ConfigMode == "infrastructure" {
if req.Provider == "" {
return nil, errors.New("请选择 DNS 服务商")
}
if req.Domain == "" {
return nil, errors.New("请输入根域名")
}
// 根据服务商校验认证信息
switch req.Provider {
case "cloudflare":
if req.Token == "" {
return nil, errors.New("请输入 API Token")
}
case "aliyun":
if req.AuthUsername == "" || req.AuthPassword == "" {
return nil, errors.New("请输入 AccessKey ID 和 Secret")
}
case "tencent":
if req.AuthUsername == "" || req.AuthPassword == "" {
return nil, errors.New("请输入 SecretId 和 SecretKey")
}
}
}
// DDNS 全功能模式校验
if req.Type == "DDNS" && req.ConfigMode == "fullservice" {
if req.DDNSConfigID == "" {
return nil, errors.New("请选择 DDNS 配置")
}
if req.RecordType == "" {
return nil, errors.New("请选择记录类型")
}
// 根据记录类型校验字段
switch req.RecordType {
case "A", "AAAA":
if req.Subdomain == "" {
return nil, errors.New("请输入主机记录")
}
if req.TargetIP == "" {
return nil, errors.New("请输入目标 IP")
}
if req.Port <= 0 {
return nil, errors.New("请输入检测端口")
}
case "TXT":
if req.TXTRecordName == "" {
return nil, errors.New("请输入 TXT 记录名称")
}
if req.TXTValue == "" {
return nil, errors.New("请输入 TXT 记录值")
}
case "CNAME":
if req.CNAMETarget == "" {
return nil, errors.New("请输入目标域名")
}
}
}
// 生成 ID
if req.ID == "" {
req.ID = uuid.New().String()
}
// 如果是 DDNS 全功能模式,先创建 DNS 记录
if req.Type == "DDNS" && req.ConfigMode == "fullservice" {
// 使用事务确保原子性
tx := s.store.DB().Begin()
defer func() {
if r := recover(); r != nil {
tx.Rollback()
}
}()
// 1. 获取关联的 DDNS 配置
var ddnsConfig model.Service
if err := tx.Where("id = ?", req.DDNSConfigID).First(&ddnsConfig).Error; err != nil {
tx.Rollback()
return nil, fmt.Errorf("获取 DDNS 配置失败:%w", err)
}
// 2. 确定记录类型、名称和值
var recordType, name, value string
switch req.RecordType {
case "A", "AAAA":
recordType = req.RecordType
name = req.Subdomain
value = req.TargetIP
case "TXT":
recordType = req.RecordType
name = req.TXTRecordName
value = req.TXTValue
case "CNAME":
recordType = req.RecordType
name = req.Subdomain
value = req.CNAMETarget
default:
tx.Rollback()
return nil, errors.New("不支持的记录类型")
}
// 3. 创建 DNS Provider
providerConfig := dnsprovider.ProviderConfig{
Provider: dnsprovider.ProviderType(ddnsConfig.Provider),
Domain: ddnsConfig.Domain,
APIToken: ddnsConfig.Token,
AccessKeyID: ddnsConfig.AuthUsername,
AccessKeySecret: ddnsConfig.AuthPassword,
SecretId: ddnsConfig.AuthUsername,
SecretKey: ddnsConfig.AuthPassword,
}
provider, err := dnsprovider.NewDNSProvider(providerConfig)
if err != nil {
tx.Rollback()
return nil, fmt.Errorf("创建 DNS Provider 失败:%w", err)
}
// 4. 构建并添加 DNS 记录
dnsRecord := &dnsprovider.DNSRecord{
Type: dnsprovider.RecordType(recordType),
Name: name,
Value: value,
TTL: req.TTL,
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
_, err = provider.AppendRecords(ctx, ddnsConfig.Domain, []libdns.Record{dnsRecord.ToLibdnsRecord()})
if err != nil {
tx.Rollback()
return nil, fmt.Errorf("创建 DNS 记录失败:%w", err)
}
// 5. DNS 记录创建成功,保存到数据库
if err := tx.Create(req).Error; err != nil {
tx.Rollback()
return nil, err
}
tx.Commit()
return req, nil
}
// 其他情况直接保存
if err := s.store.DB().Create(req).Error; err != nil {
return nil, err
}
return req, nil
}
// UpdateService 更新服务
func (s *ServiceService) UpdateService(id string, updates map[string]interface{}) (*model.Service, error) {
var service model.Service
if err := s.store.DB().Where("id = ?", id).First(&service).Error; err != nil {
return nil, errors.New("服务不存在")
}
if err := s.store.DB().Model(&service).Updates(updates).Error; err != nil {
return nil, err
}
return &service, nil
}
// DeleteService 删除服务
func (s *ServiceService) DeleteService(id string) error {
var service model.Service
if err := s.store.DB().Where("id = ?", id).First(&service).Error; err != nil {
return errors.New("服务不存在")
}
return s.store.DB().Delete(&service).Error
}
// TestServiceConnectivity 测试服务连通性
func (s *ServiceService) TestServiceConnectivity(id string) (string, error) {
var service model.Service
if err := s.store.DB().Where("id = ?", id).First(&service).Error; err != nil {
return "", errors.New("服务不存在")
}
addr := net.JoinHostPort(service.Address, strconv.Itoa(service.Port))
if service.Type == "STUN" {
// STUN 使用 UDP 测试
conn, err := net.DialTimeout("udp", addr, 5*time.Second)
if err != nil {
return "unreachable", nil
}
conn.Close()
return "ok", nil
}
// 其他使用 TCP 测试
conn, err := net.DialTimeout("tcp", addr, 5*time.Second)
if err != nil {
return "unreachable", nil
}
conn.Close()
return "ok", nil
}
+128
View File
@@ -0,0 +1,128 @@
package service
import (
"encoding/json"
"errors"
"fmt"
"git.zkcoi.com/zkcoi/meshray/internal/model"
sqlite "git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
)
// SettingsService 系统设置服务
type SettingsService struct {
store *sqlite.Store
logger *zap.Logger
}
// NewSettingsService 创建系统设置服务
func NewSettingsService(store *sqlite.Store, logger *zap.Logger) *SettingsService {
return &SettingsService{
store: store,
logger: logger,
}
}
// GetSettings 获取系统设置(单例)
func (s *SettingsService) GetSettings() (*model.SystemSetting, error) {
var setting model.SystemSetting
// 尝试查找 ID=1 的记录
result := s.store.DB().First(&setting, 1)
if result.Error != nil {
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
// 如果不存在,创建默认设置
setting = model.SystemSetting{
ID: 1,
ServerIP: "",
ServerPort: 51820,
LogLevel: "info",
LogFormat: "console",
MaxBackups: 7,
MaxAge: 30,
Theme: "light",
Language: "zh-CN",
}
if err := s.store.DB().Create(&setting).Error; err != nil {
return nil, fmt.Errorf("创建默认设置失败:%w", err)
}
s.logger.Info("已创建默认系统设置")
return &setting, nil
}
return nil, fmt.Errorf("查询设置失败:%w", result.Error)
}
return &setting, nil
}
// UpdateSettings 更新系统设置
func (s *SettingsService) UpdateSettings(updates map[string]interface{}) (*model.SystemSetting, error) {
// 先获取现有设置
setting, err := s.GetSettings()
if err != nil {
return nil, err
}
// JSON 序列化 updates,确保数据有效
data, err := json.Marshal(updates)
if err != nil {
return nil, fmt.Errorf("序列化更新数据失败:%w", err)
}
// 反序列化到临时对象,过滤无效字段
var validUpdates map[string]interface{}
if err := json.Unmarshal(data, &validUpdates); err != nil {
return nil, fmt.Errorf("解析更新数据失败:%w", err)
}
// 移除敏感字段和不可变字段
delete(validUpdates, "id")
delete(validUpdates, "created_at")
delete(validUpdates, "updated_at")
// 构建 GORM 的 Updates map
updatesMap := make(map[string]interface{})
for key, value := range validUpdates {
// 驼峰转蛇形转换(可选,这里直接使用前端传来的键名)
updatesMap[key] = value
}
// 执行更新
if err := s.store.DB().Model(&setting).Updates(updatesMap).Error; err != nil {
return nil, fmt.Errorf("更新设置失败:%w", err)
}
s.logger.Info("系统设置已更新", zap.Any("updates", updates))
// 返回最新数据
return s.GetSettings()
}
// ResetSettings 重置为默认设置
func (s *SettingsService) ResetSettings() (*model.SystemSetting, error) {
_, err := s.GetSettings()
if err != nil {
return nil, err
}
defaultSetting := model.SystemSetting{
ID: 1,
ServerIP: "",
ServerPort: 51820,
LogLevel: "info",
LogFormat: "console",
MaxBackups: 7,
MaxAge: 30,
Theme: "light",
Language: "zh-CN",
}
if err := s.store.DB().Save(&defaultSetting).Error; err != nil {
return nil, fmt.Errorf("重置设置失败:%w", err)
}
s.logger.Info("系统设置已重置为默认值")
return &defaultSetting, nil
}
+120
View File
@@ -0,0 +1,120 @@
package service
import (
"errors"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"go.uber.org/zap"
"golang.org/x/crypto/bcrypt"
"gorm.io/gorm"
)
// SystemConfigService 系统配置服务
type SystemConfigService struct {
store *sqlite.Store
logger *zap.Logger
}
// NewSystemConfigService 创建系统配置服务
func NewSystemConfigService(store *sqlite.Store, logger *zap.Logger) *SystemConfigService {
return &SystemConfigService{
store: store,
logger: logger,
}
}
// GetWGMode 获取当前 WG 运行模式
func (s *SystemConfigService) GetWGMode() (string, error) {
config, err := s.getConfig("wg_mode")
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
// 如果不存在,返回默认值 "auto"
return "auto", nil
}
return "", err
}
return config.Value, nil
}
// SetWGMode 设置 WG 运行模式(需要重启 MeshRay 才能生效)
func (s *SystemConfigService) SetWGMode(mode string) error {
// 验证模式值
validModes := map[string]bool{
"auto": true,
"kernel": true,
"userspace": true,
}
if !validModes[mode] {
return errors.New("无效的 WG 模式,仅支持:auto, kernel, userspace")
}
return s.setConfig("wg_mode", mode)
}
// GetConfig 获取任意配置项
func (s *SystemConfigService) GetConfig(key string) (string, error) {
config, err := s.getConfig(key)
if err != nil {
return "", err
}
return config.Value, nil
}
// SetConfig 设置任意配置项
func (s *SystemConfigService) SetConfig(key, value string) error {
return s.setConfig(key, value)
}
// getConfig 内部方法:获取配置
func (s *SystemConfigService) getConfig(key string) (*model.SystemConfig, error) {
var config model.SystemConfig
err := s.store.DB().Where("key = ?", key).First(&config).Error
if err != nil {
return nil, err
}
return &config, nil
}
// setConfig 内部方法:设置配置
func (s *SystemConfigService) setConfig(key, value string) error {
var config model.SystemConfig
err := s.store.DB().Where("key = ?", key).First(&config).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
// 创建新配置
config = model.SystemConfig{
Key: key,
Value: value,
}
return s.store.DB().Create(&config).Error
}
// 更新现有配置
config.Value = value
return s.store.DB().Save(&config).Error
}
// ChangePassword 修改密码(P1
func (s *SystemConfigService) ChangePassword(userID uint, oldPassword, newPassword string) error {
// 获取用户
var user model.User
if err := s.store.DB().First(&user, userID).Error; err != nil {
return errors.New("用户不存在")
}
// 验证旧密码
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(oldPassword)); err != nil {
return errors.New("原密码错误")
}
// 加密新密码
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
if err != nil {
s.logger.Error("加密新密码失败", zap.Error(err))
return errors.New("密码加密失败")
}
// 更新密码
return s.store.DB().Model(&user).Update("password_hash", string(hashedPassword)).Error
}
+221
View File
@@ -0,0 +1,221 @@
package service
import (
"crypto/rand"
"encoding/base64"
"errors"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"golang.org/x/crypto/bcrypt"
"gorm.io/gorm"
)
// 字符集用于生成随机密码(保留用于特殊场景)
const passwordChars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789!@#$%^&*"
// generateRandomPassword 使用 crypto/rand 生成安全的随机密码
func generateRandomPassword(length int) string {
b := make([]byte, length)
_, err := rand.Read(b)
if err != nil {
// 极端情况下回退到简单方案(几乎不会发生)
return "REPLACE_WITH_SECURE_PASSWORD"
}
// 使用 Base64 编码,确保包含各种字符
encoded := base64.StdEncoding.EncodeToString(b)
// 截取所需长度(Base64 编码后长度为 4/3 倍)
if len(encoded) >= length {
return encoded[:length]
}
return encoded
}
// GenerateRandomPassword 生成随机密码(公开函数)
func GenerateRandomPassword(length int) string {
return generateRandomPassword(length)
}
// UserService 用户服务
type UserService struct {
store *sqlite.Store
}
// NewUserService 创建用户服务实例
func NewUserService(store *sqlite.Store) *UserService {
return &UserService{store: store}
}
// InitializeAdmin 初始化管理员账户(首次启动时调用)
func (s *UserService) InitializeAdmin() (username, password string, err error) {
// 检查是否已存在 admin 用户
var existingUser model.User
if err := s.store.DB().Where("role = ?", "admin").First(&existingUser).Error; err == nil {
// 已存在,返回当前用户名(密码不返回)
return existingUser.Username, "", nil
}
// 生成随机密码
randomPassword := generateRandomPassword(16)
// 密码加密
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(randomPassword), bcrypt.DefaultCost)
if err != nil {
return "", "", err
}
// 创建 admin 用户
adminUser := &model.User{
Username: "admin",
PasswordHash: string(hashedPassword),
Email: "",
Role: "admin",
Status: "active",
}
if err := s.store.DB().Create(adminUser).Error; err != nil {
return "", "", err
}
return "admin", randomPassword, nil
}
// Authenticate 验证用户登录
func (s *UserService) Authenticate(username, password string) (*model.User, error) {
var user model.User
if err := s.store.DB().Where("username = ?", username).First(&user).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("用户不存在")
}
return nil, err
}
// 验证密码
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil {
return nil, errors.New("密码错误")
}
// 检查用户状态
if user.Status != "active" {
return nil, errors.New("账户已被禁用")
}
return &user, nil
}
// ChangePasswordRequest 修改密码请求
type ChangePasswordRequest struct {
OldPassword string `json:"old_password"`
NewPassword string `json:"new_password"`
}
// ChangePassword 修改用户密码
func (s *UserService) ChangePassword(userID uint, req *ChangePasswordRequest) error {
// 1. 查询用户
var user model.User
if err := s.store.DB().First(&user, userID).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New("用户不存在")
}
return err
}
// 2. 验证旧密码
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.OldPassword)); err != nil {
return errors.New("原密码错误")
}
// 3. 验证新密码强度
if len(req.NewPassword) < 6 {
return errors.New("密码长度不能少于 6 位")
}
// 4. 加密新密码
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(req.NewPassword), bcrypt.DefaultCost)
if err != nil {
return err
}
// 5. 更新密码
user.PasswordHash = string(hashedPassword)
if err := s.store.DB().Save(&user).Error; err != nil {
return err
}
return nil
}
// GetUserByID 根据 ID 获取用户
func (s *UserService) GetUserByID(userID uint) (*model.User, error) {
var user model.User
if err := s.store.DB().First(&user, userID).Error; err != nil {
return nil, err
}
return &user, nil
}
// UpdateUser 更新用户信息
func (s *UserService) UpdateUser(userID uint, updates map[string]interface{}) (*model.User, error) {
var user model.User
if err := s.store.DB().First(&user, userID).Error; err != nil {
return nil, err
}
if err := s.store.DB().Model(&user).Updates(updates).Error; err != nil {
return nil, err
}
return &user, nil
}
// UpdateAdminProfile 更新管理员资料(仅允许 admin 用户调用)
func (s *UserService) UpdateAdminProfile(username, email, password string) error {
var user model.User
if err := s.store.DB().Where("role = ?", "admin").First(&user).Error; err != nil {
return errors.New("管理员账户不存在")
}
updates := make(map[string]interface{})
if username != "" && username != user.Username {
// 检查新用户名是否已被使用
var existing model.User
if err := s.store.DB().Where("username = ? AND role = ?", username, "admin").First(&existing).Error; err == nil {
return errors.New("用户名已存在")
}
updates["username"] = username
}
if email != "" && email != user.Email {
updates["email"] = email
}
if password != "" {
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return err
}
updates["password_hash"] = string(hashedPassword)
}
if len(updates) == 0 {
return nil // 没有需要更新的
}
return s.store.DB().Model(&user).Updates(updates).Error
}
// ResetAdminPassword 重置管理员密码(命令行工具使用)
func (s *UserService) ResetAdminPassword(newPassword string) error {
var user model.User
if err := s.store.DB().Where("role = ?", "admin").First(&user).Error; err != nil {
return errors.New("管理员账户不存在")
}
// 加密新密码
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
if err != nil {
return err
}
return s.store.DB().Model(&user).Update("password_hash", string(hashedPassword)).Error
}