Files
zeromesh/internal/service/network.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
}