import os
import time
import numpy as np
from PIL import Image, ImageOps
from pycocotools.coco import COCO
import matplotlib.pyplot as plt

from coco2voc_aux import annotations_to_seg
import cv2




def resize_and_pad_mask(mask, output_size=(512, 512), fill=0):

    mask_pil = Image.fromarray(mask)

    original_width, original_height = mask_pil.size

    new_height = int(original_height * (output_size[0] / original_width))

    mask_resized = mask_pil.resize((output_size[0], new_height), Image.NEAREST)
    padding_top = (output_size[1] - new_height) // 2
    padding_bottom = output_size[1] - new_height - padding_top

    mask_padded = ImageOps.expand(mask_resized, (0, padding_top, 0, padding_bottom), fill=fill)

    return np.array(mask_padded)




def coco2png(annotations_file: str, target_folder: str, n: int = 1500, visualize = False):
    """
    Converts COCO style annotations to PNG masks and maps left and right anatomical structures to the same index.
    :param annotations_file: COCO annotations file
    :param target_folder: path to the folder where the results will be saved
    :param n: Number of image annotations to convert. Default is 1500.
    :return: PNG masks are saved to the target folder
    """

    coco_instance = COCO(annotations_file)
    coco_imgs = coco_instance.imgs

    n = min(n, len(coco_imgs))

    mask_target_path = os.path.join(target_folder)
    os.makedirs(mask_target_path, exist_ok=True)

    start = time.time()


    for i, img_id in enumerate(coco_imgs):
        if i >= n:
            break
        img_info = coco_imgs[img_id] 

        annotation_ids = coco_instance.getAnnIds(imgIds=img_id)
        annotations = coco_instance.loadAnns(annotation_ids)
        if not annotations:
            print(img_info['file_name'])
            continue

        class_seg, instance_seg, id_seg = annotations_to_seg(annotations, coco_instance)

        class_seg_resized = resize_and_pad_mask(class_seg)

        output_filename = os.path.splitext(img_info['file_name'])[0] #+ '.png'
        
        
        np.save(os.path.join(mask_target_path, output_filename), class_seg_resized)
        # output_path = os.path.join(mask_target_path, output_filename)
        # class_seg_uint8 = class_seg.astype(np.uint8)

        # # Create an image from the numpy array
        # img = Image.fromarray(class_seg_uint8)

        # # Save the image as PNG
        # img.save(output_path)



        if visualize:
            if i < 5:  
                plt.figure(figsize=(10, 5))
                plt.title(f'Mask for Image ID: {img_id}')
                plt.imshow(class_seg, cmap='jet')
                plt.axis('off')
                plt.show()

        if i % 100 == 0 and i > 0:
            print(f"{i} annotations processed in {int(time.time() - start)} seconds")

    return

json_path = 'inter_intra/alex/alex_intra_inter.json'


output_path = 'inter_intra/alex/inter_norm_size_masks_alex'
n = 100

coco2png(json_path, output_path, n=n,  visualize = True)
