import numpy as np
from PYME.IO.events import EVENTS_DTYPE
from PYME.IO.image import ImageStack
from PYME.Analysis import piecewiseMapping as piecewise_mapping
import logging

logger = logging.getLogger(__name__)

def convert_state_updates_to_laser_events(og_events):
    """
    Replaces state update events involving lasers with lasername On/Off events, e.g.
    'l635.On' or 'l635.Off'
    """
    from PYME.IO.events import EventLogger
    events = []
    replace_inds = []
    for ind in range(len(og_events)):
        if og_events[ind]['EventName'] == b'ProtocolTask':
            if b'update' in og_events[ind]['EventDescr']:
                if b"'Lasers." in og_events[ind]['EventDescr']:
                    replace_inds.append(ind)
                    laser_states = [s.strip(b' ').strip(b'{').strip(b'}') for s in og_events[ind]['EventDescr'].split(b',') if b'.On' in s]
                    for state in laser_states:
                        _, ev_name, io = state.split(b"'")
                        if b'False' in io:
                            ev_name = ev_name.replace(b'On', b'Off')
                        # 'EventName'][j], events_array['EventDescr'][j], events_array['Time'
                        events.append((ev_name, '', og_events[ind]['Time']))

    return np.concatenate([np.delete(og_events, replace_inds), 
                           EventLogger.list_to_array(events)])

def GenerateShiftedPMFromEventList(events, metadata, x0, y0, eventName=b'ProtocolFocus', dataPos=1, shift=0):
    """
    Parameters
    ----------
    events:
    metadata:
    x0: why?
    y0:
    eventName:
    dataPos: int
        position in comma-separated event['EventDesc'] str of the float which makes 'y' for this mapping
    shift: float
        time in seconds to shift all `eventName` events. Can be positive or negative

    Returns
    -------
    map: piecewiseMap
    """
    from PYME.Analysis.piecewiseMapping import piecewiseMap, times_to_frames
    x = []
    y = []

    secsPerFrame = metadata.getEntry('Camera.CycleTime')

    for e in events[events['EventName'] == eventName]:
        #if e['EventName'] == eventName:
        #print(e)
        x.append(e['Time'])
        y.append(float(e['EventDescr'].decode('ascii').split(', ')[dataPos]))
        
    x = np.array(x) + shift
    y = np.array(y)
        
    I = np.argsort(x)
    
    x = x[I]
    y = y[I]

    #print array(x) - metadata.getEntry('StartTime'), timeToFrames(array(x), events, metadata)

    return piecewiseMap(y0, times_to_frames(x, events, metadata), y, secsPerFrame, xIsSecs=False)

def median_delay_interval(events):
    delay_change_times = []
    for e in events:
        if e['EventName'] == b'pump-delay':
            delay_change_times.append(e['Time'])
    return np.median(np.diff(np.sort(delay_change_times)))

