164 lines
4.0 KiB
Go
164 lines
4.0 KiB
Go
package service
|
|
|
|
import (
|
|
"fmt"
|
|
"math/rand"
|
|
"net"
|
|
|
|
"zeromesh/internal/model"
|
|
"zeromesh/internal/repository"
|
|
)
|
|
|
|
type NetworkService struct {
|
|
netRepo *repository.NetworkRepo
|
|
nodeRepo *repository.NodeRepo
|
|
userRepo *repository.UserRepo
|
|
}
|
|
|
|
func NewNetworkService(netRepo *repository.NetworkRepo, nodeRepo *repository.NodeRepo, userRepo *repository.UserRepo) *NetworkService {
|
|
return &NetworkService{netRepo: netRepo, nodeRepo: nodeRepo, userRepo: userRepo}
|
|
}
|
|
|
|
func randomPrivateSubnet() string {
|
|
const (
|
|
_10 = iota // 10.x.y.0/24
|
|
_172 // 172.16-31.y.0/24
|
|
_192 // 192.168.y.0/24
|
|
)
|
|
class := rand.Intn(3)
|
|
var b1, b2 byte
|
|
switch class {
|
|
case _10:
|
|
b1 = byte(rand.Intn(256))
|
|
b2 = byte(rand.Intn(256))
|
|
return fmt.Sprintf("10.%d.%d.0/24", b1, b2)
|
|
case _172:
|
|
b1 = byte(16 + rand.Intn(16)) // 16-31
|
|
b2 = byte(rand.Intn(256))
|
|
return fmt.Sprintf("172.%d.%d.0/24", b1, b2)
|
|
default:
|
|
b2 = byte(rand.Intn(256))
|
|
return fmt.Sprintf("192.168.%d.0/24", b2)
|
|
}
|
|
}
|
|
|
|
func (s *NetworkService) CreateNetwork(userID uint, name, ipRange string) (*model.Network, error) {
|
|
user, err := s.userRepo.FindByID(userID)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("user not found")
|
|
}
|
|
if user.QuotaNetworks >= 0 {
|
|
used, _ := s.netRepo.CountByUser(userID)
|
|
if used >= int64(user.QuotaNetworks) {
|
|
return nil, fmt.Errorf("network quota exceeded (%d)", user.QuotaNetworks)
|
|
}
|
|
}
|
|
if ipRange == "" {
|
|
ipRange = randomPrivateSubnet()
|
|
}
|
|
id := rand.Uint32()
|
|
network := &model.Network{
|
|
UserID: userID,
|
|
NetworkID: id,
|
|
Name: name,
|
|
IPRange: ipRange,
|
|
MTU: 2800,
|
|
Multicast: true,
|
|
Private: true,
|
|
}
|
|
if err := s.netRepo.Create(network); err != nil {
|
|
return nil, err
|
|
}
|
|
return network, nil
|
|
}
|
|
|
|
func (s *NetworkService) ListNetworks(userID uint) ([]model.Network, error) {
|
|
if userID == 0 {
|
|
return s.netRepo.List()
|
|
}
|
|
return s.netRepo.ListByUser(userID)
|
|
}
|
|
|
|
func (s *NetworkService) GetNetwork(networkID uint32, userID uint) (*model.Network, error) {
|
|
if userID == 0 {
|
|
return s.netRepo.FindByNetworkID(networkID)
|
|
}
|
|
return s.netRepo.FindByNetworkIDForUser(networkID, userID)
|
|
}
|
|
|
|
func (s *NetworkService) DeleteNetwork(networkID uint32, userID uint) error {
|
|
net, err := s.GetNetwork(networkID, userID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return s.netRepo.Delete(net.NetworkID)
|
|
}
|
|
|
|
func (s *NetworkService) AuthorizeMember(networkID uint32, nodeID string, userID uint) error {
|
|
network, err := s.GetNetwork(networkID, userID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
node, err := s.nodeRepo.FindByNodeID(nodeID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if node.UserID != userID && userID != 0 {
|
|
return fmt.Errorf("node does not belong to user")
|
|
}
|
|
ip, err := s.allocateIP(network.IPRange, networkID, nodeID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
member := &model.NetworkMember{
|
|
NetworkID: networkID,
|
|
NodeID: nodeID,
|
|
IPAddress: ip,
|
|
Authorized: true,
|
|
Label: node.Name,
|
|
}
|
|
if err := s.netRepo.AddMember(member); err != nil {
|
|
return fmt.Errorf("add member: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *NetworkService) DeauthorizeMember(networkID uint32, nodeID string, userID uint) error {
|
|
_, err := s.GetNetwork(networkID, userID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return s.netRepo.RemoveMember(networkID, nodeID)
|
|
}
|
|
|
|
func (s *NetworkService) ListMembers(networkID uint32, userID uint) ([]model.NetworkMember, error) {
|
|
_, err := s.GetNetwork(networkID, userID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return s.netRepo.FindMembers(networkID)
|
|
}
|
|
|
|
func (s *NetworkService) allocateIP(cidr string, networkID uint32, nodeID string) (string, error) {
|
|
_, ipNet, err := net.ParseCIDR(cidr)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
ones, bits := ipNet.Mask.Size()
|
|
ip := ipNet.IP.To4()
|
|
if ip == nil {
|
|
return "", fmt.Errorf("IPv4 only")
|
|
}
|
|
// simple allocation: use last 2 bytes of nodeID hash
|
|
h := uint32(0)
|
|
for _, b := range []byte(nodeID) {
|
|
h = h*31 + uint32(b)
|
|
}
|
|
hostBits := bits - ones
|
|
maxHosts := (1 << uint(hostBits)) - 2
|
|
hostOffset := int(h%uint32(maxHosts)) + 1
|
|
ip[2] = byte(hostOffset >> 8)
|
|
ip[3] = byte(hostOffset)
|
|
return ip.String(), nil
|
|
}
|