#!/usr/bin/env python3
from flask import Flask
from flask import Response
from flask import redirect
from flask import render_template
from flask import request
from flask import url_for
from mirascan.scanner import Scanner
import argparse
import json
import mimetypes
import os
import socket
import urllib.parse
import yaml
import copy


app = Flask(__name__)

errors_count = 0
errors_limit = 10

def dictDeepMerge(target, *args):
    if len(args) > 1:
        for obj in args:
            dictDeepMerge(target, obj)
        return target
    obj = args[0]
    if not isinstance(obj, dict):
        return obj
    for k, v in obj.items():
        if k in target and isinstance(target[k], dict):
            dictDeepMerge(target[k], v)
        else:
            target[k] = copy.deepcopy(v)
    return target

@app.errorhandler(500)
def handle(e):
    global errors_count
    global errors_limit
    if errors_count < errors_limit:
        errors_count += 1
        scanner.reconnect()
        return Response('Please, try again later', status=500)
    else:
        return Response('Something definetly went wrong', status=500)


@app.route('/scan', methods=['POST'])
def scan():
    if request.form.get('host'):
        ids = dict()
        hosts = [x.strip() for x in request.form.get('host').split(',')]
        for host in hosts:
            try:
                host = socket.gethostbyname(host)
            except socket.gaierror:
                ids[host] = "Host couldn't be resolved"
                continue

            resp = scanner.Scan(host)
            ids[host] = resp['id']

        ref = request.headers.get('Referer')
        if urllib.parse.urlparse(ref).path == url_for('dashboard'):
            return redirect(url_for('dashboard'))
        return Response(json.dumps(ids),
                        status=201,
                        mimetype='text/json')
    return Response("host data wasn't found in your request", status=400)


@app.route('/status/<taskid>')
def status(taskid):
    return Response(json.dumps(scanner.GetTaskStatus(taskid)),
                    status=200,
                    mimetype='text/json')


@app.route('/report/<taskid>/<fmt>')
def ftmreport(taskid, fmt):
    if fmt.lower() in ['arf', 'xml']:
        mime = mimetypes.types_map['.xml']
    elif fmt.lower() in ['cpe', 'nbe', 'txt']:
        mime = mimetypes.types_map['.txt']
    elif fmt.lower() in ['html', 'csv', 'latex', 'pdf']:
        mime = mimetypes.types_map['.'+fmt.lower()]
    else:
        return Response('Unknown format', status=400)
    return Response(scanner.GetReport(taskid, rformat=fmt.upper()),
                    status=200, mimetype=mime)


@app.route('/report/merged/<reportid>/<fmt>')
def mfmtreport(reportid, fmt):
    if fmt.lower() in ['arf', 'xml']:
        mime = mimetypes.types_map['.xml']
    elif fmt.lower() in ['cpe', 'nbe', 'txt']:
        mime = mimetypes.types_map['.txt']
    elif fmt.lower() in ['html', 'csv', 'latex', 'pdf']:
        mime = mimetypes.types_map['.'+fmt.lower()]
    elif fmt.lower() == 'json':
        return Response(scanner.GetCvesFromMergedReport(reportid),
                        status=200, mimetype='application/json')
    else:
        return Response('Unknown format', status=400)
    return Response(scanner._GetReport(reportid, rformat=fmt.upper()),
                    status=200, mimetype=mime)


@app.route('/report/json/<taskid>')
def getJsonReport(taskid):
    return Response(scanner.GetCvesFromMergedReport(taskid),
                    status=200, mimetype='application/json')


@app.route('/report/merge/<tasklabel>')
def mergeReport(tasklabel):
    username = request.form.get('username',
                                config['credential']['username'])
    key = request.form.get('key', config['credential']['key'])
    ovalpath = config['oscap']['oval']['path']
    ovaldef = request.form.get('oval',
                               config['oscap']['oval']['def'])
    localoval = '/'.join([ovalpath, os.path.basename(ovaldef)])
    localoval = os.path.normpath(localoval)
    scanner.UpdateResults(tasklabel, localoval, username=username, key=key)
    return Response("Done", status=200)


