#!/opt/alt/python35/bin/python3.5

import argparse
import json
import logging
import os
import pwd
import random
import re
import string
import sys


class CertCollection:

    ssl_cache_path = '/var/cache/imunify360-webshield/ssl.cache'
    domain_wildcard_regex = re.compile(r'^[*\w\d-]{,63}\.([\w\d.-]{,63})\w$')
    cache_user = 'imunify360-webshield'

    def __init__(self, exclude=None, logger=None):
        if logger is None:
            self._logger = logging.getLogger(self.__class__.__name__)
            self._logger.addHandler(logging.NullHandler())
        else:
            self._logger = logger

        self._entities = {}
        self._domains = set()

        if exclude is None:
            exclude = tuple()
        self._exclude = set(exclude)

        self._load()

    @staticmethod
    def _generate(len=8):
        sample = string.ascii_letters + string.digits
        return ''.join(random.sample(sample, len))

    def _load(self):
        """
        Reads PEM certificates from cache and keeps them in dictionary by
        domain names. If 'exclude' is not None, then skip certificates from
        'exclude'.
        """
        counter = 0
        try:
            with open(self.ssl_cache_path, 'r') as f:
                for line in f:

                    if ';;' not in line:
                        continue

                    domain, data = line.split(';;', 1)

                    if not domain:
                        continue

                    if domain in self._exclude:
                        self._logger.info('%s is to be excluded. Skip', domain)
                        continue

                    self._domains.add(domain)
                    self._logger.debug("Loaded %s from cache", domain)
                    self._entities[domain] = data
                    counter += 1
            self._logger.info("Loaded %d certificates from cache", counter)
        except FileNotFoundError:
            self._logger.info("No cache found. Starting from scratch")

    def get(self, as_json=False):
        if as_json:
            json.dump(sorted(self._entities.keys()), sys.stdout)
        else:
            for key in sorted(self._entities.keys()):
                print(key)

    def items(self):
        return self._entities.items()

    def add(self, filehandle):
        """
        Adds a key, a cert, and, possibly, a chain to certificate cache
        :param filehandle: obj -> open file object, stdin or regular file
        """
        entities = json.load(filehandle)
        if not isinstance(entities, (list, tuple)):
            raise TypeError("An array is expected, not a dictionary")

        domains = set()
        for entity in entities:

            if 'domain' not in entity:
                self._logger.error("Entity has no domain. Skipping.")
                continue

            if 'key' not in entity or 'certificate' not in entity:
                self._logger.error("Imcomplete certificate data. Skipping.")
                continue

            domain = entity['domain']

            if domain in self._domains:
                self._logger.warning(
                    "%s is already in cache. Skip", domain)
                continue
            if domain in domains:
                self._logger.warning(
                    "%s has already been added. Skip", domain)
                continue
            text = []
            text.append(entity['key'].encode('unicode_escape').decode('utf-8'))
            text.append(
                entity['certificate'].encode('unicode_escape').decode('utf-8'))
            if 'chain' in entity:
                text.append(
                    entity['chain'].encode('unicode_escape').decode('utf-8'))
            self._entities[domain] = ''.join(text)
            self._logger.info("Added %s to cache", domain)
            domains.add(domain)

    def save(self):
        suffix = self._generate()
        temp = os.path.extsep.join([self.ssl_cache_path, suffix])
        counter = 0

        with open(temp, 'w') as f:
            for k, v in self._entities.items():
                f.write(';;'.join([k, v]))
                if v[-1] != '\n':
                    f.write('\n')
                counter += 1

        user = pwd.getpwnam(self.cache_user)
        os.chmod(temp, 0o600)
        os.chown(temp, user.pw_uid, user.pw_gid)
        os.rename(temp, self.ssl_cache_path)
        self._logger.info("Saved %d certificates to cache", counter)


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument('-a', '--add', nargs='?', const=sys.stdin,
                        type=argparse.FileType('r'),
                        help="Append certificate to cache")
    parser.add_argument('-R', '--remove', nargs='+',
                        help="Remove specified domains from cache")
    parser.add_argument('-l', '--log', type=argparse.FileType('a'),
                        default=sys.stdout,
                        help="Logging destination. STDOUT by default")
    parser.add_argument('-V', '--verbose', action='store_true',
                        help="Make logging a bit move verbose")
    parser.add_argument('-d', '--debug', action='store_true',
                        help="Enable debug output")
    parser.add_argument('-j', '--json', action='store_true',
                        help="Output in JSON format")
    return parser.parse_args()


def set_logging(args):
    logger = logging.getLogger('im360-ssl-cache')
    if args.debug:
        logger.setLevel(logging.DEBUG)
    elif args.verbose:
        logger.setLevel(logging.INFO)
    else:
        logger.setLevel(logging.WARNING)

    ch = logging.StreamHandler(args.log)
    formatter = logging.Formatter('%(asctime)s [%(levelname)s]: %(message)s')
    ch.setFormatter(formatter)
    logger.addHandler(ch)
    return logger


def main():
    args = parse_args()
    logger = set_logging(args)
    if not any([args.add, args.remove]):
        return CertCollection(logger=logger).get(args.json)
    cc = CertCollection(logger=logger, exclude=args.remove)
    if args.add:
        cc.add(args.add)
    cc.save()


if __name__ == '__main__':
    main()
