# 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()