"""Demand-driven footnote resolution from the text layer — no vision finder.

The tables are already parsed. For each corrected table we read the markers it
references (REFERENCE_MARKER, on the corrected markdown). For each referenced
marker we then resolve its DEFINITION deterministically from the page text
layer: a line that BEGINS with that marker followed by prose. That start-of-line
test excludes table rows (which start with a row label) and reference headers
(which carry the marker at the END, e.g. "Loans Past Due ... (1) :"), so the
parsed table cannot interfere — there is no vision step to mis-locate.

The footnote box is the bounding box of the definition line plus its wrapped
continuation lines. Output: an annotated PDF per doc (blue = table som_region,
orange = resolved footnote definition) and a PNG per candidate page.

Run:
    uv run python experiments/que270/demand_driven_textlayer.py
"""

from __future__ import annotations

import glob
import json
import re
import sys
from collections import defaultdict
from pathlib import Path
from typing import Dict, List, Optional, Tuple

import fitz
from loguru import logger

sys.path.insert(0, str(Path(__file__).resolve().parent))
import annotate_regions as ar  # noqa: E402  (reuse colours + _rect)

from quber.agents.completeness import page_words  # noqa: E402

CORPUS = Path("experiments/que270/out/strip_test_corpus")
OUT_DIR = Path("experiments/que270/out/demand_driven_textlayer")

# A marker REFERENCE attached to a label ("Income Tax Provision(1)"); lookbehind
# excludes a bare parenthesized negative value ("(11)") in a numeric cell.
REFERENCE_MARKER = re.compile(r"(?<=[A-Za-z%)])\((\d{1,2})\)")
REFERENCE_SYMBOL = re.compile(r"(?<=[A-Za-z%)])([*†‡§])")

# A footnote DEFINITION line: begins with a marker, then prose. The trailing
# prose check (a letter after the marker) rejects a numeric row that merely
# opens with a parenthesized value.
DEFINITION_LINE = re.compile(r"^\s*(?:\((\d{1,2})\)|(\d{1,2})[.)]|([*†‡§]))\s*(?=[A-Za-z])")


def referenced_markers(corrected_markdown: str) -> set:
    m = set(REFERENCE_MARKER.findall(corrected_markdown))
    m.update(REFERENCE_SYMBOL.findall(corrected_markdown))
    return m


def _group_lines(words) -> List[list]:
    out: List[list] = []
    for w in sorted(words, key=lambda w: (w[1] + w[3]) / 2):
        yc = (w[1] + w[3]) / 2
        if out and abs(yc - (out[-1][0][1] + out[-1][0][3]) / 2) <= 4.0:
            out[-1].append(w)
        else:
            out.append([w])
    for ln in out:
        ln.sort(key=lambda w: w[0])
    return out


def _line_marker(line) -> Optional[str]:
    text = " ".join(w[4] for w in line)
    m = DEFINITION_LINE.match(text)
    if not m:
        return None
    return next((g for g in m.groups() if g is not None), None)


def resolve_definitions(words, page_w: float, page_h: float) -> Dict[str, Tuple[Tuple[float, float, float, float], str]]:
    """Marker -> (normalized box, text) for every footnote definition on the page.

    A definition is a line that begins with a marker; immediately-following lines
    that do not themselves begin with a marker are folded in as wrapped
    continuation, so a box covers the whole multi-line note.
    """
    lines = _group_lines(words)
    defs: Dict[str, List[list]] = {}
    current: Optional[str] = None
    prev_bottom = None
    for line in lines:
        marker = _line_marker(line)
        top = min(w[1] for w in line)
        if marker is not None:
            current = marker
            defs.setdefault(marker, [line])
            if defs[marker][0] is not line:
                defs[marker].append(line)
            prev_bottom = max(w[3] for w in line)
        elif current is not None and prev_bottom is not None and 0 <= top - prev_bottom <= 14:
            defs[current].append(line)  # wrapped continuation
            prev_bottom = max(w[3] for w in line)
        else:
            current = None
    out: Dict[str, Tuple[Tuple[float, float, float, float], str]] = {}
    for marker, mlines in defs.items():
        flat = [w for ln in mlines for w in ln]
        x0 = min(w[0] for w in flat) / page_w
        y0 = min(w[1] for w in flat) / page_h
        x1 = max(w[2] for w in flat) / page_w
        y1 = max(w[3] for w in flat) / page_h
        text = " ".join(w[4] for ln in mlines for w in ln)
        out[marker] = ((x0, y0, x1, y1), text)
    return out


