Initial commit

This commit is contained in:
2026-06-30 15:14:37 +08:00
commit 15dab96872
311 changed files with 95639 additions and 0 deletions
+249
View File
@@ -0,0 +1,249 @@
package service
import (
"errors"
"fmt"
"net"
"strings"
"time"
"git.zkcoi.com/zkcoi/meshray/internal/model"
"git.zkcoi.com/zkcoi/meshray/internal/store/sqlite"
"gorm.io/gorm"
)
// PendingJoinService 待审核服务
type PendingJoinService struct {
store *sqlite.Store
}
// NewPendingJoinService 创建待审核服务实例
func NewPendingJoinService(store *sqlite.Store) *PendingJoinService {
return &PendingJoinService{store: store}
}
// ListPendingJoins 获取待审核列表
func (s *PendingJoinService) ListPendingJoins(networkID uint, status string, page, size int) ([]model.PendingJoin, int64, error) {
query := s.store.DB().Model(&model.PendingJoin{})
// 按网络 ID 筛选
if networkID > 0 {
query = query.Where("network_id = ?", networkID)
}
// 按状态筛选
if status != "" && status != "all" {
query = query.Where("status = ?", status)
}
// 统计总数
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
// 分页查询
var result []model.PendingJoin
offset := (page - 1) * size
err := query.Order("created_at DESC").Offset(offset).Limit(size).Find(&result).Error
if err != nil {
return nil, 0, err
}
return result, total, nil
}
// GetPendingJoinByID 根据 ID 获取待审核记录
func (s *PendingJoinService) GetPendingJoinByID(id uint) (*model.PendingJoin, error) {
var record model.PendingJoin
err := s.store.DB().First(&record, id).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("记录不存在")
}
return nil, err
}
return &record, nil
}
// ApproveResult 审核通过结果
type ApproveResult struct {
Device *model.Device `json:"device"`
PrivateKey string `json:"private_key"` // 仅首次返回
Network *model.Network `json:"network"`
ConfigText string `json:"config_text"` // WireGuard 配置文本
}
// ApproveJoin 审核通过
func (s *PendingJoinService) ApproveJoin(id uint) (*ApproveResult, error) {
var record model.PendingJoin
if err := s.store.DB().First(&record, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("记录不存在")
}
return nil, err
}
if record.Status != "pending" {
return nil, errors.New("该申请已处理")
}
// 1. 查询 MeshSeed 获取网络信息
var meshSeed model.MeshSeed
if err := s.store.DB().Where("seed_id = ?", record.SeedID).First(&meshSeed).Error; err != nil {
return nil, fmt.Errorf("查询 MeshSeed 失败:%w", err)
}
// 2. 查询网络详情
var network model.Network
if err := s.store.DB().First(&network, meshSeed.NetworkID).Error; err != nil {
return nil, fmt.Errorf("查询网络失败:%w", err)
}
// 3. 生成 WireGuard 密钥对(调用全局函数,定义在 device.go 中)
privateKey, publicKey, err := generateWireGuardKeys()
if err != nil {
return nil, fmt.Errorf("生成密钥失败:%w", err)
}
// 4. 分配 IP 地址
ipAddress, err := s.allocateIPAddress(&network)
if err != nil {
return nil, fmt.Errorf("分配 IP 失败:%w", err)
}
// 5. 创建设备记录
device := &model.Device{
Name: record.DeviceName,
NetworkID: network.ID,
PublicKey: publicKey,
VirtualIP: ipAddress,
Status: "active",
LastSeen: time.Now(),
}
if err := s.store.DB().Create(device).Error; err != nil {
return nil, fmt.Errorf("创建设备失败:%w", err)
}
// 6. 更新审核状态
now := time.Now()
record.Status = "approved"
record.ApprovedAt = &now
if err := s.store.DB().Save(&record).Error; err != nil {
return nil, fmt.Errorf("更新审核状态失败:%w", err)
}
// 7. 生成 WireGuard 配置文本(使用 pending join 专用方法,读取 SystemSetting
configText := s.pendingJoinGenerateConfig(device, &network, privateKey)
return &ApproveResult{
Device: device,
PrivateKey: privateKey,
Network: &network,
ConfigText: configText,
}, nil
}
// RejectJoin 审核拒绝
func (s *PendingJoinService) RejectJoin(id uint, reason string) error {
var record model.PendingJoin
if err := s.store.DB().First(&record, id).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New("记录不存在")
}
return err
}
if record.Status != "pending" {
return errors.New("该申请已处理")
}
now := time.Now()
record.Status = "rejected"
record.RejectedAt = &now
record.Reason = reason
return s.store.DB().Save(&record).Error
}
// allocateIPAddress 分配 IP 地址
func (s *PendingJoinService) allocateIPAddress(network *model.Network) (string, error) {
// 解析子网
_, ipNet, err := net.ParseCIDR(network.SubnetIPv4)
if err != nil {
return "", fmt.Errorf("解析子网失败:%w", err)
}
// 获取已使用的 IP
var devices []model.Device
if err := s.store.DB().Where("network_id = ?", network.ID).Find(&devices).Error; err != nil {
return "", fmt.Errorf("查询设备列表失败:%w", err)
}
usedIPs := make(map[string]bool)
for _, device := range devices {
usedIPs[device.VirtualIP] = true
}
// 从 .2 开始分配(.1 通常是网关)
ip := ipNet.IP.To4()
if ip == nil {
return "", errors.New("暂不支持 IPv6 地址分配")
}
for i := 2; i < 254; i++ {
candidateIP := make(net.IP, len(ip))
copy(candidateIP, ip)
candidateIP[3] = byte(i)
if !usedIPs[candidateIP.String()] {
return candidateIP.String(), nil
}
}
return "", errors.New("IP 地址已用尽")
}
// pendingJoinGenerateConfig 生成 WireGuard 配置文本(从 SystemSetting 获取服务端信息)
func (s *PendingJoinService) pendingJoinGenerateConfig(device *model.Device, network *model.Network, privateKey string) string {
// 从 SystemSetting 获取服务端信息
var settings model.SystemSetting
if err := s.store.DB().First(&settings, 1).Error; err != nil {
// 使用默认值
settings.ServerPort = 51820
}
var sb strings.Builder
sb.WriteString("[Interface]\n")
sb.WriteString(fmt.Sprintf("PrivateKey = %s\n", privateKey))
sb.WriteString(fmt.Sprintf("Address = %s/32\n", device.VirtualIP))
sb.WriteString(fmt.Sprintf("MTU = %d\n\n", network.MTU))
sb.WriteString("[Peer]\n")
if settings.ServerPublicKey != "" {
sb.WriteString(fmt.Sprintf("PublicKey = %s\n", settings.ServerPublicKey))
} else {
sb.WriteString("PublicKey = <SERVER_PUBLIC_KEY>\n")
}
if settings.ServerIP != "" {
sb.WriteString(fmt.Sprintf("Endpoint = %s:%d\n", settings.ServerIP, settings.ServerPort))
} else {
sb.WriteString(fmt.Sprintf("Endpoint = <SERVER_IP>:%d\n", settings.ServerPort))
}
sb.WriteString(fmt.Sprintf("AllowedIPs = %s\n", network.SubnetIPv4))
sb.WriteString("PersistentKeepalive = 25\n")
return sb.String()
}
// DeleteExpired 删除过期的待审核记录
func (s *PendingJoinService) DeleteExpired() error {
return s.store.DB().Where("expire_at < ? AND status = ?", time.Now(), "pending").Delete(&model.PendingJoin{}).Error
}
// CountPending 统计待审核数量
func (s *PendingJoinService) CountPending() (int64, error) {
var count int64
err := s.store.DB().Model(&model.PendingJoin{}).Where("status = ?", "pending").Count(&count).Error
return count, err
}