import base64
import oss2
import boto3
import hashlib
import logging
from pathlib import Path
import io
import json
import os
import imghdr
import re
import time
import uuid
import numpy as np
from PIL import Image
import cv2
import requests
import sys
import tempfile
from urllib.request import urlopen, Request
from urllib.parse import urlparse
import shutil
from ipaddress import ip_address

from cairosvg import svg2png

try:
    from tqdm.auto import tqdm  # automatically select proper tqdm submodule if available
except ImportError:
    try:
        from tqdm import tqdm
    except ImportError:
        # fake tqdm if it's not installed
        class tqdm(object):  # type: ignore[no-redef]

            def __init__(self, total=None, disable=False,
                         unit=None, unit_scale=None, unit_divisor=None):
                self.total = total
                self.disable = disable
                self.n = 0
                # ignore unit, unit_scale, unit_divisor; they're just for real tqdm

            def update(self, n):
                if self.disable:
                    return

                self.n += n
                if self.total is None:
                    sys.stderr.write("\r{0:.1f} bytes".format(self.n))
                else:
                    sys.stderr.write("\r{0:.1f}%".format(100 * self.n / float(self.total)))
                sys.stderr.flush()

            def close(self):
                self.disable = True

            def __enter__(self):
                return self

            def __exit__(self, exc_type, exc_val, exc_tb):
                if self.disable:
                    return

                sys.stderr.write('\n')

