#!/usr/bin/env python3
"""
Trích xuất text từ JSON Cloud Vision → Markdown (V5 — Robust column detection).
- Detect columns dựa trên x_min (cạnh trái) → chính xác hơn x_mid
- Phân biệt TOC (row-by-row) vs 2-cột độc lập (column-by-column)
- Split multi-line blocks thành từng dòng riêng

Usage:
  python3 json_to_markdown_v5.py <json_dir> <output_dir> [--dry-run] [--debug]
"""

import os
import sys
import json
import glob
import re
import math
from collections import Counter

# ─── Constants ───────────────────────────────────────────────────────────────

SCANNER_PATTERNS = [
    re.compile(r'scanned\s+with', re.IGNORECASE),
    re.compile(r'cam\s*scanner', re.IGNORECASE),
]

# Column detection
MIN_COLUMN_GAP = 0.06        # Minimum X gap between column clusters (x_min based)
MIN_COLUMN_WIDTH = 0.05      # Minimum width of column cluster
MIN_BLOCKS_PER_COLUMN = 3    # Minimum items per column
COLUMN_X_MAX_OVERLAP = 0.03  # Max overlap between column X ranges

Y_OVERLAP_THRESHOLD = 0.15


# ─── Helpers ─────────────────────────────────────────────────────────────────

def extract_paragraph_text(para):
    text = ''
    for word in para.get('words', []):
        for sym in word.get('symbols', []):
            text += sym.get('text', '')
            prop = sym.get('property', {})
            if prop.get('detectedBreak'):
                bt = prop['detectedBreak']['type']
                if bt in ('SPACE', 'SURE_SPACE'):
                    text += ' '
                elif bt in ('EOL_SURE_SPACE', 'LINE_BREAK'):
                    text += '\n'
    return text


def get_bounding_box(obj):
    verts = obj.get('boundingBox', {}).get('normalizedVertices', [])
    if not verts:
        return (0, 0, 0, 0)
    x_min = min(v.get('x', 0) for v in verts)
    y_min = min(v.get('y', 0) for v in verts)
    x_max = max(v.get('x', 0) for v in verts)
    y_max = max(v.get('y', 0) for v in verts)
    return (x_min, y_min, x_max, y_max)


def is_scanner_artifact(text):
    for pat in SCANNER_PATTERNS:
        if pat.search(text):
            return True
    return False


def y_overlap(y1_min, y1_max, y2_min, y2_max, threshold=Y_OVERLAP_THRESHOLD):
    overlap_start = max(y1_min, y2_min)
    overlap_end = min(y1_max, y2_max)
    overlap = max(0, overlap_end - overlap_start)
    r1 = max(y1_max - y1_min, 0.001)
    r2 = max(y2_max - y2_min, 0.001)
    return overlap / min(r1, r2)


# ─── Column Detection ────────────────────────────────────────────────────────

