"""
calibration.py - tools for calibration camera array from target images
"""

import numpy as np
import cv2
import glob
import sys
import os
import json

# numpy to json helper
class NumpyEncoder(json.JSONEncoder):
    def default(self, obj):
        if isinstance(obj, np.ndarray):
            return obj.tolist()
        return json.JSONEncoder.default(self, obj)

# draws a grid on an image
def draw_grid(img, line_color=(0, 255, 0), thickness=1, type_=cv2.LINE_AA, pxstep=50):
    '''(ndarray, 3-tuple, int, int) -> void
    draw gridlines on img
    line_color:
        BGR representation of colour
    thickness:
        line thickness
    type:
        8, 4 or cv2.LINE_AA
    pxstep:
        grid line frequency in pixels
    '''
    x = pxstep
    y = pxstep
    while x < img.shape[1]:
        cv2.line(img, (x, 0), (x, img.shape[0]), color=line_color, lineType=type_, thickness=thickness)
        x += pxstep

    while y < img.shape[0]:
        cv2.line(img, (0, y), (img.shape[1], y), color=line_color, lineType=type_, thickness=thickness)
        y += pxstep

# runs the calibration on a specific camera
def calibrate_camera(image_path, file_pattern, board_shape=(10,10), save_grids=False):

    # termination criteria
    criteria = (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 30, 0.001)

    # prepare object points, like (0,0,0), (1,0,0), (2,0,0) ....,(6,5,0)
    objp = np.zeros((board_shape[0] * board_shape[1], 3), np.float32)
    objp[:, :2] = np.mgrid[0:board_shape[0], 0:board_shape[1]].T.reshape(-1, 2)

    # Arrays to store object points and image points from all the images.
    objpoints = []  # 3d point in real world space
    imgpoints = []  # 2d points in image plane.

    # Get the images matching the path and pattern for the camera
    images = glob.glob(os.path.join(image_path, '*' + file_pattern + '.png'))

    # Detect the corners
    for fname in images:
        print(fname)
        img = cv2.imread(fname)
        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)

        # Find the chess board corners
        ret, corners = cv2.findChessboardCorners(gray, board_shape, None)

        # If found, add object points, image points (after refining them)
        if ret == True:
            objpoints.append(objp)

            corners2 = cv2.cornerSubPix(gray, corners, (11, 11), (-1, -1), criteria)
            imgpoints.append(corners2)

            # Draw and display the corners
            img = cv2.drawChessboardCorners(img, board_shape, corners2, ret)
            scale = 1
            if img.shape[0] > 2048:
                scale = 4
            elif img.shape[0] > 1024:
                scale = 2
            cv2.imshow('img', cv2.resize(img, (img.shape[1] / scale, img.shape[0] / scale)))
            cv2.imwrite(fname[:-4] + '_corners.jpg', img)
            cv2.waitKey(100)

    # Do the calibration
    ret, mtx, dist, rvecs, tvecs = cv2.calibrateCamera(objpoints, imgpoints, gray.shape[::-1], None, None)

    # show the RMS re-projection error
    print("RMS error: ", ret)

    # Save out the results
    if ret < 0.5:
        print('Calibration okay.')
        output = {}
        output['mtx'] = mtx
        output['dist'] = dist
        output['rvecs'] = rvecs
        output['tvecs'] = tvecs
        output['image_width'] = img.shape[1]
        output['image_height'] = img.shape[0]
        output['rms_error'] = ret
        with open(os.path.join(image_path, 'calibration_' + file_pattern + '.json'), 'w') as outfile:
            json.dump(output, outfile, sort_keys=True, indent=4, cls=NumpyEncoder)

    if save_grids:

        # revise the calibration matrix
        h, w = img.shape[:2]
        newcameramtx, roi = cv2.getOptimalNewCameraMatrix(mtx, dist, (w, h), 1, (w, h))

        # build the distortion correction mapping
        mapx, mapy = cv2.initUndistortRectifyMap(mtx, dist, None, newcameramtx, (w, h), 5)

        # undistort the calibration images
        for fname in images:
            img = cv2.imread(fname)
            dst = cv2.remap(img, mapx, mapy, cv2.INTER_LINEAR)
            cv2.imwrite(fname[:-4] + '_remap.jpg', dst)

        # invert the undistortion brute-force like (need some extra pixel room in the arrays)
        mapx_inv = np.zeros((h + 100, w + 100), np.float32)
        mapy_inv = np.zeros((h + 100, w + 100), np.float32)
        for i in range(0, mapx.shape[0]):
            for j in range(0, mapx.shape[1]):
                mapx_inv[int(mapy[i, j]), int(mapx[i, j])] = j
                mapy_inv[int(mapy[i, j]), int(mapx[i, j])] = i

        # create a qualitative example of the distortion of a grid
        grid_img = 0 * (img.copy())
        draw_grid(grid_img)
        dst = cv2.remap(grid_img, mapx_inv, mapy_inv, cv2.INTER_LINEAR)
        cv2.imwrite(os.path.join(image_path, 'distortion_grid.jpg'), dst[0:1080 - 1, 0:1920 - 1])

# main function
if __name__=="__main__":

    if len(sys.argv) < 2:
        print('Please input path to images as first argument')
        exit()

    cameras = ['LT', 'LB', 'MT', 'MB', 'RT', 'RB']

    for cam in cameras:

        calibrate_camera(sys.argv[1], cam)