#coding=utf-8 
import rosbag
import cv2
import  math
import glob
from cv_bridge import CvBridge
import sys
import os
import yaml


R  = '\033[31m' # red
G  = '\033[32m' # green
O  = '\033[33m' # orange
W  = '\033[0m'  # white (normal)


def intervalCheck(interval, precision,time=[]):
    diff_time = { }
    for i in range(len(time) - 1):
        if math.fabs(time[i + 1] - time[i] - interval) > interval * precision:
            diff_time[time[i]] = time[i + 1]

    # diff_time.
    return diff_time


def acquireTime(topicName, bag):
    # init timestamp list
    t_lidar, t_imu = [], []

    for topic, msg, t in bag.read_messages():
        if topic == topicName[0]:
            t_lidar.append(msg.header.stamp.to_sec())
        if topic == topicName[1]:
            t_imu.append(msg.header.stamp.to_sec())


    return t_lidar, t_imu, bag.get_start_time(), bag.get_end_time()



def checkTime(topicName, bagFiles, threshLidarMsgsLossNum,threshImuMsgsLossNum,precision,result):
    """
    time check:
        - lidar time check
        - imu time check
        - camera-imu time check
    :param topicName:
    :param bag:
    :return:
    """
    lastEnd = 0
    for i in range(len(bagFiles)):
        nowFile = bagFiles[i]
        print nowFile
        try:
            bag = rosbag.bag.Bag(nowFile, "r")
        except:
            print("open rosbag file error, check you path")
            exit(-1)
        str_path = os.path.dirname(bagFiles[i])
        nowFile = bagFiles[i][len(str_path) + 1: ]
        result[nowFile] = []

        t_lidar, t_imu, startTime, endTime = acquireTime(topicName, bag)
        

        # init timestamp list
        duration = endTime - startTime
        lidar_fps = 595
        imu_fps = int(duration / 0.005)
      
        prefix = W + nowFile + ": " + W
        print prefix
        lidarflag = True

        # fps check
        if math.fabs(len(t_lidar) - lidar_fps) >= threshLidarMsgsLossNum:
            result[nowFile].append(prefix + R + "lidar records: " + str(len(t_lidar)) + ", ERROR: frame loss occur, loss: " + str(math.fabs(len(t_lidar) - lidar_fps)) + R)

        # if math.fabs(len(t_imu) - imu_fps) > threshImuMsgsLossNum:
        #     result[nowFile].append(prefix + R + "IMU records: " + str(len(t_imu)) + ", ERROR: frame loss occur, loss: " + str(math.fabs(len(t_imu) - imu_fps)) + R)
        # elif math.fabs(len(t_imu) - imu_fps) <= threshImuMsgsLossNum and math.fabs(len(t_imu) - imu_fps) >= 1:
        #     result[nowFile].append(prefix + O + "IMU records: " + str(len(t_imu)) + ", WARNING: frame loss occur, loss: " + str(math.fabs(len(t_imu) - imu_fps)) + O)

        if len(result[nowFile]) <= 1:
            result[nowFile].append(prefix + G + "NORMAL: number of lidar data is normal" + G)

        # interval check
        # diff_dict = intervalCheck(0.005,precision, t_imu)
        # if diff_dict:
        #     for key, value in diff_dict.items():
        #         if math.fabs(value - key) <= 0.05:
        #             result[nowFile].append(prefix + O + "IMU INTERVAL WARNING: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + O)
        #         else:
        #             result[nowFile].append(prefix + R + "IMU INTERVAL ERROR: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + R)
        # else:
        #     result[nowFile].append(prefix + G + "NORMAL: interval of IMU is normal" + G)
        # result[nowFile].append(prefix + G + "NORMAL: interval of IMU is normal" + G)
        diff_dict = intervalCheck(0.1,precision, t_lidar)
        if diff_dict:
            for key, value in diff_dict.items():
                lidarflag = False
                result[nowFile].append(prefix + R + "LIDAR INTERVAL ERROR: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + R)

 
        if lidarflag:
            result[nowFile].append(prefix + G + "NORMAL: interval of LIDAR is normal" + G)

        if i != 0 and math.fabs(lastEnd - startTime) > 1E-2:
            result[nowFile].append(prefix + R + "ERROR: bag time discontinue. last bag end time: " + str(lastEnd) + ", now bag end time: " + str(startTime) + " diff: " + str(math.fabs(lastEnd - startTime)))
        lastEnd = endTime

    #     for i in range(len(result[nowFile])):
    #         print result[nowFile][i]
    #     print " "
    #     # break
    # for key, value in result.items():
    #     for i in range(len(result[key])):
    #         print(result[key][i])


def checkRosbag(bagFile, reportPath,threshLidarMsgsLossNum,precision,threshImuMsgsLossNum):
    """
    check imu and camera frame in a bag file, including:
        - timestamp
        - sensor output frequency
    :param bagFile:
    :param topic:
    :return:
    """
    topicName = ["/velodyne_points", "/imu_raw"]
    result_time = { }
    bagFiles = glob.glob(bagFile + "*.bag")
    bagFiles.sort()
    print("execute time check procedure. It will take few seconds")
    checkTime(topicName, bagFiles, threshLidarMsgsLossNum,threshImuMsgsLossNum,precision,result_time)
    # print(result_time)
    # print(W + "extracting features. It will take few minutes" + W)
    # # result_feature = { }
    # # extractFeature(featureNum, bagFiles, result_feature)
    BagFlag=True

    error_list = []
    warning_count, error_count = 0, 0

    for key, value in result_time.items():
        for i in range(len(result_time[key])):
            if "ERROR" in result_time[key][i]:
                error_count += 1
                error_list.append(result_time[key][i])
            elif "WARN" in result_time[key][i]:
                warning_count += 1
            # print result_time[key][i]
            result_time[key][i] = result_time[key][i].replace(W, '')
            result_time[key][i] = result_time[key][i].replace(R, '')
            result_time[key][i] = result_time[key][i].replace(O, '')
            result_time[key][i] = result_time[key][i].replace(G, '')
            
            # f.write(result_time[key][i] + "\n")

    
    if((error_count>0)):
        BagFlag=False

    with open(reportPath, "w+") as f:
        f.write("SUMMARY: " + str(warning_count) + " warning(s) " + str(error_count) + " errors!" + " please check\n")
        f.write('--------------------------------------------------------------------------------------------------------\n\n')
        [(k,result_time[k]) for k in sorted(result_time.keys())]
        for key, value in result_time.items():
            # result_time[key].sort()
            for i in range(len(result_time[key])):
                # print result_time[key][i]
                f.write(result_time[key][i] + "\n")

        # for key, value in result_feature.items():
        #     for i in range(len(result_feature[key])):
        #         f.write(result_feature[key][i] + "\n")

        print W + "  "
    return BagFlag
    # # first check timestamp and fps


if __name__ == "__main__":
    optFile = os.getcwd() + "/lidar_imu_check.yaml"
    with open(optFile) as f:
        opt = yaml.safe_load(f)

    threshLidarMsgsLossNum = opt["threshLidarMsgsLossNum"]
    threshImuMsgsLossNum   =opt["threshImuMsgsLossNum"]
    precision=opt["precision"]
    bagPath = opt["bag_path"]
    reportPath = opt["report_path"] + "report.txt"
    print bagPath
    print reportPath
    
    BagFlag=checkRosbag(bagPath, reportPath,threshLidarMsgsLossNum,precision,threshImuMsgsLossNum)
    print(BagFlag)