py-vipman/vipman.py
2021-06-14 14:56:03 +02:00

315 lines
12 KiB
Python

# Copyright (C) 2021 Marius Schellenberger
import argparse
import logging
import os
import signal
import subprocess
import sys
from time import sleep
import uuid
import yaml
from pyroute2 import NDB, config
config.cache_expire = -1
from scapy.all import ARP, Ether, sendp, conf
conf.verb = 0
import etcd3
parser = argparse.ArgumentParser(description='vipman - etcd based virtual ip manager')
parser.add_argument('-c', '--config', default='/etc/vipman/vipman.yaml', help='config file')
args = parser.parse_args()
class ConfigError(Exception):
err: str
def __init__(self, err: str):
self.err = err
def __str__(self):
return self.err
class Config:
def __init__(self, file: str):
with open(file, 'r') as f:
c = yaml.safe_load(f)
vips_key = 'virtualIPs'
if not vips_key in c or c[vips_key] == None or len(c[vips_key]) == 0:
raise ConfigError(f"missing {vips_key} config")
self.vips = c[vips_key]
for vip in self.vips:
if len(vip.split('/')) != 2:
raise ConfigError(f"{vips_key}: invalid CIDR: '{vip}'")
iface_key = 'interface'
if not iface_key in c or c[iface_key] == None or c[iface_key] == '':
raise ConfigError(f"missing {iface_key} config")
self.iface = c[iface_key]
label_key = 'label'
if not label_key in c or c[label_key] == None or c[label_key] == '':
self.label = 'vipman'
else:
self.label = c[label_key]
leader_key = 'leaderHook'
if leader_key in c and c[leader_key] != None and len(c[leader_key]) > 0:
if not os.path.isabs(c[leader_key][0]):
raise ConfigError(f"{leader_key} executable path must be absolute")
self.leader = c[leader_key]
else:
self.leader = None
follower_key = 'followerHook'
if follower_key in c and c[follower_key] != None and len(c[follower_key]) > 0:
if not os.path.isabs(c[follower_key][0]):
raise ConfigError(f"{follower_key} executable path must be absolute")
self.follower = c[follower_key]
else:
self.follower = None
etcd_key = 'etcd'
if not etcd_key in c or c[etcd_key] == None:
raise ConfigError(f"missing {etcd_key} config")
etcd = c[etcd_key]
endpoints_key = 'endpoints'
if not endpoints_key in etcd or etcd[endpoints_key] == None or len(etcd[endpoints_key]) == 0:
raise ConfigError(f"missing {etcd_key}.{endpoints_key} config")
self.etcd_endpoints = etcd[endpoints_key]
prefix_key = 'prefix'
if not prefix_key in etcd or etcd[prefix_key] == None or etcd[prefix_key] == '':
self.etcd_prefix = '/vipman'
else:
prefix = etcd[prefix_key]
if not prefix.startswith('/'):
prefix = '/' + prefix
self.etcd_prefix = prefix.rstrip('/')
cluster_key = 'clusterName'
if not cluster_key in etcd or etcd[cluster_key] == None or etcd[cluster_key] == '':
raise ConfigError(f"missing {etcd_key}.{cluster_key} config")
self.etcd_cluster = etcd[cluster_key].rstrip('/')
self.etcd_path = self.etcd_prefix + '/' + self.etcd_cluster
user_key = 'username'
if not user_key in etcd or etcd[user_key] == None or etcd[user_key] == '':
self.etcd_user = None
else:
self.etcd_user = etcd[user_key]
password_key = 'password'
if not password_key in etcd or etcd[password_key] == None or etcd[password_key] == '':
self.etcd_password = None
else:
self.etcd_password = etcd[password_key]
tls_key = 'tls'
if tls_key in etcd and etcd[tls_key] != None:
tls = etcd[tls_key]
ca_key = 'ca'
if not ca_key in tls or tls[ca_key] == None or tls[ca_key] == '':
self.etcd_ca = None
else:
self.etcd_ca = tls[ca_key]
class Network():
interface: str
label: str
def __init__(self, interface: str, label: str, vips):
self.interface = interface
self.label = interface + ':' + label
self.vips = vips
def hasIP(self, cidr: str, ip) -> bool:
try:
ignored = ip.addresses[cidr]
return True
except KeyError:
return False
return False
def addIP(self, cidr: str):
ip = NDB()
if self.hasIP(cidr, ip):
ip.close()
return
try:
(addr, prefix) = cidr.split('/')
mac = ip.interfaces[self.interface]['address']
(ip.interfaces[self.interface].add_ip(address=addr, prefixlen=prefix, label=self.label).commit())
arp = ARP(psrc=addr, hwsrc=mac, pdst=addr)
sendp(Ether(dst='ff:ff:ff:ff:ff:ff') / arp)
finally:
ip.close()
def delIP(self, cidr: str):
ip = NDB()
if not self.hasIP(cidr, ip):
ip.close()
return
try:
(ip.interfaces[self.interface].del_ip(cidr).commit())
finally:
ip.close()
def addIPs(self):
for ip in self.vips:
self.addIP(ip)
def delIPs(self):
for ip in self.vips:
self.delIP(ip)
class Controller:
c: Config
n: Network
stepdown_flag: bool
stop_flag: bool
fail_flag: bool
lock_name = 'vipman'
def __init__(self, c: Config, n: Network):
self.c = c
self.n = n
self.stepdown_flag = False
self.stop_flag = False
self.fail_flag = False
self.leader = None
self.uuid = uuid.uuid1().bytes
self.uuid_str = self.uuid.hex()
self.log = logging.getLogger('controller')
self.lp = None
self.fp = None
def hook(self, leader: bool):
if leader and not self.leader:
self.leader = leader
if self.c.leader != None:
self.lp = subprocess.Popen(self.c.leader, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
if self.lp != None:
try:
out, err = self.lp.communicate(timeout=1)
if self.lp.returncode != 0:
errs = out.decode() + err.decode()
self.log.error(f"msg=\"error running hook\" leaderHook=\"{self.c.leader}\" error=\"{errs}\"")
self.lp = None
except subprocess.TimeoutExpired:
pass
if not leader and (self.leader or self.leader == None):
self.leader = leader
if self.c.follower != None:
self.fp = subprocess.Popen(self.c.follower, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
if self.fp != None:
try:
out, err = self.fp.communicate(timeout=1)
if self.fp.returncode != 0:
errs = out.decode() + err.decode()
self.log.error(f"msg=\"error running hook\" followerHook=\"{self.c.follower}\" error=\"{errs}\"")
self.fp = None
except subprocess.TimeoutExpired:
pass
def stepdown(self, sig, ignored):
if self.leader:
s = signal.strsignal(sig)
self.log.info(f"msg=\"signal received\" signal=\"{s}\" type=\"stepdown\"")
self.stepdown_flag = True
def stop(self, sig, ignored):
s = signal.strsignal(sig)
self.log.info(f"msg=\"signal received\" signal=\"{s}\" type=\"stop\"")
self.stop_flag = True
def run(self, etcd):
lock = etcd3.Lock(self.lock_name, ttl=10, etcd_client=etcd)
lock.key = self.c.etcd_path + '/' + self.lock_name
lock.uuid = self.uuid
while True:
if self.stop_flag:
lock.release()
return
if self.stepdown_flag:
self.stepdown_flag = False
self.hook(False)
lock.release()
sleep(5)
continue
if self.fail_flag:
self.fail_flag = False
if self.leader:
lock.release()
else:
sleep(5)
lock.acquire(timeout=5)
else:
if not lock.is_acquired():
lock.acquire(timeout=5)
if lock.is_acquired():
self.hook(True)
while True:
if not lock.is_acquired():
break
if self.stop_flag:
lock.release()
return
if self.stepdown_flag:
self.stepdown_flag = False
self.hook(False)
lock.release()
sleep(5)
break
self.log.info(f"msg=\"leader\" id=\"{self.uuid_str}\"")
try:
self.n.addIPs()
except Exception as e:
self.log.error(f"msg=\"error adding vips\" error=\"{e}\"")
lock.refresh()
sleep(5)
else:
self.hook(False)
if self.stop_flag:
return
self.log.info(f"msg=\"follower\" id=\"{self.uuid_str}\"")
try:
self.n.delIPs()
except Exception as e:
self.log.error(f"msg=\"error removing vips\" error=\"{e}\"")
def loop(self):
signal.signal(signal.SIGHUP, self.stepdown)
signal.signal(signal.SIGINT, self.stop)
signal.signal(signal.SIGTERM, self.stop)
while True:
for e in self.c.etcd_endpoints:
if self.stop_flag:
try:
self.n.delIPs()
except Exception as e:
self.log.error(f"msg=\"error removing vips\" error=\"{e}\"")
return
hp = e.split(':')
host = hp[0]
port = 2379
if len(hp) == 2:
port = hp[1]
self.log.info(f"msg=\"connecting to etcd\" host=\"{host}:{port}\"")
etcd = etcd3.client(host=host, port=port, ca_cert=self.c.etcd_ca,
user=self.c.etcd_user, password=self.c.etcd_password)
try:
(ok, ignored) = etcd.get(self.c.etcd_path)
if ok == None:
self.run(etcd)
except Exception as e:
self.fail_flag = True
self.log.error(f"msg=\"control loop error\" error=\"{e}\"")
sleep(1)
def main():
logging.basicConfig(level=logging.INFO,
format='%(asctime)s %(levelname).1s %(filename)s:%(lineno)d %(name)s: %(message)s')
log = logging.getLogger('main')
log.info(f"msg=\"starting vipman\" version=\"{__version__}\"")
cfg = None
try:
cfg = Config(args.config)
except Exception as e:
log.error(f"msg=\"error parsing config file\" error=\"{e}\"")
sys.exit(1)
n = Network(cfg.iface, cfg.label, cfg.vips)
c = Controller(cfg, n)
c.loop()
log.info(f"msg=\"stopped vipman\"")
__version__ = 'v1.0'
if __name__ == '__main__':
main()