"""Photo clean-up and Tesseract OCR for the glasses server.

Each photo is cleaned up two ways and both versions are read in parallel.
The one Tesseract is more confident about wins: flattening the lighting
suits flash photos with a bright spot, the black-and-white threshold suits
evenly lit pages.
"""
import io
import os
import re
import shutil
import time
from concurrent.futures import ThreadPoolExecutor

import cv2
import numpy as np
import pytesseract
from PIL import Image, ImageOps

# One core per Tesseract run, so the two versions run side by side on 2 cores.
os.environ.setdefault("OMP_THREAD_LIMIT", "1")

TARGET_LONG_SIDE = 2400   # px; ESP32-CAM UXGA (1600) is enlarged, phone photos shrunk
WINDOWS_TESSERACT = r"C:\Program Files\Tesseract-OCR\tesseract.exe"
HYPHENS = "-\u2010\u00ad"
CYRILLIC_LANGS = {"rus", "uzb_cyrl", "ukr", "bel", "bul", "kaz", "kir", "mkd", "srp", "tgk", "tat"}
SENTENCE_END = re.compile(r"[.!?…:;»\"”')\]]$")

_pool = ThreadPoolExecutor(max_workers=2)


def configure(tesseract_cmd="", tessdata_prefix=""):
    """Point pytesseract at the Tesseract binary and language data."""
    cmd = tesseract_cmd or shutil.which("tesseract") or WINDOWS_TESSERACT
    if os.path.exists(cmd) or shutil.which(cmd):
        pytesseract.pytesseract.tesseract_cmd = cmd
    if tessdata_prefix:
        os.environ["TESSDATA_PREFIX"] = tessdata_prefix


def available_languages():
    """Installed language packs, e.g. ['eng', 'rus', 'uzb']. Empty if Tesseract is missing."""
    try:
        return sorted(l for l in pytesseract.get_languages(config="") if l != "osd")
    except Exception:
        return []


def load_gray(data: bytes, rotate: int = 0) -> np.ndarray:
    """Decode a JPEG/PNG, honour the phone's EXIF orientation, rotate clockwise by `rotate` degrees."""
    img = ImageOps.exif_transpose(Image.open(io.BytesIO(data))).convert("L")
    if rotate:
        img = img.rotate(-rotate, expand=True)
    return np.asarray(img)


def resize(gray: np.ndarray) -> np.ndarray:
    h, w = gray.shape
    scale = min(2.0, max(0.25, TARGET_LONG_SIDE / max(h, w)))
    if abs(scale - 1) < 0.05:
        return gray
    interp = cv2.INTER_CUBIC if scale > 1 else cv2.INTER_AREA
    return cv2.resize(gray, None, fx=scale, fy=scale, interpolation=interp)


