# !/usr/bin/python3

# Built-in libraries
import gc
import logging
from pathlib import Path
import io
import json
import re
import os
import hashlib
import time
from multiprocessing import Process, cpu_count
import socket
import traceback
from concurrent.futures import ProcessPoolExecutor
from aiohttp import web
import asyncio

import oss2
import boto3

# Libraries of this project
import libs.networks as networks
import libs.preprocessing as preprocessing
import libs.postprocessing as postprocessing

CPU_COUNT = cpu_count()

BASE_DIR = Path(__file__).parent
ENV = 'live' if str(BASE_DIR).find('/alidata/www/') > -1 else 'sandbox'
ZK_DIR  = (str(BASE_DIR.parent.parent) + '/s/') if ENV == 'live' else (str(BASE_DIR) + '/s/')
with open(ZK_DIR + 'Zk.ai-ps.json') as f:
    ZK_CONF = json.loads(f.read())


LOG_DIR = str(BASE_DIR.parent.parent) + '/log/ai-ps-biz/'
today = time.strftime("%Y-%m-%d", time.localtime())

matting_log_file = LOG_DIR + '/matting.log' + today

logging.basicConfig(level=logging.INFO, filename=matting_log_file, filemode="a+", format="%(asctime)-15s %(levelname)-8s %(message)s")

def __save_content_aliyun_oss(oss_path, bytes_content, orign_img_url):
    auth = oss2.Auth(ZK_CONF['ALIYUN']['ACCESS_KEY'], ZK_CONF['ALIYUN']['SECRET_KEY'])
    bucket = oss2.Bucket(auth, endpoint = ZK_CONF['ALIYUN']['ENDPOINT'], bucket_name = ZK_CONF['ALIYUN']['BUCKET'])
    
    try:
        resp = bucket.put_object(oss_path, bytes_content)
        status_code = resp.status
        if status_code == 200:
            return {
                "result": "success",
                "img_url": ZK_CONF['ALIYUN']['OSS_DOMAIN'] + oss_path,
                "orign_img_url": orign_img_url,
                "status_code": status_code,
                "etag": resp.etag
            }
        else:
            return {
                "result": "fail",
                "orign_img_url": orign_img_url,
                "reason": "upload oss fail"
            }
    except oss2.exceptions.ServerError as e:
            return {
                "result": "fail",
                "reason": "upload oss fail",
                "orign_img_url": orign_img_url,
                "error_code": e.status,
                "request_id": e.request_id
            }

def __save_content_jcloud_oss(oss_path, bytes_content, orign_img_url):
    s3 = boto3.client(  
        's3',  
        aws_access_key_id = ZK_CONF['JCLOUD']['ACCESS_KEY'],  
        aws_secret_access_key = ZK_CONF['JCLOUD']['SECRET_KEY'],  
        endpoint_url = ZK_CONF['JCLOUD']['ENDPOINT']  
    )
    resp = s3.put_object(Bucket = ZK_CONF['JCLOUD']['BUCKET'], Key = oss_path, Body = bytes_content)
    meta_data = resp['ResponseMetadata']
    status_code = meta_data['HTTPStatusCode']
    if status_code == 200:
        return {
            "result": "success",
            "img_url": ZK_CONF['JCLOUD']['OSS_DOMAIN'] + oss_path,
            "orign_img_url": orign_img_url,
            "status_code": status_code,
            "etag": resp['ETag']
        }
    else:
        return {
            "result": "fail",
            "orign_img_url": orign_img_url,
            "reason": "upload oss fail"
        }

def __get_str_hash(data_str):
    return hashlib.md5(data_str.encode('utf-8')).hexdigest()
     
def __save_image_to_oss(pil_img, img_url, oss_base_path):
    """
    :param img: PIL image
    """
    img_byte_arr = io.BytesIO()
    pil_img.save(img_byte_arr, format='PNG')
    img_byte_str = img_byte_arr.getvalue()

    img_hash = __get_str_hash(img_url)
    oss_path = oss_base_path + img_hash + '.png'
    hostname = os.uname()[1]
    if 'jc-' in hostname:
        return __save_content_jcloud_oss(oss_path, img_byte_str, img_url)
    else:
        return __save_content_aliyun_oss(oss_path, img_byte_str, img_url)    

def __is_url(vail_str):
    regex = re.compile(
        r'^(?:http|ftp)s?://' # http:// or https://
        r'(?:(?:[A-Z0-9](?:[A-Z0-9-]{0,61}[A-Z0-9])?\.)+(?:[A-Z]{2,6}\.?|[A-Z0-9-]{2,}\.?)|' #domain...
        r'localhost|' #localhost...
        r'\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})' # ...or ip
        r'(?::\d+)?' # optional port
        r'(?:/?|[/?]\S+)$', re.IGNORECASE)
    return re.match(regex, vail_str)

