#!/usr/bin/env python3
"""Unified field detector: crossword/letter cells (lattice+edge-coverage),
boxed answers, score boxes, and writing lines. Empty-region aware.
Outputs fields.json and render/fld_pNN.png overlays of accepted fields."""
import fitz, json
from collections import defaultdict
from statistics import median

doc = fitz.open("data/Level2_WhatChristiansBelieve.pdf")

def merge_iv(ivs, gap=1.5):
    ivs = sorted(ivs)
    out = []
    for a, b in ivs:
        if out and a <= out[-1][1] + gap:
            out[-1][1] = max(out[-1][1], b)
        else:
            out.append([a, b])
    return out

def coverage(ivs, lo, hi):
    if hi <= lo: return 0.0
    s = 0.0
    for a, b in ivs:
        s += max(0.0, min(b, hi) - max(a, lo))
    return s / (hi - lo)

def snap_groups(vals, tol=2.5):
    """cluster scalar coords; return {rounded_center: members}"""
    vals = sorted(vals)
    groups = []
    for v in vals:
        if groups and v - groups[-1][-1] <= tol:
            groups[-1].append(v)
        else:
            groups.append([v])
    return {round(sum(g)/len(g),1): g for g in groups}

def has_text(page, rect, pad=0.5):
    r = fitz.Rect(rect.x0-pad, rect.y0-pad, rect.x1+pad, rect.y1+pad)
    t = page.get_text("text", clip=r).strip()
    return len(t.replace("_","").replace(".","").replace("·","").replace("-","").strip()) > 0

def cell_has_letter(page, rect, pad=-1.0):
    """True if a grid cell holds a printed LETTER (given answer / word-search).
    A digit-only cell is a crossword clue number -> still fillable, returns False."""
    r = fitz.Rect(rect.x0-pad, rect.y0-pad, rect.x1+pad, rect.y1+pad)
    t = page.get_text("text", clip=r)
    return any(ch.isalpha() for ch in t)

def rect_squares(page):
    """Square-ish rectangles drawn as 're' items -> candidate letter cells."""
    out=[]
    for dr in page.get_drawings():
        for it in dr["items"]:
            if it[0]=="re":
                r=it[1]; w,h=r.width,r.height
                if 11<=w<=46 and 11<=h<=46 and abs(w-h)<8:
                    out.append(fitz.Rect(r))
    return out

def lattice_cells(page, mod_hint=None):
    """Reconstruct cells from stroke segments via edge coverage on a lattice."""
    Hsegs=defaultdict(list); Vsegs=defaultdict(list); rawH=[]; rawV=[]
    for dr in page.get_drawings():
        for it in dr["items"]:
            if it[0]=="l":
                p1,p2=it[1],it[2]
                if abs(p1.y-p2.y)<1.2 and abs(p1.x-p2.x)>=2:
                    rawH.append((min(p1.x,p2.x),max(p1.x,p2.x),(p1.y+p2.y)/2))
                elif abs(p1.x-p2.x)<1.2 and abs(p1.y-p2.y)>=2:
                    rawV.append((min(p1.y,p2.y),max(p1.y,p2.y),(p1.x+p2.x)/2))
            elif it[0]=="re":
                r=it[1]
                rawH.append((r.x0,r.x1,r.y0)); rawH.append((r.x0,r.x1,r.y1))
                rawV.append((r.y0,r.y1,r.x0)); rawV.append((r.y0,r.y1,r.x1))
    if not rawH or not rawV: return []
    yc=snap_groups([h[2] for h in rawH]); xc=snap_groups([v[2] for v in rawV])
    def nearest(c,val):
        best=min(c,key=lambda k:abs(k-val)); return best if abs(best-val)<=2.5 else None
    for x0,x1,y in rawH:
        k=nearest(yc,y)
        if k is not None: Hsegs[k].append((x0,x1))
    for y0,y1,x in rawV:
        k=nearest(xc,x)
        if k is not None: Vsegs[k].append((y0,y1))
    Hsegs={k:merge_iv(v) for k,v in Hsegs.items()}
    Vsegs={k:merge_iv(v) for k,v in Vsegs.items()}
    xs=sorted(Vsegs); ys=sorted(Hsegs)
    if mod_hint: mod=mod_hint
    else:
        diffs=[round(b-a) for arr in (xs,ys) for a,b in zip(arr,arr[1:]) if 12<=b-a<=42]
        if not diffs: return []
        mod=median(sorted(diffs))
    # gridline pairs ~one module apart (skip spurious intermediate lines from
    # unrelated rects e.g. right-margin score boxes)
    def cell_pairs(arr):
        out=[]
        for a in range(len(arr)):
            for b in range(a+1,len(arr)):
                d=arr[b]-arr[a]
                if d>mod*1.35: break
                if mod*0.7<=d<=mod*1.35: out.append((arr[a],arr[b]))
        return out
    xpairs=cell_pairs(xs); ypairs=cell_pairs(ys)
    cells=[]
    for x0,x1 in xpairs:
        for y0,y1 in ypairs:
            ix0,ix1=x0+0.8,x1-0.8; iy0,iy1=y0+0.8,y1-0.8
            cov=sorted([coverage(Hsegs[y0],ix0,ix1),coverage(Hsegs[y1],ix0,ix1),
                        coverage(Vsegs[x0],iy0,iy1),coverage(Vsegs[x1],iy0,iy1)])
            # all 4 edges fairly covered, OR 3 strong edges (interlocking shared edge)
            if cov[0]>=0.55 or (cov[1]>=0.7 and cov[0]>=0.25):
                cells.append(fitz.Rect(x0,y0,x1,y1))
    return cells

