Files
Meshray-Manager/internal/service/pending_join.go
T
2026-06-30 15:14:37 +08:00

250 lines
7.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}