def map_lasers_and_delay(im, events, metadata, drop_uncertain):
    from scipy.ndimage.morphology import binary_erosion, binary_dilation
    from PYME.IO.MetaDataHandler import DictMDHandler
    from PYME.IO.tabular import RecArraySource

    pump_on = piecewise_mapping.bool_map_between_events(events, metadata, 
                                                        b'Lasers.l635.On', b'Lasers.l635.Off', 
                                                        default=False)
    probe_on = piecewise_mapping.bool_map_between_events(events, metadata, 
                                                            b'Lasers.l750.On', b'Lasers.l750.Off', 
                                                            default=False)
    delay = piecewise_mapping.GeneratePMFromEventList(events, metadata, 
                                                        metadata['StartTime'], metadata['PicosecondDelayer.Delay_ps'], 
                                                        eventName=b'pump-delay', dataPos=0)
    delay_time = median_delay_interval(events)
    # generate a 'delay' mapping where everything is shifted forwards 1/2 cycle
    # forward_negative_control_delay = GenerateShiftedPMFromEventList(events, metadata,
    #                                                         metadata['StartTime'], metadata['PicosecondDelayer.Delay_ps'], 
    #                                                         eventName=b'pump-delay', dataPos=0,
    #                                                         shift=0.5*delay_time)
    # generate a 'delay' mapping where everything is shifted backwards 1/2 cycle
    
    backward_negative_control_delay = GenerateShiftedPMFromEventList(events, metadata,
                                                            metadata['StartTime'], metadata['PicosecondDelayer.Delay_ps'], 
                                                            eventName=b'pump-delay', dataPos=0,
                                                            shift=-0.5*delay_time)

    # convert mapping functions into per-frame arrays
    frames = np.arange(im.data_xytc.shape[2])
    
    delay = delay(frames)
    cycle = np.cumsum(np.diff(delay, prepend=[0]) != 0)

    # forward_negative_control_delay = forward_negative_control_delay(frames)
    # forward_negative_control_cycle = np.cumsum(np.diff(forward_negative_control_delay, prepend=[0]) != 0)
    backward_negative_control_delay = backward_negative_control_delay(frames)
    backward_negative_control_cycle = np.cumsum(np.diff(backward_negative_control_delay, prepend=[0]) != 0)
    
    pump_on = pump_on(frames)
    # fig, ax = plt.subplots()
    # plt.plot(frames, cycle)
    # plt.ylabel('cycle', fontsize=20)
    # ax2 = ax.twinx()
    # ax2.plot(frames, delay, color='k')
    # ax2.set_ylabel('Delay [ps]', fontsize=20)
    # plt.tight_layout()
    # plt.show()
    probe_on = probe_on(frames)
    both_lasers = np.logical_and(pump_on, probe_on)
    pump_only = np.logical_and(pump_on, ~probe_on)
    probe_only = np.logical_and(~pump_on, probe_on)

    # create booleans from our delay arrays
    delays = np.unique(delay)
    assert(len(delays) == 2)
    phase0_raw = delay == delays[0]
    phase1_raw = delay == delays[1]
    # same for the negative control delay arrays
    # forward_nc_phase0_raw = forward_negative_control_delay == delays[0]
    # forward_nc_phase1_raw = forward_negative_control_delay == delays[1]
    backward_nc_phase0_raw = backward_negative_control_delay == delays[0]
    backward_nc_phase1_raw = backward_negative_control_delay == delays[1]
    
    # erode each bool array so we drop edge frames
    both_lasers = binary_erosion(both_lasers, iterations=drop_uncertain)
    pump_only = binary_erosion(pump_only, iterations=drop_uncertain)
    probe_only = binary_erosion(probe_only, iterations=drop_uncertain)

    no_lasers = ~np.logical_or(binary_dilation(pump_on, iterations=drop_uncertain),
                               binary_dilation(probe_on, iterations=drop_uncertain))
    
    phase0 = binary_erosion(phase0_raw, iterations=drop_uncertain)
    phase1 = binary_erosion(phase1_raw, iterations=drop_uncertain)
    # same for the negative control phase arrays
    # forward_nc_phase0 = binary_erosion(forward_nc_phase0_raw, iterations=drop_uncertain)
    # forward_nc_phase1 = binary_erosion(forward_nc_phase1_raw, iterations=drop_uncertain)
    backward_nc_phase0 = binary_erosion(backward_nc_phase0_raw, iterations=drop_uncertain)
    backward_nc_phase1 = binary_erosion(backward_nc_phase1_raw, iterations=drop_uncertain)
    

    frame_dt = [('filtered', [('both_lasers', bool), ('pump_only', bool), ('probe_only', bool), ('no_lasers', bool),
                             ('phase0', bool), ('phase1', bool),
                             ('backward_nc_phase0', bool), ('backward_nc_phase1', bool),
                             ('forward_nc_phase0', bool), ('forward_nc_phase1', bool)]),
                 ('raw', [('pump_on', bool), ('probe_on', bool),
                           ('phase0', bool), ('phase1', bool),
                           ('backward_nc_phase0', bool), ('backward_nc_phase1', bool),
                        #    ('forward_nc_phase0', bool), ('forward_nc_phase1', bool),
                           ('cycle', int),
                           ('forward_negative_control_cycle', int),
                           ('backward_negative_control_cycle', int)])]
    frames = np.zeros(len(both_lasers), dtype=frame_dt)
    frames['filtered']['both_lasers'] = both_lasers
    frames['filtered']['pump_only'] = pump_only
    frames['filtered']['probe_only'] = probe_only
    frames['filtered']['no_lasers'] = no_lasers
    frames['filtered']['phase0'] = phase0
    frames['filtered']['phase1'] = phase1
    # frames['filtered']['forward_nc_phase0'] = forward_nc_phase0
    # frames['filtered']['forward_nc_phase1'] = forward_nc_phase1
    frames['filtered']['backward_nc_phase0'] = backward_nc_phase0
    frames['filtered']['backward_nc_phase1'] = backward_nc_phase1
    frames['raw']['pump_on'] = pump_on
    frames['raw']['probe_on'] = probe_on
    frames['raw']['phase0'] = phase0_raw
    frames['raw']['phase1'] = phase1_raw
    frames['raw']['cycle'] = cycle
    # frames['raw']['forward_negative_control_cycle'] = forward_negative_control_cycle
    frames['raw']['backward_negative_control_cycle'] = backward_negative_control_cycle

    frames = RecArraySource(frames)
    frames.mdh = DictMDHandler(metadata)
    frames.mdh.setEntry('MapLasersAndDelay.NFramesDrop', drop_uncertain)
    frames.mdh.setEntry('MapLasersAndDelay.MedianDelayInterval', delay_time)
    return frames

