"""Per-document comparison of main vs branch extractions: table counts and grid shapes per page."""
import json, glob, collections, sys
OUT='/home/mande/repo/quber/data/que370-review/sample'
docs=json.load(open(f'{OUT}/sample10.json'))
shape=lambda t: f"{len(t.get('corrected_grid') or [])}x{len((t.get('corrected_grid') or [[]])[0])}"
raw=lambda t: f"{len(t.get('cell_grid') or [])}x{len((t.get('cell_grid') or [[]])[0])}"
tot_same=tot_pages=0; rows=[]
for o in docs:
    d=f"{OUT}/{o['folder']}-{o['key']}"
    fm=glob.glob(f'{d}/main/*.tables.json'); fb=glob.glob(f'{d}/branch/*.tables.json')
    if not (fm and fb): rows.append((o['folder'],o['filing'],o['pages'],'pending')); continue
    m=json.load(open(fm[0])); b=json.load(open(fb[0]))
    mp=collections.defaultdict(list); bp=collections.defaultdict(list)
    for t in m: mp[t['page']].append(t)
    for t in b: bp[t['page']].append(t)
    pages=sorted(set(mp)|set(bp)); same=0; count_diff=[]; shape_diff=[]
    for pg in pages:
        ms=[shape(t) for t in sorted(mp[pg], key=lambda t:t['som_region'][1])]; bs=[shape(t) for t in sorted(bp[pg], key=lambda t:t['som_region'][1])]
        mr=[raw(t) for t in sorted(mp[pg], key=lambda t:t['som_region'][1])]; br=[raw(t) for t in sorted(bp[pg], key=lambda t:t['som_region'][1])]
        if ms==bs and mr!=br: shape_diff.append((pg,'raw',mr,br)); continue
        if ms==bs: same+=1
        elif len(ms)!=len(bs): count_diff.append((pg,len(ms),len(bs)))
        else: shape_diff.append((pg,ms,bs))
    tot_same+=same; tot_pages+=len(pages)
    rows.append((o['folder'],o['filing'],o['pages'],f"tables {len(m)}->{len(b)} | table pages {len(pages)} | identical {same} | count changes {count_diff} | shape-only changes {len(shape_diff)}"))
    if '-v' in sys.argv and shape_diff: print(o['folder'], "shape-only:", shape_diff)
for r in rows: print(f"{r[0]:5s} {r[1]:5s} {r[2]:4d}p | {r[3]}")
print(f"\nTOTAL table pages {tot_pages} | identical count and shapes {tot_same}")