def __verify_post_sign(sign, img_urls_str):
    data_hash = __get_str_hash(ZK_CONF['TOKEN'] + img_urls_str)
    logging.info("post sign is %s and token is %s and img_urs_str is %s", sign, ZK_CONF['TOKEN'], img_urls_str)
    return sign == data_hash

def process_img_urls(img_urls, oss_base_path = 'ai_matting/', model_name="u2net",
            preprocessing_method_name="bbd-fastrcnn",
            postprocessing_method_name="rtb-bnb"):
    """
    Processes the file.
    :param img_urls: 待处理图片，可以是单张，或者多张
    :param model_name: Model to use.
    :param postprocessing_method_name: Method for image preprocessing
    :param preprocessing_method_name: Method for image post-processing
    """
    model = networks.model_detect(model_name)  # Load model

    if not model:
        model_name = 'u2net' # If the model line is wrong, select the model with better quality.
        model = networks.model_detect(model_name)  # Load model

    #preprocessing_method = preprocessing.method_detect(preprocessing_method_name)
    preprocessing_method = False
    postprocessing_method = postprocessing.method_detect(postprocessing_method_name)
    process_id = os.getpid()
    if isinstance(img_urls, str):
        logging.info("start %s matting imgurl url is %s", process_id, img_urls)
        image = model.process_image(img_urls, preprocessing_method, postprocessing_method)
        logging.info("end %s matting imgulr and start upoad oss %s", process_id, img_urls)
        ret = __save_image_to_oss(image, img_urls, oss_base_path)
        gc.collect()
        return ret

    elif isinstance(img_urls , list):
        ret_list = []
        for img_url in img_urls:
            if __is_url(img_url) is False:
                ret = {
                    "result": "fail",
                    "reason": "file is not img url"
                }
                ret_list.append(ret)  
            else:
                logging.info("start %s  matting imgurl list and imgurl is %s", process_id, img_url)
                image = model.process_image(img_url, preprocessing_method, postprocessing_method)
                logging.info("end %s  matting imgurl list and imgurl is %s", process_id, img_url)
                ret = __save_image_to_oss(image, img_url, oss_base_path)
                gc.collect() 
                ret_list.append(ret)    

        return ret_list

def process_matting(matting_data):
    img_urls = matting_data['imgUrls']
    oss_base_path = 'ai_matting/' + today + '/'
    return process_img_urls(img_urls, oss_base_path)

async def do_matting(request):
    form = await request.post()
    resp = json.dumps(
        {
            "result": "fail",
            "reason": "格式解析失败"
        }   
    )
    if request.method == 'POST':
        img_urls_str = form['imgUrls']
        sign = form['sign']
        try:
            is_vaild = __verify_post_sign(sign, img_urls_str)
            if is_vaild is False:
                return web.Response(text = json.dumps(
                    {
                        "result": "fail",
                        "reason": "非法的请求内容"
                    } 
                ))

            data_json = json.loads(img_urls_str)    
            process_resp = process_matting(data_json)
            resp = json.dumps(process_resp)
            return web.Response(text = resp)
        except ValueError as e:
            return web.Response(text = resp)
    else:
        return web.Response(text = resp)

def mk_socket(host="127.0.0.1", port=9090, reuseport=False):
    sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    if reuseport:
        SO_REUSEPORT = 15
        sock.setsockopt(socket.SOL_SOCKET, SO_REUSEPORT, 1)
    sock.bind((host, port))
    return sock

async def handle(request):
    name = request.match_info.get('name', "Anonymous")
    pid = os.getpid()
    text = "{:.2f}: Hello {}! Process {} is treating you\n".format(
        time.time(), name, pid)
    #time.sleep(5)  # intentionally blocking sleep to simulate CPU load
    return web.Response(text=text)

async def start_server():
    try:
        host = "0.0.0.0"
        port=8899
        reuseport = True
        app = web.Application()
        app.router.add_get('/matting', do_matting, name="matting")
        app.router.add_post('/matting', do_matting, name="matting")
        runner = web.AppRunner(app)
        await runner.setup()
        sock = mk_socket(host, port, reuseport=reuseport)
        srv = web.SockSite(runner, sock)
        await srv.start()
        return srv, app, runner
    except Exception:
        traceback.print_exc()
        raise

async def finalize(srv, app, runner):
    sock = srv.sockets[0]
    app.loop.remove_reader(sock.fileno())
    sock.close()

    #await handler.finish_connections(1.0)
    await runner.cleanup()
    srv.close()
    await srv.wait_closed()
    await app.finish()

def init():
    loop = asyncio.get_event_loop()
    srv, app, runner = loop.run_until_complete(start_server())
    try:
        loop.run_forever()
    except KeyboardInterrupt:
        loop.run_until_complete((finalize(srv, app, runner)))


if __name__ == '__main__':
    #main()
    with ProcessPoolExecutor() as executor:
        for i in range(0, CPU_COUNT):
            executor.submit(init)