from io import BytesIO
from sys import platform
from numpy import (asarray as np_asarray, uint8 as np_uint8)
from PIL import Image, ImageDraw
from threading import Thread
from var import bannedids, bot
import pytesseract
import cv2

if platform == 'win32':
    pytesseract.pytesseract.tesseract_cmd = r"C:\Program Files\Tesseract-OCR\tesseract.exe"

def imgtostring_mtd1(io: BytesIO):
    io.seek(0)
    img = cv2.imdecode(np_asarray(bytearray(io.read()), dtype=np_uint8), cv2.IMREAD_COLOR)
    img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
    img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)
    result = pytesseract.image_to_string(img, config='--psm 6 --oem 3')
    d = pytesseract.image_to_data(img, output_type=pytesseract.Output.DICT)
    n_boxes = len(d['text'])
    for i in range(n_boxes):
        (x, y, w, h) = (d['left'][i], d['top'][i], d['width'][i], d['height'][i])
        img = cv2.rectangle(img, (x, y), (x + w, y + h), (0, 255, 0), 1)
    newim = BytesIO()
    newim.write(cv2.imencode('.png', img)[1])
    newim.seek(0)
    return result, newim


def imgtostring_mtd2(io: BytesIO):
    io.seek(0)
    img = Image.open(io)
    pix = img.load()
    for y in range(img.size[1]):
        for x in range(img.size[0]):
            if pix[x, y][0] < 102 or pix[x, y][1] < 102 or pix[x, y][2] < 102:
                pix[x, y] = (0, 0, 0, 255)
            else:
                pix[x, y] = (255, 255, 255, 255)
    result = pytesseract.image_to_string(img, config='--psm 6 --oem 3')
    d = pytesseract.image_to_data(img, output_type=pytesseract.Output.DICT)
    n_boxes = len(d['text'])
    draw = ImageDraw.Draw(img)
    for i in range(n_boxes):
        (x, y, w, h) = (d['left'][i], d['top'][i], d['width'][i], d['height'][i])
        draw.rectangle(((x, y), (x + w, y + h)), outline=(0, 255, 0))
    newim = BytesIO()
    img.save(newim, 'png')
    newim.seek(0)
    return result, newim

def cmd(update, context):
    if update.message.from_user.id in bannedids:
        return
    method = update.message.text.split(' ')
    if len(method) == 1:
        method = 1
    else:
        method = 2
    if not update.message.reply_to_message:
        update.message.reply_text('Reply to a photo or a sticker to perform ocr operation.')
        return
    if update.message.reply_to_message is not None:
        if len(update.message.reply_to_message.photo) > 0:
            def do():
                fileid = update.message.reply_to_message.photo[
                    len(update.message.reply_to_message.photo) - 1].file_id
                m = update.message.reply_text('Processing image...')
                file = bot.getFile(fileid)
                img = BytesIO()
                img.write(file.download_as_bytearray())
                if method == 1:
                    result = imgtostring_mtd1(img)
                else:
                    result = imgtostring_mtd2(img)
                update.message.reply_photo(result[1])
                m.edit_text('Result:\n' + result[0])
                return

            Thread(target=do).start()
        if update.message.reply_to_message.sticker is not None:
            def do():
                fileid = update.message.reply_to_message.sticker.file_id
                m = update.message.reply_text('Processing image...')
                file = bot.getFile(fileid)
                img = BytesIO()
                img.write(file.download_as_bytearray())
                if method == 1:
                    result = imgtostring_mtd1(img)
                else:
                    result = imgtostring_mtd2(img)
                update.message.reply_photo(result[1])
                m.edit_text('Result:\n' + result[0])
                return

            Thread(target=do).start()
