"""Deskew a scanned document image with OpenCV — only run when Tesseract's
first pass on the original image already came back too thin to use (see
OcrPipeline::recognizePage), not on every upload.

Usage: python deskew.py <input path> <output path>

Finds the dominant text-block orientation via the minimum-area bounding
rectangle of thresholded foreground pixels, then rotates the image to
straighten it.

Exit codes (the PHP caller — DeskewService — branches on these):
  0 = corrected; the straightened, cropped image was written to <output path>
  1 = the image couldn't be read
  2 = malformed invocation (wrong argument count)
  3 = no correction needed (skew below the noise threshold); nothing written
"""

import sys

import cv2
import numpy as np

# Below this, the correction is noise (JPEG artifacts, antialiasing) rather
# than real skew, and rotating would just blur the page for no gain.
MIN_CORRECTABLE_ANGLE_DEGREES = 0.3

# Kept around every edge of the cropped result so a straightened line of
# text isn't left touching the image border.
CROP_MARGIN_PX = 15

EXIT_CORRECTED = 0
EXIT_UNREADABLE = 1
EXIT_BAD_ARGS = 2
EXIT_NO_SKEW = 3


def detect_skew_angle(gray: np.ndarray) -> float:
    inverted = cv2.bitwise_not(gray)
    thresh = cv2.threshold(inverted, 0, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1]

    coords = np.column_stack(np.where(thresh > 0))
    if coords.size == 0:
        return 0.0

    angle = cv2.minAreaRect(coords)[-1]

    # cv2.minAreaRect returns an angle in (-90, 0]; normalize to the
    # nearest-to-horizontal rotation needed to straighten the page.
    if angle < -45:
        angle = -(90 + angle)
    else:
        angle = -angle

    return angle


def crop_to_content(image: np.ndarray, margin: int = CROP_MARGIN_PX) -> np.ndarray:
    """Trims the blank canvas rotation adds around the page's corners
    (expanding the canvas to fit a tilted rectangle leaves whitespace
    triangles at each corner) down to just the page content plus a small
    margin. A no-op if the page is blank or fills the whole frame already.
    """
    gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
    non_white = cv2.threshold(gray, 250, 255, cv2.THRESH_BINARY_INV)[1]

    coords = cv2.findNonZero(non_white)
    if coords is None:
        return image

    x, y, w, h = cv2.boundingRect(coords)
    x0, y0 = max(x - margin, 0), max(y - margin, 0)
    x1 = min(x + w + margin, image.shape[1])
    y1 = min(y + h + margin, image.shape[0])

    return image[y0:y1, x0:x1]


def main() -> int:
    if len(sys.argv) != 3:
        print("usage: deskew.py <input path> <output path>", file=sys.stderr)
        return EXIT_BAD_ARGS

    input_path, output_path = sys.argv[1], sys.argv[2]

    image = cv2.imread(input_path)
    if image is None:
        print(f"could not read image: {input_path}", file=sys.stderr)
        return EXIT_UNREADABLE

    gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
    angle = detect_skew_angle(gray)

    if abs(angle) < MIN_CORRECTABLE_ANGLE_DEGREES:
        return EXIT_NO_SKEW

    (h, w) = image.shape[:2]
    center = (w // 2, h // 2)
    rotation_matrix = cv2.getRotationMatrix2D(center, angle, 1.0)

    # Bounding box grows when rotating a non-square image; expand the
    # canvas so the rotation doesn't crop content off the edges.
    cos = abs(rotation_matrix[0, 0])
    sin = abs(rotation_matrix[0, 1])
    new_w = int((h * sin) + (w * cos))
    new_h = int((h * cos) + (w * sin))
    rotation_matrix[0, 2] += (new_w / 2) - center[0]
    rotation_matrix[1, 2] += (new_h / 2) - center[1]

    # A plain white fill (rather than e.g. BORDER_REPLICATE) for the
    # corners the expanded canvas adds — matches an actual scanned page's
    # background and is what makes crop_to_content's whitespace detection
    # below reliable.
    rotated = cv2.warpAffine(
        image,
        rotation_matrix,
        (new_w, new_h),
        flags=cv2.INTER_CUBIC,
        borderMode=cv2.BORDER_CONSTANT,
        borderValue=(255, 255, 255),
    )
    rotated = crop_to_content(rotated)

    cv2.imwrite(output_path, rotated)
    return EXIT_CORRECTED


if __name__ == "__main__":
    sys.exit(main())
