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, time=[]):
    diff_time = { }
    for i in range(len(time) - 1):
        if math.fabs(time[i + 1] - time[i] - interval) > interval * 1E-3:
            diff_time[time[i]] = time[i + 1]

    # diff_time.
    return diff_time


def acquireTime(topicName, bag):
    # init timestamp list
    t_cam00, t_cam01, t_cam02, t_cam03, t_imu = [], [], [], [], []
    isImuFirst = True
    startTimeImu = 0

    for topic, msg, t in bag.read_messages():
        if topic == topicName[0]:
            t_cam00.append(msg.header.stamp.to_sec())
        if topic == topicName[1]:
            t_cam01.append(msg.header.stamp.to_sec())
        if topic == topicName[2]:
            t_cam02.append(msg.header.stamp.to_sec())
        if topic == topicName[3]:
            t_cam03.append(msg.header.stamp.to_sec())
        if topic == topicName[4]:
            t_imu.append(msg.header.stamp.to_sec())
            if isImuFirst:
                isImuFirst = False
                startTimeImu = msg.header.stamp.to_sec()

    # print bag.get_start_time(), bag.get_end_time()
    return t_cam00, t_cam01, t_cam02, t_cam03, t_imu, t_cam00[0], t_cam00[-1], startTimeImu


def extractFeature(thres, bagFiles, result):
    for i in range(len(bagFiles)):
        nowFile = bagFiles[i]
        try:
            bag = rosbag.bag.Bag(nowFile, "r")
        except:
            print("open rosbag file error, check you path")
            exit(-1)
        # nowFile = os.path.splitext(bagFiles[i])[0]
        str_path = os.path.dirname(bagFiles[i])
        nowFile = bagFiles[i][len(str_path) + 1: ]
        result[nowFile] = []
        prefix = W + nowFile + ": " + W
        bridge = CvBridge()
        for topic, msg, t in bag.read_messages():
            if "cam" in topic:
                cv_image = bridge.imgmsg_to_cv2(msg, "bgr8")
                orb = cv2.ORB_create()
                kp1, des1 = orb.detectAndCompute(cv_image, None)
                # keyp_without_size = copy.copy(training_image)
                # cv2.drawKeypoints(cv_image, kp1, cv_image, color = (0, 255, 0))
                # cv2.imshow("img", cv_image)
                k = cv2.waitKey(1)
                if len(kp1) <= thres:
                    result[nowFile].append(prefix + R + "ERROR: feature num is less than threshold. feature num: " + str(len(kp1)) + " time: " + str(msg.header.stamp.to_sec()) + R)
        if len(result[nowFile]) == 0:
            result[nowFile].append(prefix + G + "NORMAL: feature detection normal" + G)

        for i in range(len(result[nowFile])):
            print (result[nowFile][i])
        # break


