# !/usr/bin/python3

# Built-in libraries
import torch
import gc
from pathlib import Path
import json
import os
from aiohttp import web
import numpy as np
import aiohttp_cors
import time
import base64
from PIL import Image
from io import BytesIO
import os.path as osp
import torchvision.transforms as transforms
import cv2

# Libraries of this project
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 libs.matting_face.model import BiSeNet

port=8920
models_dir = Path(__file__).parent.joinpath("models/matting_face")
class AiPsMattingFace:
    def __init__(self):
        self.app_tool = AppTool('ps_matting_face.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

    def process_matting(self, img_binary_data, post_method="rtb-bnb"):
        image = Image.open(BytesIO(img_binary_data))
        width, height = image.size
        if width != height:
           raise Exception('图片宽高必须相等')

        stride = width / 512
        image512 = image.resize((512, 512), Image.BILINEAR)
        parsing_anno = self.evaluate(image512)
        transparent_image, transparent_neck_image = self.get_img_by_parse_data(image, stride, parsing_anno)

        face_image = Image.fromarray(cv2.cvtColor(transparent_image, cv2.COLOR_BGRA2RGBA))
        neck_image = Image.fromarray(cv2.cvtColor(transparent_neck_image, cv2.COLOR_BGRA2RGBA))

        if post_method == 'rtb-bnb':
            postprocessing_method = postprocessing.method_detect(post_method)
            face_image = postprocessing_method.run(None, face_image, image)
            neck_image = postprocessing_method.run(None, neck_image, image)

        face_img_url = self.save_image_to_oss(face_image)

        # 裁剪图像
        bbox = neck_image.getbbox()
        neck_image = neck_image.crop(bbox)
        neck_img_url = self.save_image_to_oss(neck_image)

        gc.collect()

        return self.app_tool.success_result({
            'face_img_url': face_img_url,
            'neck_img_url': neck_img_url,
        })

    def save_image_to_oss(self, image):
        date_h = time.strftime("%Y-%m-%d/%H", time.localtime())
        oss_base_path = 'ai_matting_face/' + date_h + '/'
        uuid = self.app_tool.uuid()
        ret = self.app_tool.save_image_to_oss(image, uuid, oss_base_path)
        return ret['img_url']

    async def do_matting_face(self, request):
        form = await request.post()
        resp = json.dumps(self.app_tool.fail_result("empty avgs"))
        if request.method != 'POST':
            return web.Response(text = resp)
        sign = None if 'sign' not in form else form['sign']
        img_bytes = None if 'imgBytes' not in form else form['imgBytes']
        is_self_ip = self.app_tool.is_self_ip(request)

        try:
            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 %s", is_self_ip)

            self.app_tool._init_log()
            process_resp = self.process_matting(base64.b64decode(img_bytes))
            self.logger.info("process_matting end adn ret %s", process_resp)
            return web.Response(text = json.dumps(process_resp))
        except ValueError as e:
            return web.Response(text = json.dumps(self.app_tool.fail_result("error")))


    def evaluate(self, image, cp='79999_iter.pth'):
        n_classes = 19
        net = BiSeNet(n_classes=n_classes)
        save_pth = models_dir.joinpath( cp)

        to_tensor = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
        ])

        if torch.cuda.is_available():
            net.cuda()
            net.load_state_dict(torch.load(save_pth))
            net.eval()
            with torch.no_grad():
                img = to_tensor(image)
                img = torch.unsqueeze(img, 0)
                img = img.cuda()
                out = net(img)[0]
                parsing = out.squeeze(0).cpu().numpy().argmax(0)
                return parsing
        else:
            state_dict = torch.load(save_pth, map_location=torch.device("cpu"))
            for key in state_dict.keys():
                # For each parameter in the state_dict, move it to the CPU device
                state_dict[key] = state_dict[key].to(torch.device("cpu"))
            net.load_state_dict(state_dict)
            net.eval()
            with torch.no_grad():
                img = to_tensor(image)
                img = torch.unsqueeze(img, 0)
                out = net(img)[0]
                parsing = out.squeeze(0).cpu().numpy().argmax(0)
                return parsing

    def get_img_by_parse_data(self, image, stride, parsing_anno):
        part_colors = [[255, 255, 255],#
                               [0, 0, 0], #脸
                               [0, 0, 0],#右眉
                               [0, 0, 0],#左眉
                               [0, 0, 0],#右眼
                               [0, 0, 0],#左眼
                               [0, 0, 0],#眼镜
                               [0, 0, 0],#右耳
                               [0, 0, 0],#左耳
                               [0, 0, 0],#耳环
                               [0, 0, 0],#鼻子
                               [0, 0, 0],#嘴
                               [0, 0, 0],#上唇
                               [0, 0, 0], #下唇
                               [255, 0, 0],#脖子
                               [255, 255, 255],#项链
                               [255, 255, 255],#衣服
                               [0, 0, 0],#头发
                               [0, 0, 0],#帽子
                               [255, 255, 255],
                               [255, 255, 255],
                               [255, 255, 255], #
                               [255, 255, 255],#
                               [255, 255, 255] #
                              ]

        im = np.array(image)
        vis_im = im.copy().astype(np.uint8)
        vis_parsing_anno = parsing_anno.copy().astype(np.uint8)
        vis_parsing_anno = cv2.resize(vis_parsing_anno, None, fx=stride, fy=stride, interpolation=cv2.INTER_NEAREST)
        vis_parsing_anno_color = np.zeros((vis_parsing_anno.shape[0], vis_parsing_anno.shape[1], 3)) + 255

        num_of_class = np.max(vis_parsing_anno)
        for pi in range(1, num_of_class + 1):
            index = np.where(vis_parsing_anno == pi)
            vis_parsing_anno_color[index[0], index[1], :] = part_colors[pi]

        vis_parsing_anno_color = vis_parsing_anno_color.astype(np.uint8)
        vis_im2 = cv2.addWeighted(cv2.cvtColor(vis_im, cv2.COLOR_RGB2BGR), 1, vis_parsing_anno_color, 0, 0)

        transparent_image = np.zeros((vis_parsing_anno_color.shape[0], vis_parsing_anno_color.shape[1], 4), dtype=np.uint8)
        transparent_neck_image = np.zeros((vis_parsing_anno_color.shape[0], vis_parsing_anno_color.shape[1], 4), dtype=np.uint8)

        # 找到不是白色的像素索引
        non_white_indices = np.all(vis_parsing_anno_color != [255, 255, 255], axis=-1)
        # 找到脖子像素索引
        neck_indices = np.all(vis_parsing_anno_color == [255, 0, 0], axis=-1)
        # 将不是白色的像素复制到 transparent_image，并将透明通道设置为不透明（255）
        transparent_image[non_white_indices, :3] = vis_im2[non_white_indices]
        transparent_image[non_white_indices, 3] = 255
        # 将脖子像素复制到 transparent_neck_image，并将透明通道设置为不透明（255）
        transparent_neck_image[neck_indices, :3] = vis_im2[neck_indices]
        transparent_neck_image[neck_indices, 3] = 255

        return [transparent_image, transparent_neck_image]

    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', port))
        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_face', self.do_matting_face, name="matting_face")
        app.router.add_post('/matting_face', self.do_matting_face, name="matting_face")

        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(client_max_size=1024**2*4)
        app.router.add_get('/matting_face', self.do_matting_face, name="matting_face")
        app.router.add_post('/matting_face', self.do_matting_face, name="matting_face")

        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 = port)


if __name__ == '__main__':
    
    cpu_cnt = cpu_count()
    process_cnt =  1 if cpu_cnt == 1 else int(cpu_cnt / 2)
    AiPs = AiPsMattingFace()
    if torch.cuda.is_available():
        AiPs.main()
    else:
        AiPs.main_multi(process_cnt)