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 = \n") } if settings.ServerIP != "" { sb.WriteString(fmt.Sprintf("Endpoint = %s:%d\n", settings.ServerIP, settings.ServerPort)) } else { sb.WriteString(fmt.Sprintf("Endpoint = :%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 }