def checkTime(topicName, bagFiles, result = { }):
    """
    time check:
        - four cameras 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]
        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_cam00, t_cam01, t_cam02, t_cam03, t_imu, startTime, endTime, imuStart = acquireTime(topicName, bag)

        # init timestamp list
        # endTime = t_cam00[-1]
        duration = endTime - startTime
        # print duration
        if len(t_cam00) != 0:
            cam_fps = int((t_cam00[-1] - t_cam00[0]) / 0.1)
        else:
            cam_fps = 0
        # if (len(t_imu) != 0):
        #     imu_fps = int((t_imu[-1] - imuStart) / 0.005)
        # else:
        #     imu_fps = 0

        prefix = W + nowFile + ": " + W
        cam00, cam01, cam02, cam03 = True, True, True, True

        # fps check
        if math.fabs(len(t_cam00) - cam_fps) >= 3:
            result[nowFile].append(prefix + R + "camera 00 records: " + str(len(t_cam00)) + ", ERROR: frame loss occur, loss: " + str(math.fabs(len(t_cam00) - cam_fps)) + R)

        if math.fabs(len(t_cam01) - cam_fps) >= 3:
            result[nowFile].append(prefix + R + "camera 01 records: " + str(len(t_cam01)) + ", ERROR: frame loss occur, loss: " + str(math.fabs(len(t_cam00) - cam_fps)) + R)

        if math.fabs(len(t_cam02) - cam_fps) >= 3:
            result[nowFile].append(prefix + R + "camera 02 records: " + str(len(t_cam02)) + ", ERROR: frame loss occur, loss: " + str(math.fabs(len(t_cam00) - cam_fps)) + R)

        if math.fabs(len(t_cam03) - cam_fps) >= 3:
            result[nowFile].append(prefix + R + "camera 03 records: " + str(len(t_cam03)) + ", ERROR: frame loss occur, loss: " + str(math.fabs(len(t_cam00) - cam_fps)) + R)

        # if math.fabs(len(t_imu) - imu_fps) > 10:
        #     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) <= 3 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 camera records data are normal" + G)

        # interval check
        # diff_dict = intervalCheck(0.005, 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)

        diff_dict = intervalCheck(0.1, t_cam00)
        if diff_dict:
            for key, value in diff_dict.items():
                cam00 = False
                result[nowFile].append(prefix + O + "CAMERA 00 INTERVAL WARNING: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + O)

        diff_dict = intervalCheck(0.1, t_cam01)
        if diff_dict:
            for key, value in diff_dict.items():
                cam01 = False
                result[nowFile].append(prefix + O + "CAMERA 01 INTERVAL WARNING: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + O)

        diff_dict = intervalCheck(0.1, t_cam02)
        if diff_dict:
            for key, value in diff_dict.items():
                cam02 = False
                result[nowFile].append(prefix + O + "CAMERA 02 INTERVAL WARNING: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + O)

        diff_dict = intervalCheck(0.1, t_cam03)
        if diff_dict:
            for key, value in diff_dict.items():
                cam03 = False
                result[nowFile].append(prefix + O + "CAMERA 03 INTERVAL WARNING: t = " + str(key) + ", t+1 = " + str(value) + ", INTERVAL: " + str(value - key) + O)

        if cam00 and cam01 and cam02 and cam03:
            result[nowFile].append(prefix + G + "NORMAL: interval of CAMERAS are normal" + G)

        if i != 0 and math.fabs(lastEnd - startTime) > 1.1E-1:
            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, featureNum, reportPath):
    """
    check imu and camera frame in a bag file, including:
        - timestamp
        - sensor output frequency
    :param bagFile:
    :param topic:
    :return:
    """
    topicName = ["/cam00/image_raw", "/cam01/image_raw", "/cam02/image_raw", "/cam03/image_raw", "/imu_raw"]
    result_time = { }
    bagFiles = glob.glob(bagFile + "*.bag")
    bagFiles.sort()
    print("execute time check procedure. It will take few minutes")
    checkTime(topicName, bagFiles, result_time)
    # print(W + "extracting features. It will take few minutes" + W)
    # result_feature = { }
    # extractFeature(featureNum, bagFiles, result_feature)

    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
            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")

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

            # f.write(result_feature[key][i] + "\n")

    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')
        for key, value in result_time.items():
            for i in range(len(result_time[key])):
                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 error_list
    # # first check timestamp and fps


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

    featureNum = opt["feature_threshold"]
    bagPath = opt["bag_path"]
    reportPath = opt["report_path"] + "cam_imu_report.txt"
    print (featureNum)
    print (bagPath)
    print (reportPath)
    error_list = checkRosbag(bagPath, featureNum, reportPath)
    if len(error_list) == 0:
        print ("SUMMARY: " + G + "Data records are normal!")
    else:
        print ("SUMMARY: " + R + str(len(error_list)) + " errors!" + " please check\n")
        for i in range(len(error_list)):
            print (error_list[i])
