Initial commit
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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):
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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. 加密 MeshSeed(AES-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 加密 MeshSeed(AES-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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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. 生成随机 SeedID(16 字节随机数)
|
||||
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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
// 计算 listenPort:51820 + 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
|
||||
}
|
||||
@@ -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(¬ification).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(¬ification)
|
||||
}
|
||||
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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user