class AppTool:
    def __init__(self, log_name = 'ai_ps_defalut.log'):
        self.BASE_DIR = Path(__file__).parent.parent
        self.ENV = 'live' if str(self.BASE_DIR).find('/alidata/www/') > -1 else 'sandbox'
        self.ZK_CONF = {}

        ZK_DIR  = (str(self.BASE_DIR.parent.parent) + '/s/') if self.ENV == 'live' else (str(self.BASE_DIR) + '/s/')
        with open(ZK_DIR + 'Zk.ai-ps.json') as f:
            self.ZK_CONF = json.loads(f.read())
        
        self.hostname = os.uname()[1]
        self.logger = logging
        self.inpaint_model = None
        self.matting_modle = None
        self._init_conf()
        self._init_log(log_name)

        
    def zk(self):
        return self.ZK_CONF

    def _init_log(self, log_name = 'matting.log'):
        LOG_DIR = str(self.BASE_DIR.parent.parent) + '/log/ai-ps-biz/'
        Path(LOG_DIR).mkdir(parents=True, exist_ok=True)

        today = time.strftime("%Y-%m-%d", time.localtime())
        matting_log_file = LOG_DIR + '/' + log_name + '.' + today

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


    def _init_conf(self):
        if 'jc-' in self.hostname:
            self.OSS_ACCESS_KEY = self.ZK_CONF['JCLOUD']['ACCESS_KEY']
            self.OSS_SECRET_KEY = self.ZK_CONF['JCLOUD']['SECRET_KEY']
            self.OSS_ENDPOINT = self.ZK_CONF['JCLOUD']['ENDPOINT']
            self.OSS_BUCKET = self.ZK_CONF['JCLOUD']['BUCKET']
            self.OSS_INTERNAL_ENDPOINT = self.ZK_CONF['JCLOUD']['INTERNAL_ENDPOINT']
            self.OSS_DOMAIN =  self.ZK_CONF['JCLOUD']['OSS_DOMAIN'] 
            self.OSS_INTERNAL_DOMAIN =  self.ZK_CONF['JCLOUD']['OSS_INTERNAL_DOMAIN'] 
            self.IDC = 'jcloud'
        elif 'ay-' in self.hostname:
            self.OSS_ACCESS_KEY = self.ZK_CONF['ALIYUN']['ACCESS_KEY']
            self.OSS_SECRET_KEY = self.ZK_CONF['ALIYUN']['SECRET_KEY']
            self.OSS_ENDPOINT = self.ZK_CONF['ALIYUN']['ENDPOINT']
            self.OSS_BUCKET = self.ZK_CONF['ALIYUN']['BUCKET']
            self.OSS_INTERNAL_ENDPOINT = self.ZK_CONF['ALIYUN']['INTERNAL_ENDPOINT']
            self.OSS_DOMAIN =  self.ZK_CONF['ALIYUN']['OSS_DOMAIN'] 
            self.OSS_INTERNAL_DOMAIN =  self.ZK_CONF['ALIYUN']['OSS_INTERNAL_DOMAIN']
            self.IDC = 'aliyun'
        else:
            self.OSS_ACCESS_KEY = self.ZK_CONF['JCLOUD']['ACCESS_KEY']
            self.OSS_SECRET_KEY = self.ZK_CONF['JCLOUD']['SECRET_KEY']
            self.OSS_ENDPOINT = self.ZK_CONF['JCLOUD']['ENDPOINT']
            self.OSS_BUCKET = self.ZK_CONF['JCLOUD']['BUCKET']
            self.OSS_DOMAIN =  self.ZK_CONF['JCLOUD']['OSS_DOMAIN'] 
            self.OSS_INTERNAL_ENDPOINT = self.ZK_CONF['JCLOUD']['INTERNAL_ENDPOINT']
            self.OSS_INTERNAL_DOMAIN =  self.ZK_CONF['JCLOUD']['OSS_INTERNAL_DOMAIN'] 
            self.IDC = 'other'    

    def save_content_aliyun_oss(self, object_key, object_content, orign_img_url):
        endpoint_url = self.OSS_ENDPOINT if self.IDC == 'other' else self.OSS_INTERNAL_ENDPOINT
        auth = oss2.Auth(self.OSS_ACCESS_KEY, self.OSS_SECRET_KEY)
        bucket = oss2.Bucket(auth, endpoint = endpoint_url, bucket_name = self.OSS_BUCKET)
        
        try:
            resp = bucket.put_object(object_key, object_content)
            status_code = resp.status
            self.logger.info("save aliyun_oss ret %s, and object key is %s ", status_code, object_key)
            if status_code == 200:
                return self.success_result({
                    "img_url": self.OSS_DOMAIN + object_key,
                    "img_internal_url": self.OSS_INTERNAL_DOMAIN + object_key,
                    "object_key": object_key,
                    "orign_img_url": orign_img_url,
                    "status_code": status_code,
                    "hostname": self.hostname,
                    "etag": resp.etag
                })

            else:
                return self.fail_result({
                    "orign_img_url": orign_img_url,
                    "reason": "upload oss fail",
                    "hostname": self.hostname
                })

        except oss2.exceptions.ServerError as e:
                return self.fail_result({
                    "reason": "upload oss fail",
                    "orign_img_url": orign_img_url,
                    "error_code": e.status,
                    "request_id": e.request_id,
                    "hostname": self.hostname
                })

    def save_content_jcloud_oss(self, object_key, object_content, orign_img_url):
        endpoint_url = self.OSS_ENDPOINT if self.IDC == 'other' else self.OSS_INTERNAL_ENDPOINT
        s3 = boto3.client(  
            's3',  
            aws_access_key_id = self.OSS_ACCESS_KEY,  
            aws_secret_access_key = self.OSS_SECRET_KEY,  
            endpoint_url = endpoint_url 
        )
        resp = s3.put_object(Bucket = self.OSS_BUCKET, Key = object_key, Body = object_content)
        meta_data = resp['ResponseMetadata']
        status_code = meta_data['HTTPStatusCode']
        self.logger.info("save jcloud_oss ret %s, and object key is %s ", status_code, object_key)
        if status_code == 200:
            return self.success_result({
                "img_url": self.OSS_DOMAIN + object_key,
                "img_internal_url": self.OSS_INTERNAL_DOMAIN + object_key,
                "object_key": object_key,
                "orign_img_url": orign_img_url,
                "status_code": status_code,
                "hostname": self.hostname,
                "etag": resp['ETag']
            })

        else:
            return self.fail_result({
                "orign_img_url": orign_img_url,
                "hostname": self.hostname,
                "reason": "upload oss fail"
            })

    def upload_file_aliyun_oss(self, object_key, file_path, new_bucket_name = None):
        endpoint_url = self.OSS_ENDPOINT if self.IDC == 'other' else self.OSS_INTERNAL_ENDPOINT
        oss_domain = self.OSS_DOMAIN
        bucket_name = self.OSS_BUCKET

        if new_bucket_name is not None:
            bucket_name = new_bucket_name
            parts = urlparse(self.OSS_ENDPOINT)
            oss_domain = ('http://%s.%s/' % (bucket_name,  parts.hostname))

        auth = oss2.Auth(self.OSS_ACCESS_KEY, self.OSS_SECRET_KEY)
        bucket = oss2.Bucket(auth, endpoint = endpoint_url, bucket_name = bucket_name)
        
        try:
            resp = bucket.put_object_from_file(object_key, file_path)
            status_code = resp.status
            self.logger.info("save aliyun_oss ret %s, and object key is %s and bucket_name is %s and oss domain is %s ", status_code, object_key, bucket_name, oss_domain)
            if status_code == 200:

                return self.success_result({
                    "subset_url": oss_domain + object_key,
                    "object_key": object_key,
                    "status_code": status_code,
                    "hostname": self.hostname,
                    "etag": resp.etag,
                    "bucket": bucket_name
                })
            else:
                return self.fail_result({
                    "reason": "upload oss fail",
                    "hostname": self.hostname
                })

        except oss2.exceptions.ServerError as e:
                return self.fail_result({
                    "reason": "upload oss fail",
                    "error_code": e.status,
                    "request_id": e.request_id,
                    "hostname": self.hostname
                })

    def upload_file_jcloud_oss(self, object_key, file_path, new_bucket_name = None):
        endpoint_url = self.OSS_ENDPOINT if self.IDC == 'other' else self.OSS_INTERNAL_ENDPOINT
        oss_domain = self.OSS_DOMAIN
        bucket_name = self.OSS_BUCKET

        if new_bucket_name is not None:
            bucket_name = new_bucket_name
            parts = urlparse(self.OSS_ENDPOINT)
            oss_domain = ('http://%s.%s/' % (bucket_name,  parts.hostname))

        s3 = boto3.client(  
            's3',  
            aws_access_key_id = self.OSS_ACCESS_KEY,  
            aws_secret_access_key = self.OSS_SECRET_KEY,  
            endpoint_url = endpoint_url 
        )

        try:
            s3.upload_file(Bucket = bucket_name, Key = object_key, Filename = file_path)
            self.logger.info("save jcloud_oss end, and object key is %s ", object_key)
            return self.success_result({
                "subset_url": oss_domain + object_key,
                "object_key": object_key,
                "hostname": self.hostname,
                "bucket": bucket_name
            })

        except IOError:
            return self.fail_result({
                "hostname": self.hostname,
                "reason": "upload oss fail"
            })
            
    def del_oss_object(self, object_key):
        if self.IDC == 'aliyun':
            auth = oss2.Auth(self.OSS_ACCESS_KEY, self.OSS_SECRET_KEY)
            bucket = oss2.Bucket(auth, endpoint = self.OSS_ENDPOINT, bucket_name = self.OSS_BUCKET)
            resp = bucket.delete_object(object_key)
            status_code = resp.status
            if status_code == 200:
                return self.success_result()
            else:
                return self.fail_result("delete oss fail") 
        else:
            s3 = boto3.client(  
                's3',  
                aws_access_key_id = self.OSS_ACCESS_KEY,  
                aws_secret_access_key = self.OSS_SECRET_KEY,  
                endpoint_url = self.OSS_ENDPOINT  
            )

            resp = s3.delete_object(Bucket = self.OSS_BUCKET, Key = object_key)
            meta_data = resp['ResponseMetadata']
            status_code = meta_data['HTTPStatusCode']
            if status_code == 204 or status_code == 200:
                return self.success_result()
            else:
                return self.fail_result()

    def save_image_to_oss_by_bytes(self, img_byte_str, img_url, oss_base_path, img_format = 'PNG'):
        oss_path = oss_base_path + self.uuid() + '.' + img_format.lower()

        if self.IDC == 'aliyun':
            return self.save_content_aliyun_oss(oss_path, img_byte_str, img_url)   
        else:
            return self.save_content_jcloud_oss(oss_path, img_byte_str, img_url)

    def save_image_to_oss(self, pil_img, img_url, oss_base_path, img_format = 'PNG', denoise=None):
        """
        :param img: PIL image
        """
        if img_format in ['JPG', 'JPEG']:
            pil_img = pil_img.convert("RGB")
            format = 'JPG'
        else:
            format = img_format

        if denoise is not None and format == 'PNG':
            pil_img = self.image_auto_center_denosie(pil_img)

        img_byte_arr = io.BytesIO()
        pil_img.save(img_byte_arr, format=img_format)
        img_byte_str = img_byte_arr.getvalue()
        
        # from urllib import parse
        # img_path = parse.urlsplit(img_url).hostname + parse.urlsplit(img_url).path

        #img_hash = self.get_str_hash(img_url)
        img_hash = self.uuid()
        oss_path = oss_base_path + img_hash + '.' + format.lower()

        if self.IDC == 'aliyun':
            return self.save_content_aliyun_oss(oss_path, img_byte_str, img_url)   
        else:
            return self.save_content_jcloud_oss(oss_path, img_byte_str, img_url)
    
    def upload_to_oss_by_file(self, local_file, oss_path, bucket = None):
        if self.IDC == 'aliyun':
            return self.upload_file_aliyun_oss(oss_path, local_file, bucket)   
        else:
            return self.upload_file_jcloud_oss(oss_path, local_file, bucket)

    def image_auto_center_denosie(self, pil_image, area_threshold = 1500):
        """
        去除图像噪点，并且图像居中。
        """
        # 将 PIL 图像转换为 OpenCV 格式
        image = cv2.cvtColor(np.array(pil_image), cv2.COLOR_RGBA2BGRA)  # 如果图像是RGB模式

        # 提取透明通道
        alpha_channel = image[:, :, 3]

        # 连通组件分析
        num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(alpha_channel)

        # 设置阈值，筛选孤立的区域
        min_area_threshold = area_threshold  # 调整此阈值以适应你的图像和需求
        for label in range(1, num_labels):
            area = stats[label, cv2.CC_STAT_AREA]
            if area < min_area_threshold:
                alpha_channel[labels == label] = 0

        # 找到非透明区域的边界框
        alpha_channel = image[:, :, 3]
        non_transparent_pixels = np.where(alpha_channel != 0)
        min_row, max_row = np.min(non_transparent_pixels[0]), np.max(non_transparent_pixels[0])
        min_col, max_col = np.min(non_transparent_pixels[1]), np.max(non_transparent_pixels[1])

        # 计算需要平移的像素数量
        image_center_x = image.shape[1] // 2
        image_center_y = image.shape[0] // 2
        shift_x = image_center_x - (max_col + min_col) // 2
        shift_y = image_center_y - (max_row + min_row) // 2

        # 创建新的图像，大小与原图像相同
        new_image = np.zeros_like(image)

        # 计算平移后的非透明区域位置
        new_min_row = min_row + shift_y
        new_max_row = max_row + shift_y
        new_min_col = min_col + shift_x
        new_max_col = max_col + shift_x

        # 将原图像的非透明区域放置到新图像的相应位置
        new_image[new_min_row:new_max_row, new_min_col:new_max_col] = image[min_row:max_row, min_col:max_col]

        # 保存新图像
        image_pil = Image.fromarray(new_image)
        image_pil = Image.fromarray(cv2.cvtColor(new_image, cv2.COLOR_BGRA2RGBA))

        return image_pil

    def build_post_sign(self, post_data_str):
        return self.get_str_hash(self.ZK_CONF['TOKEN'] + post_data_str)

    def get_str_hash(self, data_str):
        if isinstance(data_str, bytes):
            return hashlib.md5(data_str).hexdigest()
        else:
            return hashlib.md5(data_str.encode('utf-8')).hexdigest()

    def uri_validator(self, url):
        try:
            result = urlparse(url)
            return all([result.scheme, result.netloc])
        except:
            return False

    def is_url(self, 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(self, sign, post_data_str):
        data_hash = self.get_str_hash(self.ZK_CONF['TOKEN'] + post_data_str)
        return sign == data_hash

    def verify_post_bytes_sign(self, sign, data_bytes):
        byts_hash = self.get_str_hash(data_bytes)
        data_hash = self.get_str_hash(self.ZK_CONF['TOKEN'] + byts_hash)
        return sign == data_hash

    def check_is_json(self, json_str):
        try:
            json_object = json.loads(json_str)
        except ValueError as e:
            return False

        return json_object

    def get_image_ext(self, img_bytes):
        w = imghdr.what("", img_bytes)
        if w is None:
            w = "jpeg"
        return w

    def locate_imgfile_tmp(self, img_data, img_type='file'):    
        ##try get img by url
        if self.is_url(img_data):
            response = requests.get(img_data, verify=False, timeout=10)
            img = Image.open(io.BytesIO(response.content))
        elif img_type == 'bytes':
            img = Image.open(io.BytesIO(img_data))
        elif img_type == 'base64':
            img = Image.open(io.BytesIO(base64.b64decode(img_data)))
        elif img_type == 'svg':
            img_data = svg2png(bytestring=img_data)
            img = Image.open(io.BytesIO(img_data))
        else:
            img = Image.open(img_data) 

        return img

    def save_tmp_imagebytes(self, img_data, img_type='file'):
        img = self.locate_imgfile_tmp(img_data, img_type)
        img = img.convert("RGBA")

        width, height = img.size
        data = np.asarray(img).ravel()

        tmp = '/tmp/inpaint/'
        Path(tmp).mkdir(parents=True, exist_ok=True)
        path_name = self.uuid()
        path = str(tmp + path_name)

        with open(path, 'w') as f:
            for x in data:
                f.write(str(x) + " ")
                
        return {'path_name': path_name, 'path':path, 'width': width, 'height': height}

    def save_tmp_maskbytes(self, img_data, img_type='svg'):
        img = self.locate_imgfile_tmp(img_data, img_type)
        width, height = img.size

        tmp = '/tmp/inpaint/'
        Path(tmp).mkdir(parents=True, exist_ok=True)
        path_name = self.uuid()
        path = str(tmp + path_name)

        mask = []
        data = np.array(img)
        for x in range(height):
            for y in range(width):
                if data[x][y][0] > 0 or data[x][y][1] > 0 or data[x][y][2] > 0 or data[x][y][3] > 0:
                    mask.append(1)
                else:
                    mask.append(0)

        # Write 1/0 array of mask to file
        with open(path, 'w') as f:
            for x in mask:
                f.write(str(x) + " ")
                
        return {'path_name': path_name, 'path':path, 'width': width, 'height': height}    

    def uuid(self):
        return str(uuid.uuid4().hex)  

    def download_url_to_file(self, url, dst, hash_prefix=None, progress=True):
        r"""Download object at the given URL to a local path.

        Args:
            url (string): URL of the object to download
            dst (string): Full path where object will be saved, e.g. ``/tmp/temporary_file``
            hash_prefix (string, optional): If not None, the SHA256 downloaded file should start with ``hash_prefix``.
                Default: None
            progress (bool, optional): whether or not to display a progress bar to stderr
                Default: True

        Example:
            >>> torch.hub.download_url_to_file('https://s3.amazonaws.com/pytorch/models/resnet18-5c106cde.pth', '/tmp/temporary_file')

        """
        file_size = None
        req = Request(url, headers={"User-Agent": "torch.hub"})
        u = urlopen(req)
        meta = u.info()
        if hasattr(meta, 'getheaders'):
            content_length = meta.getheaders("Content-Length")
        else:
            content_length = meta.get_all("Content-Length")
        if content_length is not None and len(content_length) > 0:
            file_size = int(content_length[0])

        # We deliberately save it in a temp file and move it after
        # download is complete. This prevents a local working checkpoint
        # being overridden by a broken download.
        dst = os.path.expanduser(dst)
        dst_dir = os.path.dirname(dst)
        f = tempfile.NamedTemporaryFile(delete=False, dir=dst_dir)

        try:
            if hash_prefix is not None:
                sha256 = hashlib.sha256()
            with tqdm(total=file_size, disable=not progress,
                    unit='B', unit_scale=True, unit_divisor=1024) as pbar:
                while True:
                    buffer = u.read(8192)
                    if len(buffer) == 0:
                        break
                    f.write(buffer)
                    if hash_prefix is not None:
                        sha256.update(buffer)
                    pbar.update(len(buffer))

            f.close()
            if hash_prefix is not None:
                digest = sha256.hexdigest()
                if digest[:len(hash_prefix)] != hash_prefix:
                    raise RuntimeError('invalid hash value (expected "{}", got "{}")'
                                    .format(hash_prefix, digest))
            shutil.move(f.name, dst)
        finally:
            f.close()
            if os.path.exists(f.name):
                os.remove(f.name)

    @staticmethod
    def fail_result(fail_reason = None):
        resp = {
            "result": "fail"
        }
        if isinstance(fail_reason, str):
            resp['reason'] = fail_reason
        elif isinstance(fail_reason, dict):
            resp = {**resp, **fail_reason}

        return resp

    @staticmethod
    def success_result(data = None):
        resp = {
            "result": "success"
        }

        if isinstance(data, dict):
            resp = {**resp, **data}

        return resp

    @staticmethod
    def is_fail_result(json_str):
        return True if 'result' in json_str and json_str['result'] == 'fail' else False
    
    @staticmethod
    def is_self_ip(request):
        try:
            ip = request.headers.get('X-FORWARDED-FOR',None)
            if ip is False or ip is None:
                ip = request.remote

            return ip_address(ip).is_private or ip in ['112.5.138.146']

        except:
            return False

if __name__ == '__main__':
    #oss = AppTool()
    pass