def detect_columns(paragraphs):
    """Detect if page has multi-column layout using x_min clustering.
    
    Uses left-edge (x_min) position which separates columns better than
    midpoints when text spans across column boundaries.
    
    Returns (num_columns, [(col_x_min, col_x_max), ...]) or (1, [(0,1)]) if single column.
    """
    if len(paragraphs) < 8:
        return 1, [(0, 1)]
    
    # Collect x_min and x_max for each paragraph
    # Skip very wide items (> 60% page width) as they're likely full-width
    points = []
    for x_min, y_min, x_max, y_max, text in paragraphs:
        width = x_max - x_min
        if width > 0.60:
            continue
        points.append((x_min, x_max, y_min, y_max))
    
    if len(points) < 8:
        return 1, [(0, 1)]
    
    # Cluster x_min values using simple gap detection
    x_mins = sorted(set(p[0] for p in points))
    
    # Find gaps between consecutive x_min values
    gaps = []
    for i in range(len(x_mins) - 1):
        gap = x_mins[i + 1] - x_mins[i]
        if gap > MIN_COLUMN_GAP:
            gaps.append((i, gap, x_mins[i], x_mins[i + 1]))
    
    if not gaps:
        return 1, [(0, 1)]
    
    # Find the best splitting point: prefer gaps with balanced clusters
    best_split = None
    best_score = 0
    
    for idx, gap, left_x, right_x in gaps:
        left_count = idx + 1
        right_count = len(x_mins) - idx - 1
        
        if left_count >= MIN_BLOCKS_PER_COLUMN and right_count >= MIN_BLOCKS_PER_COLUMN:
            # Score: prefer larger gaps with balanced columns
            balance = min(left_count, right_count) / max(left_count, right_count)
            score = gap * 10 + balance * 5
            if score > best_score:
                best_score = score
                best_split = (left_x, right_x)
    
    if best_split is None:
        return 1, [(0, 1)]
    
    boundary = (best_split[0] + best_split[1]) / 2
    
    # Assign points to columns based on x_min
    col1_points = [p for p in points if p[0] < boundary]
    col2_points = [p for p in points if p[0] >= boundary]
    
    if len(col1_points) < MIN_BLOCKS_PER_COLUMN or len(col2_points) < MIN_BLOCKS_PER_COLUMN:
        return 1, [(0, 1)]
    
    # Compute column X ranges (using x_max from assigned points)
    col1_x_min = min(p[0] for p in col1_points)
    col1_x_max = max(p[1] for p in col1_points)  # x_max
    col2_x_min = min(p[0] for p in col2_points)
    col2_x_max = max(p[1] for p in col2_points)
    
    # Check if columns overlap too much (would indicate TOC, not 2-column)
    overlap = max(0, col1_x_max - col2_x_min)
    if overlap > 0.15:  # Significant overlap → likely TOC, not independent columns
        return 1, [(0, 1)]
    
    return 2, [(col1_x_min, col1_x_max), (col2_x_min, col2_x_max)]


# ─── Layout Processing ───────────────────────────────────────────────────────

def split_multi_line_paragraphs(paragraphs):
    result = []
    for x_min, y_min, x_max, y_max, text in paragraphs:
        lines = text.split('\n')
        lines = [l.strip() for l in lines if l.strip()]
        if len(lines) <= 1:
            result.append((x_min, y_min, x_max, y_max, lines[0] if lines else text))
        else:
            num_lines = len(lines)
            line_height = (y_max - y_min) / num_lines
            for i, line in enumerate(lines):
                ly_min = y_min + i * line_height
                ly_max = ly_min + line_height
                result.append((x_min, ly_min, x_max, ly_max, line))
    return result


def group_into_rows(fragments):
    if not fragments:
        return []
    sorted_frags = sorted(fragments, key=lambda f: (f[1], f[0]))
    rows = []
    cur = [sorted_frags[0]]
    cy_min, cy_max = sorted_frags[0][1], sorted_frags[0][3]
    
    for frag in sorted_frags[1:]:
        _, y_min, _, y_max, _ = frag
        if y_overlap(cy_min, cy_max, y_min, y_max) > 0:
            cur.append(frag)
            cy_min = min(cy_min, y_min)
            cy_max = max(cy_max, y_max)
        else:
            cur.sort(key=lambda f: f[0])
            rows.append(cur)
            cur = [frag]
            cy_min, cy_max = y_min, y_max
    
    if cur:
        cur.sort(key=lambda f: f[0])
        rows.append(cur)
    return rows


def rows_to_text(rows, gap_threshold=0.15):
    lines = []
    for row in rows:
        parts = []
        prev_x_max = 0
        for i, (x_min, y_min, x_max, y_max, text) in enumerate(row):
            if i > 0:
                gap = x_min - prev_x_max
                if gap > gap_threshold:
                    parts.append('    ')
                elif gap > 0.03:
                    parts.append('  ')
                else:
                    parts.append(' ')
            parts.append(text)
            prev_x_max = x_max
        lines.append(''.join(parts))
    return lines


