from flask import Flask, request, jsonify
from werkzeug.utils import secure_filename
import base64
import io
import os
from PIL import Image, ImageOps
import numpy as np
import cv2
from scipy.ndimage import interpolation as inter
import pytesseract
from pdf2image import convert_from_path, convert_from_bytes



def correct_skew(image, delta=1, limit=5):
    def determine_score(arr, angle):
        data = inter.rotate(arr, angle, reshape=False, order=0)
        histogram = np.sum(data, axis=1, dtype=float)
        score = np.sum((histogram[1:] - histogram[:-1]) ** 2, dtype=float)
        return histogram, score

    if len(image.shape)==3:
        gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
    else:
        gray = image
    thresh = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)[1] 

    scores = []
    angles = np.arange(-limit, limit + delta, delta)
    for angle in angles:
        histogram, score = determine_score(thresh, angle)
        scores.append(score)

    best_angle = angles[scores.index(max(scores))]

    (h, w) = image.shape[:2]
    center = (w // 2, h // 2)
    M = cv2.getRotationMatrix2D(center, best_angle, 1.0)
    corrected = cv2.warpAffine(image, M, (w, h), flags=cv2.INTER_CUBIC, \
            borderMode=cv2.BORDER_REPLICATE)

    return best_angle, corrected

def getInputData(request_data):
    if not request.json or 'fileType' not in request.json:
         return 'Error: No fileType provided', 400
    file_type = request_data.get('fileType')

    if not request.json or 'file' not in request.json:
         return 'Error: No file provided', 400
    file_data = request_data.get('file')

    if not request.json or 'ocrEngine' not in request.json:
         return 'Error: No ocrEngine provided', 400
    ocr_engine = request_data.get('ocrEngine')

    if not request.json or 'language' not in request.json:
         return 'Error: No language provided', 400
    language = request_data.get('language')

    if not request.json or 'scale' not in request.json:
         return 'Error: No scale provided', 400
    scale = request_data.get('scale')

    if not request.json or 'detectOrientation' not in request.json:
         return 'Error: No detectOrientation provided', 400
    detect_orientation = request_data.get('detectOrientation')

    if not request.json or 'region' not in request.json:
         return 'Error: No region provided', 400
    region = request_data.get('region')

    extracted_data = {
        'fileType': file_type,
        'file_data': file_data,
        'ocrEngine': ocr_engine,
        'language': language,
        'scale': scale,
        'detectOrientation': detect_orientation,
        'region': region
    }

    return extracted_data


app = Flask(__name__)


UPLOAD_FOLDER = 'files/upload/'
ALLOWED_EXTENSIONS = {'png', 'jpg', 'jpeg', 'bmp', 'webp', 'pdf', 'gif', 'tif'}

app.config['UPLOAD_FOLDER'] = UPLOAD_FOLDER


def allowed_file(ext):
    if ext in ALLOWED_EXTENSIONS:
        return True
    else:
        return False
#################################################################################################################################################################
##################################### Adjust Gamma ##############################################################################################################
#################################################################################################################################################################

def adjust_gamma(image, gamma=1.2):
    # build a lookup table mapping the pixel values [0, 255] to
    # their adjusted gamma values
    invGamma = 1.0 / gamma
    table = np.array([((i / 255.0) ** invGamma) * 255
        for i in np.arange(0, 256)]).astype("uint8")

    # apply gamma correction using the lookup table
    return cv2.LUT(image, table)

##################################################################################################################################################################
#################################### Adaptive binarization #######################################################################################################
##################################################################################################################################################################
BLOCK_SIZE = 40
DELTA = 25

def preprocessAD(image):
    image = cv2.medianBlur(image, 3)
    return 255 - image


def postprocessAD(image):
    kernel = np.ones((3,3), np.uint8)
    image = cv2.morphologyEx(image, cv2.MORPH_OPEN, kernel)
    return image

def get_block_index(image_shape, yx, block_size): 
    y = np.arange(max(0, yx[0]-block_size), min(image_shape[0], yx[0]+block_size))
    x = np.arange(max(0, yx[1]-block_size), min(image_shape[1], yx[1]+block_size))
    ymin = max(0, yx[0]-block_size)
    ymax = min(image_shape[0], yx[0]+block_size)
    xmin = max(0, yx[1]-block_size)
    xmax = min(image_shape[1], yx[1]+block_size)
    return ymin, ymax, xmin, xmax

def adaptive_median_threshold(img_in):
    med = np.median(img_in)
    img_out = np.zeros_like(img_in)
    img_out[img_in - med < DELTA] = 255
    kernel = np.ones((5,5),np.uint8)
    img_out = 255 - cv2.dilate(255 - img_out,kernel,iterations = 2)
    #img_out = 255 - cv2.dilate(img_out,kernel,iterations = 2)
    return img_out


def block_image_process(image, block_size):
    out_image = np.zeros_like(image)
    for row in range(0, image.shape[0], block_size):
        for col in range(0, image.shape[1], block_size):
            idx = (row, col)
            ymin, ymax, xmin, xmax = get_block_index(image.shape, idx, block_size)
            out_image[ymin:ymax, xmin:xmax] = adaptive_median_threshold(image[ymin:ymax, xmin:xmax])
    return out_image

def process_image(img):
    image_in = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
    image_in = preprocessAD(image_in)
    image_out = block_image_process(image_in, BLOCK_SIZE)
    image_out = postprocessAD(image_out)
    return image_out   

#############################################################################################################################################################
#################################### Soft Part Binarization #################################################################################################
#############################################################################################################################################################

def sigmoid(x, orig, rad):
    k = np.exp((x - orig) * 5 / rad)
    return k / (k + 1.)

def combine_block(img_in, mask):
    img_out = np.zeros_like(img_in)
    img_out[mask == 255] = 255
    fimg_in = img_in.astype(np.float32)

    idx = np.where(mask == 0)
    if idx[0].shape[0] == 0:
        img_out[idx] = img_in[idx]
        return img_out

    lo = fimg_in[idx].min()
    hi = fimg_in[idx].max()
    v = fimg_in[idx] - lo
    r = hi - lo

    img_in_idx = img_in[idx]
    ret3,th3 = cv2.threshold(img_in[idx],0,255,cv2.THRESH_BINARY+cv2.THRESH_OTSU)

    bound_value = np.min(img_in_idx[th3[:, 0] == 255])
    bound_value = (bound_value - lo) / (r + 1e-5)
    f = (v / (r + 1e-5))
    f = sigmoid(f, bound_value + 0.05, 0.2)

    img_out[idx] = (255. * f).astype(np.uint8)
    return img_out

def combine_block_image_process(image, mask, block_size):
    out_image = np.zeros_like(image)
    for row in range(0, image.shape[0], block_size):
        for col in range(0, image.shape[1], block_size):
            idx = (row, col)
            ymin,ymax, xmin,xmax = get_block_index(image.shape, idx, block_size)
            out_image[ymin:ymax, xmin:xmax] = combine_block(
                image[ymin:ymax, xmin:xmax], mask[ymin:ymax, xmin:xmax])
    return out_image

def combine_process(img, mask):
    image_in = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
    image_out = combine_block_image_process(image_in, mask, 20)
   # image_out = combine_postprocess(image_out)
    return image_out


#############################################################################################################################################################
#############################################################################################################################################################
#############################################################################################################################################################

def preprocess(pdfFile, image, extracted_data):
    #preprocess data to improve OCR quality
    image_np = np.array(image)
    #improve quality of photo
    if(pdfFile==False):
        #1. Adjust Gamma
        image_np = adjust_gamma(image_np)

        #2. Adaptive Binarization 
        image2 = process_image(image_np)

        #3. "Soft" Part of Binarization
        image_np = combine_process(image_np, image2)

   # cv2.imwrite('test/pokus.jpg',image_np)

    #add border to recognize text near boundaries
    color_name = 0 #White color
    value = [color_name for i in range(3)]
    image_np = cv2.copyMakeBorder(image_np, 50, 50, 50, 50, cv2.BORDER_CONSTANT, value=value)

    #cv2.imwrite('test/pokus.jpg',image_np)
    
    #img_with_border = ImageOps.expand(image_np, border=50, fill='red')
    #image_np = np.array(img_with_border)

    clahe = cv2.createCLAHE(clipLimit=1)
    if(extracted_data['detectOrientation']):
        angle, image_np = correct_skew(image_np)

    #Scanned documents are ok, CLAHE is not needed
    if(pdfFile==False):
        if len(image_np.shape)==3:
            image_np = cv2.cvtColor(image_np, cv2.COLOR_BGR2GRAY)
            image_np = clahe.apply(image_np)
        else:
            image_np = clahe.apply(image_np)
    else:
        image_np = cv2.cvtColor(image_np, cv2.COLOR_BGR2GRAY)

    #cv2.imwrite('test/pokus2.jpg',image_np)
    return image_np


def makeTransform(pdfFile, image, extracted_data):
    #preprocess to improve quality of results
    
    image_np = preprocess(pdfFile, image, extracted_data)

    custom_config = r'-l ces --oem 1 --psm 6' 
    
    text_opraveno = pytesseract.image_to_string(image_np,config=custom_config)

    return text_opraveno




# Route to receive data and return text
@app.route('/process_data', methods=['POST'])
def process_data():
    # Check if request contains JSON data
    request_data = request.json

    extracted_data = getInputData(request_data)

    encoded_image = extracted_data['file_data']

    if allowed_file(extracted_data['fileType']):
        filename = 'img.'+extracted_data['fileType']
        decoded_image = base64.b64decode(encoded_image)

        pdfFile = (extracted_data['fileType']=='pdf')
        print(pdfFile)
        if pdfFile:
            img = convert_from_bytes(decoded_image)[0]
        else:
            img = Image.open(io.BytesIO(decoded_image))


        

        image = img    

        
        if(extracted_data['region']['makeCut']):
            size = extracted_data['region']['coordinates'][0]
            image = image.crop((size['x'], size['y'], size['x'] + size['width'], size['y'] + size['height']))


        text = makeTransform(pdfFile, image, extracted_data)        
    else:
        return 'Error: No supported file format', 400





    parsed_text = {
        "html": "<p>This is converted text from the input image.</p>",
        "plainText": text
    }

    # Return the parsed text as JSON
    return jsonify({"parsedText": parsed_text})


if __name__ == '__main__':
    app.run(debug=True)