# !/usr/bin/python3

# Built-in libraries
import gc
from pathlib import Path
import json
import os
from aiohttp import web
import aiohttp_cors
import time
import base64

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

from multiprocessing import Process, cpu_count
from socket import SOL_SOCKET, SO_REUSEADDR, socket

from libs.app_tool import AppTool

from carvekit.ml.wrap.tracer_b7 import TracerUniversalB7
from carvekit.api.interface import Interface
class AiPsMatting:
    def __init__(self):
        self.app_tool = AppTool('ps_matting.log')
        self.BASE_DIR = Path(__file__).parent
        self.ZK_CONF = self.app_tool.zk()

        self.hostname = os.uname()[1]
        self.logger = self.app_tool.logger

        self.module_tracer_b7_seg_net = None
        self.interface_tracer_b7 = None
        self.module_u2net_sg_net = None

    def init_seg_net(self, module_name, post_method = None):
        if module_name == 'tracer_b7' and isinstance(self.module_tracer_b7_seg_net, TracerUniversalB7) is False:
            seg_net = TracerUniversalB7(
                device= 'cpu',
                batch_size=5,
                input_image_size=640,
                fp16=False,
            )
            self.module_tracer_b7_seg_net = seg_net
            self.interface_tracer_b7 = Interface(
                pre_pipe=None,
                post_pipe=post_method,
                seg_pipe=seg_net,
                device='cpu',
            )  
            
        if module_name == 'u2net' and self.module_u2net_sg_net is None:
            self.module_u2net_sg_net = networks.model_detect('u2net')
                

    def process_img_by_bytes(self, img_bytes, post_method="rtb-bnb", seg_mode = 'u2net', denoise = None):
        if post_method == 'rtb-bnb':
            postprocessing_method = postprocessing.method_detect(post_method)
        else:
            postprocessing_method = None

        if 'tracer_b7' == seg_mode:
            self.init_seg_net('tracer_b7', postprocessing_method)
            pil_img = self.module_tracer_b7_seg_net.load_image(img_bytes)
            image = self.interface_tracer_b7([pil_img])[0]
        else:
            self.init_seg_net('u2net')
            image = self.module_u2net_sg_net.process_image(img_bytes, preprocessing = False, postprocessing = postprocessing_method)   


        date_h = time.strftime("%Y-%m-%d/%H", time.localtime())
        oss_base_path = 'ai_matting/' + date_h + '/'
        uuid = self.app_tool.uuid()
        ret = self.app_tool.save_image_to_oss(image, uuid, oss_base_path, denoise = denoise)
        ret['seg'] = seg_mode
        gc.collect()
        return ret

    def process_img_urls(self, img_urls, oss_base_path = 'ai_matting/', seg_mode="u2net", post_method="rtb-bnb", denoise=None):
        """
        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
        """
        if post_method == 'rtb-bnb':
            postprocessing_method = postprocessing.method_detect(post_method)
        else:
            postprocessing_method = None

        if 'tracer_b7' == seg_mode:
            self.init_seg_net('tracer_b7', postprocessing_method)
        else:
            self.logger.info("start model_detect and modename is u2net")
            self.init_seg_net('u2net')
            self.logger.info("end model_detect and modename is u2net")

        self.logger.info("start matting imgurl url is %s", img_urls)
        if isinstance(img_urls, str):
            if self.app_tool.uri_validator(img_urls) is False:
                return self.app_tool.fail_result("file is not img url")
            
            self.logger.info("start matting imgurl list and imgurl is %s and seg network is %s", img_urls, seg_mode)
            if 'tracer_b7' == seg_mode:
                pil_img = self.module_tracer_b7_seg_net.load_image(img_urls)
                image = self.interface_tracer_b7([pil_img])[0]
            else:
                image = self.module_u2net_sg_net.process_image(img_urls, preprocessing = False, postprocessing = postprocessing_method) 
            
            self.logger.info("end matting imgulr and start upoad oss %s", img_urls)
            ret = self.app_tool.save_image_to_oss(image, img_urls, oss_base_path, denoise = denoise)
            ret['seg'] = seg_mode
            gc.collect()
            return ret

        elif isinstance(img_urls , list):
            ret_list = []
            for img_url in img_urls:
                if self.app_tool.uri_validator(img_url) is False:
                    ret = self.app_tool.fail_result("file is not img url")
                    ret_list.append(ret)
                    continue

                self.logger.info("start matting imgurl list and imgurl is %s and seg network is %s", img_url, seg_mode)
                if 'tracer_b7' == seg_mode:
                    pil_img = self.module_tracer_b7_seg_net.load_image(img_url)
                    image = self.interface_tracer_b7([pil_img])[0]
                else:
                    image = self.module_u2net_sg_net.process_image(img_url, preprocessing = False, postprocessing = postprocessing_method) 
                self.logger.info("end matting imgurl list and imgurl is %s", img_url)

                ret = self.app_tool.save_image_to_oss(image, img_url, oss_base_path, denoise = denoise)
                ret['seg'] = seg_mode
                gc.collect() 
                ret_list.append(ret)   
            
            return ret_list

    def process_matting(self, img_urls, post_method="rtb-bnb", seg_mode='u2net', denoise=None):
        date_h = time.strftime("%Y-%m-%d/%H", time.localtime())
        ext = 'b7/' if seg_mode in ['tracer_b7', 'b7'] else ''
        oss_base_path = 'ai_matting/' + date_h + '/' + ext
        self.logger.info("process_matting and oss_base_path is %s", oss_base_path)
        return self.process_img_urls(img_urls, oss_base_path, seg_mode=seg_mode, post_method=post_method,denoise=denoise)

    async def do_matting(self, request):
        load1, load5, load15 = os.getloadavg()
        if load1 > 15:
            error_log = "server is busy now hostname {}, load1 {}".format(self.hostname, load1)
            self.logger.error(error_log)
            resp = json.dumps(self.app_tool.fail_result(error_log))
            return web.Response(text = resp, status=503)

        form = await request.post()
        resp = json.dumps(self.app_tool.fail_result("empty avgs"))
        if request.method == 'POST':
            img_urls_str = None if 'imgUrls' not in form else form['imgUrls']
            sign = None if 'sign' not in form else form['sign']
            is_bytes = None if 'isBytes' not in form else form['isBytes']
            img_bytes = None if 'imgBytes' not in form else form['imgBytes']
            denoise = None if 'denoise' not in form else form['denoise']
            is_self_ip = self.app_tool.is_self_ip(request)
            ##segMode form
            if 'segMode' in form and form['segMode'] in ['b7', 'tracer_b7']:
                seg_mode = 'tracer_b7'
            else:
                seg_mode = 'u2net'

            if 'postMethod' in form and form['postMethod'] in ['rtb-bnb', 'rtb-bnb2']:
                post_method = form['postMethod']
            else:
                post_method = 'rtb-bnb' if seg_mode == 'u2net' else None

            try:
                self.logger.info("start process and seg_mode is %s and post_method %s", seg_mode, post_method)
                if is_bytes is not None and img_bytes is not None:
                    is_vaild = self.app_tool.verify_post_sign(sign, img_bytes)
                    self.logger.info("__verify_post_sign is %s and sign is %s" % (is_vaild, sign))
                    if is_vaild is False and is_self_ip is False:
                        return web.Response(text = json.dumps(self.app_tool.fail_result("verify sign fail for bytes")))

                    self.logger.info("start process img for data and is_self_ip is {is_self_ip}")
                    process_resp = self.process_img_by_bytes(base64.b64decode(img_bytes), post_method=post_method, seg_mode = seg_mode, denoise = denoise)
                    self.logger.info("process_matting end adn ret %s", process_resp)
                else:
                    is_vaild = self.app_tool.verify_post_sign(sign, img_urls_str)
                    self.logger.info("__verify_post_sign is %s", is_vaild)
                    if is_vaild is False and is_self_ip is False:
                        return web.Response(text = json.dumps(self.app_tool.fail_result("verify sign fail")))

                    data_json = self.app_tool.check_is_json(img_urls_str)
                    if data_json is False:
                        data_json = img_urls_str  

                    self.app_tool._init_log()
                    self.logger.info("process_matting start and data_json is %s", data_json)
                    process_resp = self.process_matting(data_json, post_method=post_method, seg_mode=seg_mode, denoise = denoise)
                    self.logger.info("process_matting end")

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

    async def del_object(self, request):
        form = await request.post()
        resp = json.dumps(
            {
                "result": "fail",
                "reason": "格式解析失败"
            }   
        )
        if request.method == 'POST':
            object_key = form['objectKey']
            sign = form['sign']
            try:
                is_vaild = self.app_tool.verify_post_sign(sign, object_key)
                self.logger.info("del objet obj key is %s and vaild is %s", object_key, is_vaild)
                if is_vaild is False:
                    return web.Response(text = json.dumps(
                        {
                            "result": "fail",
                            "reason": "非法的请求参数"
                        } 
                    ))
                
                resp = self.app_tool.del_oss_object(object_key)
                resp = json.dumps(resp)
                return web.Response(text = resp)
            except ValueError as e:
                return web.Response(text = resp)
        else:
            return web.Response(text = resp)

    def serve_multiple(self, app, workers):
        sock = socket()
        sock.setsockopt(SOL_SOCKET, SO_REUSEADDR, 1)
        sock.bind(('0.0.0.0', 8899))
        sock.set_inheritable(True)

        processes = []
        for i in range(workers):
            process = Process(target=web.run_app, name=f'worker-{i}', kwargs=dict(app=app, sock=sock))
            process.daemon = True
            process.start()
            processes.append(process)

        try:
            for process in processes:
                process.join()
        except KeyboardInterrupt:
            pass
        finally:
            for process in processes:
                process.terminate()
            sock.close()
            
    def main_multi(self, process_cnt):
        app = web.Application(client_max_size=1024**2*4)
        app.router.add_get('/matting', self.do_matting, name="matting")
        app.router.add_post('/matting', self.do_matting, name="matting")

        app.router.add_get('/del_object', self.del_object, name="del_object")
        app.router.add_post('/del_object', self.del_object, name="del_object")

        cors = aiohttp_cors.setup(app, defaults={
                "*": aiohttp_cors.ResourceOptions(
                allow_credentials=True,
                expose_headers="*",
                allow_headers="*",
            )
        })
        
        if 'ay-pdddz-sb' in self.hostname:
            process_cnt = 1
            
        self.serve_multiple(app, process_cnt)

    def main(self):
        app = web.Application()
        app.router.add_get('/matting', self.do_matting, name="matting")
        app.router.add_post('/matting', self.do_matting, name="matting")

        app.router.add_get('/del_object', self.del_object, name="del_object")
        app.router.add_post('/del_object', self.del_object, name="del_object")
        web.run_app(app, port = 8899)


if __name__ == '__main__':
    cpu_cnt = cpu_count()
    process_cnt =  1 if cpu_cnt == 1 else int(cpu_cnt / 2)
    AiPs = AiPsMatting()
    
    AiPs.main_multi(process_cnt)