def page_to_text(page_data, page_num, debug=False):
    blocks = page_data.get('blocks', [])
    
    # Step 1: Extract paragraphs
    all_paragraphs = []
    for block in blocks:
        if block.get('blockType') not in (None, 'TEXT'):
            continue
        for para in block.get('paragraphs', []):
            text = extract_paragraph_text(para)
            text_stripped = text.strip()
            if not text_stripped:
                continue
            x_min, y_min, x_max, y_max = get_bounding_box(para)
            if is_scanner_artifact(text_stripped):
                continue
            all_paragraphs.append((x_min, y_min, x_max, y_max, text_stripped))
    
    if not all_paragraphs:
        return f"## PAGE {page_num}\n\n*(trang trống)*\n"
    
    # Step 2: Split multi-line blocks
    fragments = split_multi_line_paragraphs(all_paragraphs)
    
    # Step 3: Detect columns
    num_cols, col_bounds = detect_columns(all_paragraphs)
    
    output_lines = [f"## PAGE {page_num}", ""]
    
    if debug:
        output_lines.append(f"<!-- {num_cols} col(s), bounds={col_bounds} -->")
    
    if num_cols == 1:
        # Single column or TOC: row-by-row left-to-right
        rows = group_into_rows(fragments)
        output_lines.extend(rows_to_text(rows))
    else:
        # Multi-column: read each column top-to-bottom separately
        
        for col_idx, (col_x_min, col_x_max) in enumerate(col_bounds):
            # Get fragments belonging to this column (by x_min)
            col_fragments = []
            for frag in fragments:
                x_min = frag[0]
                if col_x_min - 0.02 <= x_min <= col_x_max + 0.02:
                    col_fragments.append(frag)
            
            col_fragments.sort(key=lambda f: (f[1], f[0]))
            
            if col_idx > 0:
                output_lines.append("")
            
            col_rows = group_into_rows(col_fragments)
            col_lines = rows_to_text(col_rows, gap_threshold=0.05)
            output_lines.extend(col_lines)
    
    output_lines.append("")
    return '\n'.join(output_lines)


def json_to_markdown(json_path, output_dir, start_page, debug=False):
    with open(json_path, 'r', encoding='utf-8') as f:
        data = json.load(f)

    basename = os.path.splitext(os.path.basename(json_path))[0]
    out_path = os.path.join(output_dir, f"{basename}.md")

    page_texts = []
    for i, response in enumerate(data.get('responses', [])):
        page_num = start_page + i
        fta = response.get('fullTextAnnotation', {})
        pages = fta.get('pages', [])
        
        if pages:
            page_text = page_to_text(pages[0], page_num, debug=debug)
        else:
            raw_text = fta.get('text', '').strip()
            page_text = f"## PAGE {page_num}\n\n{raw_text if raw_text else '*(trang trống)*'}\n"
        page_texts.append(page_text)

    os.makedirs(output_dir, exist_ok=True)
    with open(out_path, 'w', encoding='utf-8') as f:
        f.write('\n'.join(page_texts))

    return out_path, len(page_texts)


def main():
    if len(sys.argv) < 3:
        print("Usage: python3 json_to_markdown_v5.py <json_dir> <output_dir> [--dry-run] [--debug]")
        sys.exit(1)

    json_dir = sys.argv[1]
    output_dir = sys.argv[2]
    dry_run = '--dry-run' in sys.argv
    debug = '--debug' in sys.argv

    json_files = sorted(
        glob.glob(os.path.join(json_dir, '*.json')),
        key=lambda f: int(re.search(r'output-(\d+)-to', os.path.basename(f)).group(1))
    )

    if not json_files:
        print(f"❌ Không tìm thấy file JSON nào trong {json_dir}")
        sys.exit(1)

    print(f"📄 Tìm thấy {len(json_files)} file JSON")
    total_pages = 0

    for jf in json_files:
        out_path, pages = json_to_markdown(jf, output_dir, total_pages + 1, debug=debug)
        total_pages += pages
        
        if dry_run:
            with open(out_path, 'r') as f:
                preview = ''.join(f.readline() for _ in range(25))
            print(f"\n{'='*60}")
            print(f"📄 {os.path.basename(out_path)} ({pages} trang) — PREVIEW:")
            print(f"{'='*60}")
            print(preview)
            print(f"... (saved at {out_path})")
        else:
            print(f"   ✅ {os.path.basename(out_path)} ({pages} trang)")

    if not dry_run:
        print(f"\n🎉 Done! {len(json_files)} file .md → {os.path.abspath(output_dir)}")
        print(f"   Tổng: {total_pages} trang")


if __name__ == '__main__':
    main()
