- go.mod module path matches repo zkcoi/Meshray-Manager (case-distinct from zkcoi/Meshray/core)
- rewrite internal imports meshray/{internal,web,pkg} -> Meshray-Manager/... (core refs kept)
- sync README.md / install.sh repo URLs; add CHANGELOG entry
250 lines
7.0 KiB
Go
250 lines
7.0 KiB
Go
package service
|
||
|
||
import (
|
||
"errors"
|
||
"fmt"
|
||
"net"
|
||
"strings"
|
||
"time"
|
||
|
||
"git.zkcoi.com/zkcoi/Meshray-Manager/internal/model"
|
||
"git.zkcoi.com/zkcoi/Meshray-Manager/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
|
||
}
|