def detect_cells(page):
    sq=rect_squares(page)
    mod=median(sorted([(r.width+r.height)/2 for r in sq])) if sq else None
    cells=list(sq)+ [c for c in lattice_cells(page,mod)]
    # dedupe by center proximity
    uniq=[]
    for c in cells:
        cx,cy=(c.x0+c.x1)/2,(c.y0+c.y1)/2
        if any(abs(cx-(u.x0+u.x1)/2)<4 and abs(cy-(u.y0+u.y1)/2)<4 for u in uniq):
            continue
        uniq.append(c)
    cells=uniq
    # connected components using each pair's OWN size (handles mixed-size grids)
    n=len(cells); adj=defaultdict(set)
    for a in range(n):
        for b in range(a+1,n):
            ra,rb=cells[a],cells[b]
            sa,sb=ra.width,rb.width; savg=(sa+sb)/2
            if abs(sa-sb)>savg*0.4: continue          # different-sized -> different grid
            dx=abs(ra.x0-rb.x0); dy=abs(ra.y0-rb.y0)
            row = dy<savg*0.35 and abs(dx-savg)<savg*0.55   # horizontal neighbour
            col = dx<savg*0.35 and abs(dy-savg)<savg*0.55   # vertical neighbour
            if row or col:
                adj[a].add(b); adj[b].add(a)
    seen=set(); grids=[]
    for a in range(n):
        if a in seen: continue
        st=[a]; comp=[]
        while st:
            k=st.pop()
            if k in seen: continue
            seen.add(k); comp.append(k); st.extend(adj[k]-seen)
        grids.append([cells[k] for k in comp])
    return [g for g in grids if len(g)>=3]

def is_dark(col):
    """A black/dark stroke (answer rule) vs a colored decorative divider."""
    if col is None: return True
    return max(col) < 0.5

def detect_lines_boxes(page):
    lines=[]; boxes=[]
    for dr in page.get_drawings():
        stroke=dr.get("color"); fill=dr.get("fill")
        for it in dr["items"]:
            if it[0]=="l":
                p1,p2=it[1],it[2]
                if abs(p1.y-p2.y)<1.5 and abs(p1.x-p2.x)>20:
                    lines.append((min(p1.x,p2.x),max(p1.x,p2.x),(p1.y+p2.y)/2))
            elif it[0]=="re":
                r=it[1]; w,h=r.width,r.height
                if h<2.5 and w>20:
                    lines.append((r.x0,r.x1,(r.y0+r.y1)/2))
                elif w>=28 and h>=8 and w<560 and h<340:
                    boxes.append(fitz.Rect(r))
    # merge collinear line segments
    lines=sorted(lines,key=lambda L:(round(L[2]/2),L[0]))
    merged=[]
    for x0,x1,y in lines:
        for o in merged:
            if abs(o[2]-y)<=2 and not (x1<o[0]-6 or x0>o[1]+6):
                o[0]=min(o[0],x0); o[1]=max(o[1],x1); o[2]=(o[2]+y)/2; break
        else:
            merged.append([x0,x1,y])
    merged=[L for L in merged if L[1]-L[0]>=30]
    return merged, boxes

def art_regions(page):
    """Raster images + dense vector illustrations -> exclusion zones."""
    regs=[]
    for im in page.get_image_info():
        regs.append(fitz.Rect(im["bbox"]))
    for dr in page.get_drawings():
        cs=sum(1 for it in dr["items"] if it[0]=="c")
        if cs>=4 and dr.get("rect"):
            regs.append(fitz.Rect(dr["rect"]))
    return regs

def hits_art(rect, arts, frac=0.35):
    a=rect.get_area()
    if a<=0: return False
    for r in arts:
        if rect.intersects(r) and (rect & r).get_area() > frac*a:
            return True
    return False

SKIP_PAGES={23}          # back cover: no fields