def split_frames(im, phase0, phase1, cycles):
    p0, p1 = [], []
    
    for cycle in np.unique(cycles)[:-1]:  # Drop the cycle
        mask0 = cycle == phase0 * cycles
        mask1 = cycle == phase1 * cycles
        if cycle % 10 == 0:
            logger.info('cycle %d' % cycle)
        if np.any(mask0):
            p0.append(np.mean(np.squeeze(im.data_xytc[:, :, :, 0])[: , :, mask0], axis=2))
        else:
            assert np.any(mask1)
            # logger.warning('Cycle %d / %d has no filtered frames to average' % (cycle, np.unique(cycles)[-1]))
            # continue  # nothing to append, go on to the next cycle
            p1.append(np.mean(np.squeeze(im.data_xytc[:, :, :, 0])[: , :, mask1], axis=2))
    
    return p0, p1

def cycle_diffs(im, events, metadata, drop_uncertain):
    from PYME.IO.MetaDataHandler import DictMDHandler

    frames = map_lasers_and_delay(im, events, metadata, drop_uncertain)
    
    phase0 = np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['phase0'])
    phase1 = np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['phase1'])
    # pump_only = frames['filtered']['pump_only']

    # drop the last two cycles, so that our forward shift can drop 1, then first/last
    p0, p1 = split_frames(im, phase0, phase1, frames['raw']['cycle'])

    # debug plotting
    # import matplotlib.pyplot as plt
    # plt.plot(np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['phase0']), label='phase0')
    # plt.plot(np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['phase1']), label='phase1')
    # plt.plot(frames['filtered']['pump_only'], label='pump-only')
    # plt.legend()
    # plt.show()
    # subtract out adjacent cents
    paired_length = min(len(p0), len(p1))
    # note that we already dropped the last bin in `split_frames`, we don't
    # necessarily need to drop the last pair, but it will keep us same length as control
    if len(p0) == len(p1):
        paired_length -= 1
    sub = []
    for ind in range(paired_length):
        sub.append(p0[ind] - p1[ind])
    
    # forward_nc_phase0 = np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['forward_nc_phase0'])
    # forward_nc_phase1 = np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['forward_nc_phase1'])
    # f_nc_p0, f_nc_p1 = split_frames(im, forward_nc_phase0, forward_nc_phase1, 
    #                                 frames['raw']['forward_negative_control_cycle'])
    # # For the negative controls, the first bin width will be too short or too long, drop it
    # # first bin may not 
    # f_nc_p0.pop(0)  # first bin is always p0
    # paired_length = min(len(f_nc_p0), len(f_nc_p1))
    # # last bin will also be wrong length, if we aren't already dropping it, do so
    # # if len(f_nc_p0) == len(f_nc_p1):
    # #     paired_length -= 1
    # f_nc_sub = []
    # for ind in range(paired_length):
    #     f_nc_sub.append(f_nc_p0[ind] - f_nc_p1[ind])
    
    backward_nc_phase0 = np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['backward_nc_phase0'])
    backward_nc_phase1 = np.logical_and(frames['filtered']['both_lasers'], frames['filtered']['backward_nc_phase1'])
    b_nc_p0, b_nc_p1 = split_frames(im, backward_nc_phase0, backward_nc_phase1, 
                                    frames['raw']['backward_negative_control_cycle'])
    # For the negative controls, the first bin width will be too short or too long, drop it
    b_nc_p0.pop(0)  # first bin is always p0
    paired_length = min(len(b_nc_p0), len(b_nc_p1))
    # last bin will also be wrong length, if we aren't already dropping it, do so
    # if len(b_nc_p0) == len(b_nc_p1):
    #     paired_length -= 1
    b_nc_sub = []
    for ind in range(paired_length):
        b_nc_sub.append(b_nc_p0[ind] - b_nc_p1[ind])
    
    mdh = DictMDHandler(im.mdh)
    mdh.setEntry('AvgFramesBySync.NFramesDrop', drop_uncertain)
    
    # given that we dropped a pair, and no more than a pair, in each case,
    # should all be the same length, so we output as different color channels
    diff_stack = ImageStack(data=[np.stack(sub, axis=2), 
                                #   np.stack(f_nc_sub, axis=2), 
                                  np.stack(b_nc_sub, axis=2)], mdh=mdh)
    return diff_stack, frames