#import acapture
import scipy
import scipy.misc
import scipy.ndimage
import imageio
import time
import cv2
import numpy as np
import imageio
from matplotlib import cm

from multiprocessing import Pool, Process, Queue
import queue
from pyqtgraph.Qt import QtCore, QtGui

DEBUG = True



class VideoTransformer:

    def __init__(self):
        self.crop = False

class CannyTransformer(VideoTransformer):

    def __init__(self, id=1, thresholds=[[3, 20], [50, 125]], cmap_type='jet'):
        super().__init__()
        self.id = id
        self.channels = len(thresholds)
        self.thresholds = thresholds
        self.cmap = cm.get_cmap(cmap_type, self.channels)(self.channels)
        self.frames = []
        self.proc_frames = []
        self.busy = False


    def add_frames(self,frames):
        for f in frames:
            self.frames.append(f)
            self.proc_frames.append(f)

    def proc(self):
        self.busy = True
        print(str(self.id) + ' : starting...')
        for i, f in enumerate(self.frames):
            print(i)
            self.proc_frames[i] = cv2.Canny(f, 25, 150)
        self.busy = False
        print(str(self.id) + ' : finished.')




class VideoHandler:

    def __init__(self, video_path, transform=None):

        self.video_path = video_path
        self.is_loaded = False
        self.current_frame = 0
        self.frames_in_file = -1
        self.video = None
        self.frames = []
        self.video_load_time = 0.0
        self.transform = transform

    def read_frame(self):

        frame = None

        if self.is_loaded:
            ret, frame = self.video.read()

        return frame

    def chunk_vid(self):


    def print_load_stats(self):
        print('Loading Time: ' + str(self.video_load_time))
        print('Loading FPS: ' + str(self.frames_in_file / self.video_load_time))

    def threaded_processor(self,threads=16):

        # Chunk the loaded video in sublists and instantiate threads to process


    def load_FVS(self):

        if DEBUG:
            print('Loading video using FVS...')

        self.frames = []
        self.frames_in_file = 0
        self.video = FileVideoStream(self.video_path,
                    transform=self.transform, queue_size=4096).start()
        self.video_load_time = time.time()
        while self.video.more():
            f = self.video.read()
            if f is not None:
                #f = cv2.resize(f, (540, 960), interpolation=cv2.INTER_AREA)
                self.frames.append(f)
                self.frames_in_file += 1
        self.video.stop()
        self.video_load_time = time.time() - self.video_load_time
        self.print_load_stats()

    def load_CV2(self):

        if DEBUG:
            print('Loading video using CV2...')

        self.frames = []
        self.frames_in_file = 0
        self.video = cv2.VideoCapture(self.video_path)
        self.video_load_time = time.time()
        while True:
            ok, frame = self.video.read()
            if ok:
                if self.transform is not None:
                    frame = self.transform(frame)
                self.frames.append(frame)
                self.frames_in_file += 1
            else:
                break
        self.video.release()
        self.video_load_time = time.time() - self.video_load_time
        self.print_load_stats()


if __name__=="__main__":

    video_path = "C:\\Users\\proberts\\Videos\\vlc-record-2019-09-30-10h56m40s-DEEPPIV_S001_S102_T001.MOV-.avi"

    vfh = VideoHandler(video_path)
    vfh.load_FVS()
    #vfh.load_CV2()

    # split frames into processors
    ips = []
    idx = 0
    ncpus = 16
    total_frames = len(vfh.frames)
    chunk_size = int(total_frames/ncpus)
    for i in range(ncpus):
        tfx = CannyTransformer(id=i)
        if idx + chunk_size < total_frames:
            tfx.add_frames(vfh.frames[idx:idx+chunk_size])
        else:
            tfx.add_frames(vfh.frames[idx:])
        ips.append(tfx)
        idx += chunk_size

    procs = []
    for i in range(ncpus):
        p = Process(target=ips[i].proc)
        p.start()
        procs.append(p)

    busy = True
    while busy:
        busy = False
        for i in range(ncpus):
            if ips[i].busy:
                busy = True

    for p in procs:
        p.terminate()

    # preprocess using Pool
    #p = Pool(16)
    #tstart = time.time()
    #p.map(tfx.proc, vfh.frames)
    #print('Proc Time: ' + str(time.time()-tstart))