def annotate_doc(tables_json: Path) -> dict:
    tables = json.loads(tables_json.read_text())
    if not tables:
        return {"doc": tables_json.stem, "candidate_pages": []}
    pdf = Path(tables[0]["source"])
    by_page: Dict[int, List[dict]] = defaultdict(list)
    for t in tables:
        by_page[t["page"]].append(t)

    doc = fitz.open(str(pdf))
    candidate_pages: List[int] = []
    findings: List[dict] = []

    for page in sorted(by_page):
        page_tables = by_page[page]
        ref_by_table: Dict[int, set] = {}
        page_refs: set = set()
        for i, t in enumerate(page_tables):
            r = referenced_markers(t.get("markdown", ""))
            ref_by_table[i + 1] = r
            page_refs |= r
        tbl_boxes = [
            (i + 1, t.get("title", ""), t["som_region"])
            for i, t in enumerate(page_tables)
            if t.get("som_region")
        ]
        fn_boxes: List[Tuple[str, Tuple[float, float, float, float]]] = []
        if page_refs:
            candidate_pages.append(page)
            pw, ph, words = page_words(pdf, page)
            defs = resolve_definitions(words, pw, ph)
            found = set()
            for marker in sorted(page_refs):
                if marker in defs:
                    box, text = defs[marker]
                    fn_boxes.append((marker, box))
                    found.add(marker)
            owners = {m: [o for o, r in ref_by_table.items() if m in r] for m in found}
            findings.append({
                "doc": tables_json.stem, "page": page,
                "referenced": sorted(page_refs), "resolved": sorted(found),
                "missing": sorted(page_refs - found), "owners": owners,
                "defs": {m: defs[m][1][:80] for m in found},
            })
        _draw(doc, page, tbl_boxes, fn_boxes, ref_by_table)

    OUT_DIR.mkdir(parents=True, exist_ok=True)
    doc.save(str(OUT_DIR / f"{tables_json.stem}.dd.pdf"))
    for page in candidate_pages:
        doc[page - 1].get_pixmap(dpi=120).save(str(OUT_DIR / f"{tables_json.stem}_p{page}.png"))
    doc.close()
    logger.info("{}: candidate pages {}", pdf.name, candidate_pages)
    return {"doc": tables_json.stem, "candidate_pages": candidate_pages, "findings": findings}


def _draw(doc, page_no, tbl_boxes, fn_boxes, ref_by_table) -> None:
    page = doc[page_no - 1]
    w, h = page.rect.width, page.rect.height
    for ordinal, _title, region in tbl_boxes:
        rect = ar._rect(region, w, h)
        page.draw_rect(rect, color=ar.TABLE_RGB, width=1.4)
        refs = ref_by_table.get(ordinal) or set()
        lab = f"T{page_no}.{ordinal}" + (f" refs={sorted(refs)}" if refs else "")
        page.insert_text((rect.x0 + 2, max(8.0, rect.y0 - 3)), lab[:80], fontsize=6, color=ar.TABLE_RGB)
    for marker, region in fn_boxes:
        rect = ar._rect(region, w, h)
        page.draw_rect(rect, color=ar.FOOTNOTE_RGB, width=1.4, dashes="[3 2] 0")
        page.insert_text((rect.x0 + 2, min(h - 2, rect.y1 + 7)), f"fn ({marker})", fontsize=6, color=ar.FOOTNOTE_RGB)


def main() -> None:
    files = sorted(glob.glob(str(CORPUS / "*" / "*.tables.json")))
    index = [annotate_doc(Path(f)) for f in files]
    OUT_DIR.mkdir(parents=True, exist_ok=True)
    (OUT_DIR / "index.json").write_text(json.dumps(index, indent=2))
    findings = [fd for d in index for fd in d.get("findings", [])]
    ref = sum(len(f["referenced"]) for f in findings)
    res = sum(len(f["resolved"]) for f in findings)
    miss = sum(len(f["missing"]) for f in findings)
    print(f"docs scanned: {len(index)}   candidate pages: {sum(len(d['candidate_pages']) for d in index)}")
    print(f"referenced markers: {ref}   resolved: {res}   missing: {miss}")
    print("\nPages with missing markers:")
    for f in findings:
        if f["missing"]:
            print(f"  {f['doc'][:30]:30} p{f['page']:<4} ref={f['referenced']} resolved={f['resolved']} missing={f['missing']}")
    print(f"\nOutput under {OUT_DIR}")


if __name__ == "__main__":
    main()
