113 lines
2.2 KiB
Go
113 lines
2.2 KiB
Go
package alloc
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net"
|
|
"strings"
|
|
"sync"
|
|
|
|
"git.giftfish.de/ston1th/haproxy-lb/pkg/db"
|
|
"github.com/mikioh/ipaddr"
|
|
)
|
|
|
|
type Alloc struct {
|
|
sync.Mutex
|
|
db *db.DB
|
|
pool []*net.IPNet
|
|
gateway *net.IPNet
|
|
}
|
|
|
|
func NewAlloc(db *db.DB, cidrs []string, gateway string) (*Alloc, error) {
|
|
pool, err := parseRange(cidrs)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
_, gw, err := net.ParseCIDR(gateway)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &Alloc{db: db, pool: pool, gateway: gw}, nil
|
|
}
|
|
|
|
func (a *Alloc) UpdateDB(db *db.DB) {
|
|
a.Lock()
|
|
a.db = db
|
|
a.Unlock()
|
|
}
|
|
|
|
func (a *Alloc) cidr(ip net.IP) string {
|
|
gw := new(net.IPNet)
|
|
*gw = *a.gateway
|
|
gw.IP = ip
|
|
return gw.String()
|
|
}
|
|
|
|
func (a *Alloc) AllocIP(ctx context.Context, name string) (addr string, err error) {
|
|
a.Lock()
|
|
defer a.Unlock()
|
|
for _, cidr := range a.pool {
|
|
c := ipaddr.NewCursor([]ipaddr.Prefix{*ipaddr.NewPrefix(cidr)})
|
|
for pos := c.First(); pos != nil; pos = c.Next() {
|
|
cidr := a.cidr(pos.IP)
|
|
if !a.db.IPExists(cidr) {
|
|
err = a.db.SetIP(cidr, name)
|
|
if err != nil {
|
|
return
|
|
}
|
|
addr = pos.IP.String()
|
|
return
|
|
}
|
|
}
|
|
}
|
|
err = fmt.Errorf("no more IPs available in pool")
|
|
return
|
|
}
|
|
|
|
func (a *Alloc) FreeIP(ctx context.Context, name string) (err error) {
|
|
a.Lock()
|
|
defer a.Unlock()
|
|
if name == "" {
|
|
return
|
|
}
|
|
ip, err := a.db.GetName(name)
|
|
if err != nil {
|
|
return
|
|
}
|
|
return a.db.DeleteIP(ip, name)
|
|
}
|
|
|
|
func parseRange(cidrs []string) (nets []*net.IPNet, err error) {
|
|
for _, cidr := range cidrs {
|
|
if !strings.Contains(cidr, "-") {
|
|
_, n, err := net.ParseCIDR(cidr)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
nets = append(nets, n)
|
|
continue
|
|
}
|
|
|
|
fs := strings.SplitN(cidr, "-", 2)
|
|
if len(fs) != 2 {
|
|
return nil, fmt.Errorf("invalid IP range %q", cidr)
|
|
}
|
|
start := net.ParseIP(strings.TrimSpace(fs[0]))
|
|
if start == nil {
|
|
return nil, fmt.Errorf("invalid IP range %q: invalid start IP %q", cidr, fs[0])
|
|
}
|
|
end := net.ParseIP(strings.TrimSpace(fs[1]))
|
|
if end == nil {
|
|
return nil, fmt.Errorf("invalid IP range %q: invalid end IP %q", cidr, fs[1])
|
|
}
|
|
|
|
for _, pfx := range ipaddr.Summarize(start, end) {
|
|
n := &net.IPNet{
|
|
IP: pfx.IP,
|
|
Mask: pfx.Mask,
|
|
}
|
|
nets = append(nets, n)
|
|
}
|
|
}
|
|
return
|
|
}
|