@app.route('/update')
def updateAll():
    if request.args.get('oval', False):
        ovalpath = config['oscap']['oval']['path']
        ovaldef = config['oscap']['oval']['def']
        localoval = '/'.join([ovalpath, os.path.basename(ovaldef)])
        localoval = os.path.normpath(localoval)
        res = scanner.UpdateData(oval=True,
                                 ovalurl=config['app'].get('ovalurl'),
                                 ovalpath=localoval)
    else:
        res = scanner.UpdateData(scap=True, feed=True, cert=True)

    if not res:
        return Response("Something went wrong", status=520)
    return Response("Ok", status=202)


@app.route('/list/tasks')
def list_active():
    if request.args.get('active', False):
        return Response(json.dumps(scanner.GetTaskList(active=True)),
                        status=200,
                        mimetype='text/json')
    return Response(json.dumps(scanner.GetTaskList()),
                    status=200,
                    mimetype='text/json')


@app.route('/dashboard')
def dashboard():
    finished = scanner._GetFilteredTasks(status='Done')
    active = scanner._GetFilteredTasks(status='Running', count=5)
    stopped = scanner._GetFilteredTasks(status='Stopped', count=5)
    return render_template('dashboard.jinja',
                           active=active,
                           finished=finished,
                           stopped=stopped,
                           baseurl=config['app']['ompurl'])


@app.route('/help')
def help():
    txt = '''Usage:
POST /scan {host=1.2.3.4} - start new scan of <host>,
            return json with taskid and status.

            Use with curl to start scan 8.8.8.8:
                curl -d 'host=8.8.8.8' http://<rest.host>/scan
GET /status/<taskid> - get status of task <taskid>
GET /report/<taskid>/<fmt> - get report of task <taskid>
            where fmt is report format. Currently
            supported formats are arf, cpe, csv, xml,
            latex, txt, html, nbe or json
GET /list/tasks - return list of tasks
            Supports ?active=True key to list active only tasks'''
    return Response(txt, status=200, mimetype=mimetypes.types_map['.txt'])


if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='Mirantis scanner')
    parser.add_argument('-c', '--config', dest='configfile', action='store',
                        help='Path to config file',
                        default='/etc/mirascan.yaml')
    args = parser.parse_args()

    configfile = args.configfile

    config = dict()
    config['oscap'] = dict()
    config['oscap']['oval'] = dict()
    config['oscap']['oval']['path'] = '/usr/share/oscap/oval'
    config['oscap']['oval']['def'] = 'ubuntu-16.04-oval.xml'

    config['openvas'] = dict()
    config['openvas']['connection'] = dict()
    config['openvas']['connection']['username'] = 'admin'
    config['openvas']['connection']['password'] = 'admin'
    config['openvas']['connection']['host'] = '127.0.0.1'
    config['openvas']['connection']['port'] = 9390

    config['credential'] = dict()
    config['credential']['username'] = 'root'
    config['credential']['key'] = None
    config['credential']['keypass'] = None
    config['credential']['port'] = 22

    config['app'] = dict()
    config['app']['listen'] = dict()
    config['app']['listen']['host'] = '0.0.0.0'
    config['app']['listen']['port'] = 54321
    config['app']['ompurl'] = 'https://127.0.0.1:4000/omp'

    if os.path.isfile(configfile):
        config = dictDeepMerge(config, yaml.safe_load(open(configfile).read()))

    scanner = Scanner(username=config['openvas']['connection']['username'],
                      passwd=config['openvas']['connection']['password'],
                      host=config['openvas']['connection']['host'],
                      port=config['openvas']['connection']['port'])

    scanner.cbhost = config['app']['listen']['host']
    scanner.cbport = config['app']['listen']['port']

    scanner.sshuser = config['credential']['username']
    scanner.sshkey = config['credential']['key']
    scanner.sshkeypass = config['credential'].get('keypass')

    app.run(host=config['app']['listen']['host'],
            port=config['app']['listen']['port'])