def variants(gray: np.ndarray) -> dict:
    h, w = gray.shape
    # Estimate the paper brightness on a small copy (dilating removes the dark text),
    # then divide it out so the flash hot-spot and dark corners even out.
    small = cv2.resize(gray, (max(1, w // 8), max(1, h // 8)), interpolation=cv2.INTER_AREA)
    paper = cv2.medianBlur(cv2.dilate(small, np.ones((7, 7), np.uint8)), 5)
    paper = cv2.resize(paper, (w, h), interpolation=cv2.INTER_LINEAR)
    flat = cv2.GaussianBlur(cv2.divide(gray, paper, scale=255), (3, 3), 0)

    binary = cv2.adaptiveThreshold(cv2.GaussianBlur(gray, (3, 3), 0), 255,
                                   cv2.ADAPTIVE_THRESH_GAUSSIAN_C, cv2.THRESH_BINARY, 31, 15)
    return {"flat": flat, "binary": binary}


def _read(img, langs, psm):
    return pytesseract.image_to_data(img, lang=langs, config=f"--psm {psm}",
                                     output_type=pytesseract.Output.DICT)


def _words(d):
    for i, text in enumerate(d["text"]):
        text = text.strip()
        conf = float(d["conf"][i])
        if text and conf >= 0:
            yield i, text, conf


def _score(d):
    """Sum of confidences of real words: rewards both more words and surer words."""
    confs = [c for _, t, c in _words(d) if any(ch.isalnum() for ch in t)]
    mean = sum(confs) / len(confs) if confs else 0.0
    return sum(confs) / 100.0, mean


def _looks_like_text(line: str) -> bool:
    letters = sum(ch.isalpha() for ch in line)
    solid = len(line.replace(" ", ""))
    return letters >= 2 and letters / solid >= 0.5


def _join_lines(lines):
    text = lines[0]
    for line in lines[1:]:
        if len(text) > 1 and text[-1] in HYPHENS and text[-2].isalpha() and line[:1].islower():
            text = text[:-1] + line           # "exam-" + "ple" -> "example"
        else:
            text += " " + line
    return text


def _merge_paragraphs(paras):
    """Tesseract often splits one paragraph into several on photos. A real paragraph
    ends with punctuation, so glue on whatever follows one that doesn't, unless it
    looks like a short heading followed by a new sentence."""
    out = []
    for p in paras:
        prev = out[-1] if out else ""
        heading = len(prev) < 40 and p[:1].isupper()
        if prev and not SENTENCE_END.search(prev) and not heading:
            out[-1] = _join_lines([prev, p])
        else:
            out.append(p)
    return out


def assemble(d) -> str:
    """Tesseract word boxes -> paragraphs separated by blank lines, noise lines dropped."""
    paras = {}
    for i, word, _ in _words(d):
        key = (d["block_num"][i], d["par_num"][i])
        paras.setdefault(key, {}).setdefault(d["line_num"][i], []).append(word)
    out = []
    for lines in paras.values():
        kept = [" ".join(ws) for ws in lines.values()]
        kept = [line for line in kept if _looks_like_text(line)]
        if kept:
            out.append(_join_lines(kept))
    return re.sub(r"[ \t]+", " ", "\n\n".join(_merge_paragraphs(out))).strip()


def single_script_langs(langs: str, text: str):
    """'eng+rus' on a page that is plainly all English -> 'eng' (and Russian -> 'rus').
    Mixed language packs misread some words, e.g. English "It is" as "№ 15"."""
    parts = langs.split("+")
    cyr = sum("Ѐ" <= ch <= "ӿ" for ch in text)
    lat = sum(ch.isascii() and ch.isalpha() for ch in text)
    if len(parts) < 2 or cyr + lat < 20:
        return None
    if cyr > 0.9 * (cyr + lat):
        keep = [l for l in parts if l in CYRILLIC_LANGS]
    elif lat > 0.9 * (cyr + lat):
        keep = [l for l in parts if l not in CYRILLIC_LANGS]
    else:
        return None
    return "+".join(keep) if keep and len(keep) < len(parts) else None


def read_page(data: bytes, langs: str = "eng+rus", psm: int = 3, rotate: int = 0) -> dict:
    """Photo bytes -> {'text', 'conf', 'variant', 'ms'}."""
    t0 = time.time()
    gray = resize(load_gray(data, rotate))
    imgs = variants(gray)
    jobs = {name: _pool.submit(_read, img, langs, psm) for name, img in imgs.items()}
    best = None
    for name, job in jobs.items():
        d = job.result()
        score, mean = _score(d)
        if best is None or score > best[0]:
            best = (score, mean, name, d)
    score, mean, name, d = best

    narrow = single_script_langs(langs, assemble(d))
    if narrow:                       # second pass with just that script's languages
        d2 = _read(imgs[name], narrow, psm)
        score2, mean2 = _score(d2)
        if score2 >= score:
            mean, name, d = mean2, f"{name} {narrow}", d2
    return {"text": assemble(d), "conf": round(mean, 1), "variant": name,
            "ms": int((time.time() - t0) * 1000)}