fields={}
for pi,page in enumerate(doc):
    if pi in SKIP_PAGES:
        fields[pi]=[]; continue
    flist=[]; used=[]
    arts=art_regions(page)
    # ---- cover: Name / Teacher (anchor to labels inside the white box) ----
    if pi==0:
        wb=None
        for dr in page.get_drawings():
            f=dr.get("fill"); r=dr.get("rect")
            if f and min(f)>0.9 and r and r.width>100 and r.height>20:
                wb=fitz.Rect(r); break
        for label,fname in (("Name:","child_name"),("Teacher:","teacher")):
            rs=[r for r in page.search_for(label) if (wb is None or wb.intersects(r))]
            if rs:
                r=rs[0]
                x1=(wb.x1-3) if wb else r.x1+240
                flist.append({"type":"text","name":fname,
                              "rect":[round(r.x1+4,1),round(r.y0-1,1),round(x1,1),round(r.y1+1,1)]})
        fields[pi]=flist; continue
    # ---- crossword/letter grids (multi-cell) ----
    grids=detect_cells(page)
    for gi,g in enumerate(grids):
        if sum(1 for c in g if cell_has_letter(page,c))>0.5*len(g):   # word search/solved
            continue
        g=sorted(g,key=lambda c:(round(c.y0/5),c.x0))
        for ci,c in enumerate(g):
            used.append(c)
            if cell_has_letter(page,c) or hits_art(c,arts): continue   # given letter -> skip; digit/empty -> field
            flist.append({"type":"char","name":f"p{pi:02d}_cw{gi}_{ci}",
                          "rect":[round(c.x0,1),round(c.y0,1),round(c.x1,1),round(c.y1,1)]})
    # ---- standalone small squares: one-char cells, or score boxes (right margin) ----
    for k,c in enumerate(rect_squares(page)):
        if any(abs(c.x0-u.x0)<3 and abs(c.y0-u.y0)<3 for u in used): continue
        used.append(c)
        if has_text(page,c) or hits_art(c,arts): continue
        if c.x0>495:    # right margin
            if c.y0<78 or c.width>34:   # top-right header logo, not a score box
                continue
            flist.append({"type":"text","name":f"p{pi:02d}_mark{k}",
                          "rect":[round(c.x0,1),round(c.y0,1),round(c.x1,1),round(c.y1,1)]})
        else:
            flist.append({"type":"char","name":f"p{pi:02d}_ch{k}",
                          "rect":[round(c.x0,1),round(c.y0,1),round(c.x1,1),round(c.y1,1)]})
    # ---- larger boxes -> free text ----
    lines,boxes=detect_lines_boxes(page)
    def in_used(rect,frac=0.3):
        a=rect.get_area()
        return a>0 and any(rect.intersects(u) and (rect&u).get_area()>frac*a for u in used)
    for k,b in enumerate(boxes):
        if b.width<46 and b.height<46: continue   # squares handled above
        if in_used(b) or has_text(page,b) or hits_art(b,arts): continue
        used.append(b)
        flist.append({"type":"text","name":f"p{pi:02d}_box{k}",
                      "rect":[round(b.x0,1),round(b.y0,1),round(b.x1,1),round(b.y1,1)]})
    # ---- writing lines -> free text ----
    for k,(x0,x1,y) in enumerate(lines):
        if y<60: continue
        if x0<66 and x1>529: continue        # full-width decorative rule / divider
        band=fitz.Rect(x0,y-13,x1,y-1.5)
        fld=fitz.Rect(x0,y-14,x1,y-1.0)
        if in_used(fld) or has_text(page,band) or hits_art(fld,arts): continue
        used.append(fld)
        flist.append({"type":"text","name":f"p{pi:02d}_line{k}",
                      "rect":[round(x0,1),round(y-14,1),round(x1,1),round(y-1.0,1)]})
    fields[pi]=flist

json.dump({str(k):v for k,v in fields.items()},open("fields.json","w"),indent=0)

# overlay: fill chars (visible) + outline text
doc2=fitz.open("data/Level2_WhatChristiansBelieve.pdf")
tot=0
for pi,page in enumerate(doc2):
    for f in fields[pi]:
        r=fitz.Rect(*f["rect"])
        if f["type"]=="char":
            page.draw_rect(r,color=(0,0,0.85),fill=(0.6,0.75,1),fill_opacity=0.45,width=0.6)
        else:
            page.draw_rect(r,color=(0,0.55,0),width=1.2,fill=(0.6,1,0.6),fill_opacity=0.30)
    page.get_pixmap(matrix=fitz.Matrix(150/72,150/72)).save(f"render/fld_p{pi:02d}.png")
    tot+=len(fields[pi])
for pi in range(doc.page_count):
    t=sum(1 for f in fields[pi] if f["type"]=="text"); c=sum(1 for f in fields[pi] if f["type"]=="char")
    print(f"p{pi:2d}: text={t:2d} char={c:2d}")
print("TOTAL",tot)
