# -*- coding: utf-8 -*-
"""把《上海校本卷》（各校月考/期中/期末整卷，2026-09-07 起录）拍平成组卷台归一化 schema，
作为新教材「上海校本卷」并入。

由 gen_组卷台.py 调用（与 zujuan_exam / zujuan_cn / zujuan_amc8 同款两处 hook）：
  data_recs, course_sub, fig_dict = zujuan_school.load(HERE)
  data += data_recs
  course['上海校本卷'] = course_sub
  FIG.update(fig_dict)
⚠ 资产 id 前缀 'xb/' 要加进 gen_组卷台.py 的 FIG_PFX；EXAM_TBS / preNarrow / AI 组卷提示词由组卷台窗口接。

## 数据来源
`数学学习系统/data/drafts/batch_school_20260907/`（八上 09-07 批）+ `batch_school_20260909/`（六上～八下 09-09 批，BATCHES 列表）：paper_extract.py（上海方言 + PE_DOC_DIRECT，清单驱动元数据）
从 `源题文件/数学资料/八上数学月考期中期末试卷合集（294套）` + `00张江集团中学试卷汇编` 提取，
exam_papers.school / source_metadata.grade / source_metadata.paper_kind 由 `校本卷20260907/_picks.json` 给定。

## 轴怎么摆（2026-09-07 与组卷台窗口 surface:3 对过口径）
  tb 教材 = 上海校本卷
  gg 年级 = 年级学期（八上/八下/七上…，与培优册名一致，前端年级识别认它）
  cd 节名 = 「2025学年·南洋模范中学·期中」（学年·学校·考次；卷商模拟卷学校位写「模拟卷」）；同年级撞名补 #2
  节点 4 元，第 4 位 = 「2025学年」→ 前端 optgroup 按学年分组
  s/src   = 考次（期中 / 期末 / 月考类：第N次月考、N月月考、阶段测试、复习卷）
  t/d     = 题型由大题头原文归成 选择题/填空题/解答题/其他（_qtype）；难度 0（校本卷不推断）
  pid     = 原卷 id（xb-<学年>-<mid|fin|mon>-<学校代号>-<g8a…>-math），「📄 原卷」按钮用

## 零答案
只读 stem_markdown（answer 早已摘进 answers_withheld.jsonl）。题干被解析污染的（故选/故答案为/∵∴开头）剔除，
与 zujuan_exam 同一套判据；答案原文 ≥4 字出现在题面里计入 leak 自查。
题面处理（表格注入 / 剥题号 / 裸上下标转 Unicode / WMF 兜底转 PNG）全部复用 zujuan_exam 的函数，别写第二份。
"""
import json
import os
import re
from collections import OrderedDict, defaultdict

from PIL import Image

from zujuan_exam import (ANS_HEAD, ANS_MARK, IMG_RE, _fig_b64, _inject_tables, _rescue_wmf,
                         _strip_lead_no, _unicode_scripts)

TB = '上海校本卷'
GRADE_ORDER = ['六上', '六下', '七上', '七下', '八上', '八下', '九上', '九下']
TYPE_ORDER = {'midterm': 1, 'final': 2, 'monthly': 0}
TYPE_NAME = {'midterm': '期中', 'final': '期末', 'monthly': '月考'}


def _qtype(raw):
    """question_type_raw 是大题头原文（会带「12*2=」「【本大题共14个小题」之类尾巴），归成四类给前端筛。"""
    t = str(raw or '')
    if '选择' in t or '单选' in t:
        return '选择题'
    if '填空' in t:
        return '填空题'
    if any(k in t for k in ('解答', '简答', '计算', '证明', '应用', '综合', '作图', '解方程', '化简', '操作', '探究')):
        return '解答题'
    return '其他'


BATCHES = ['batch_school_20260907', 'batch_school_20260909']      # 2026-09-09 第二批（上海名校 6-8 年级合集，六上～八下）；同 schools 代号表


def _find_batches(HERE):
    out = []
    for up in ('.', '..', '../..', '../../..'):
        root = os.path.normpath(os.path.join(HERE, up, '数学学习系统', 'data', 'drafts'))
        if os.path.isdir(root):
            for b in BATCHES:
                cand = os.path.join(root, b)
                if os.path.exists(os.path.join(cand, 'questions.jsonl')):
                    out.append(cand)
            if out:
                return out
    return out


def load(HERE, verbose=True):
    """返回 (data_records, course_sublist_by_grade, fig_base64_dict)。"""
    BS = _find_batches(HERE)
    if not BS:
        if verbose:
            print('  [校本] 未找到 batch_school_*/questions.jsonl，跳过')
        return [], OrderedDict(), {}
    papers, qs, asset_dir_of = {}, [], {}
    for B in BS:
        ps = {p['external_id']: p for p in
              map(json.loads, open(os.path.join(B, 'exam_papers.jsonl'), encoding='utf-8'))
              if p.get('record_status') == 'draft'}
        papers.update(ps)
        for pid in ps:
            asset_dir_of[pid] = os.path.join(B, 'assets')
        qs += [q for q in map(json.loads, open(os.path.join(B, 'questions.jsonl'), encoding='utf-8'))
               if q.get('paper_external_id') in ps]
    n_sub = sum(1 for q in qs if q.get('parent_external_id'))
    qs = [q for q in qs if not q.get('parent_external_id')]          # 小问不单独成卡（同中考批）
    polluted = {q['external_id'] for q in qs
                if ANS_MARK.search(q.get('stem_markdown') or '') or ANS_HEAD.match(q.get('stem_markdown') or '')}
    qs = [q for q in qs if q['external_id'] not in polluted]

    # 同题面精确去重（同一道题在组卷台只露一次）：公式已成 LaTeX 的卷优先留
    latex_of = defaultdict(int)
    for q in qs:
        if '\\(' in (q.get('stem_markdown') or ''):
            latex_of[q['paper_external_id']] += 1
    qs.sort(key=lambda q: (-latex_of[q['paper_external_id']], q['paper_external_id'], q.get('sort_order') or 0))
    seen_sha, kept, n_dup = set(), [], 0
    for q in qs:
        h = q.get('content_sha256')
        if h and h in seen_sha:
            n_dup += 1; continue
        if h:
            seen_sha.add(h)
        kept.append(q)

    def meta(p):
        sm = p.get('source_metadata') or {}
        return sm.get('grade') or '八上', p.get('school') or '模拟卷', sm.get('paper_kind') or TYPE_NAME.get(p['exam_type'], '')

    # 节名：学年·学校·考次；同年级内撞名补 #2
    label, per_grade = {}, defaultdict(list)
    live = sorted({q['paper_external_id'] for q in kept},
                  key=lambda pid: (-papers[pid]['year'], TYPE_ORDER.get(papers[pid]['exam_type'], 9), meta(papers[pid])[1], pid))
    for pid in live:
        g, school, kind = meta(papers[pid])
        base = '%d学年·%s·%s' % (papers[pid]['year'], school, kind)
        per_grade[g].append((pid, base))
    cs_of = {}
    for g, items in per_grade.items():
        seen = defaultdict(int)
        for i, (pid, base) in enumerate(items, 1):
            seen[base] += 1
            label[pid] = base if seen[base] == 1 else '%s#%d' % (base, seen[base])
            cs_of[pid] = i

    data, fig_refs, leak, n_scr = [], {}, 0, 0
    index_of = {}                      # 资产目录 → {资产id: 文件名}（两批各自的 assets/）
    for ad in set(asset_dir_of.values()):
        idx = {}
        with os.scandir(ad) as it:
            for e in it:
                idx[e.name.rsplit('.', 1)[0]] = e.name
        index_of[ad] = idx
    for q in kept:
        pid = q['paper_external_id']; p = papers[pid]
        asset_dir = asset_dir_of[pid]; index = index_of[asset_dir]
        g, school, kind = meta(p)
        stem = _inject_tables(q.get('stem_markdown') or '', q.get('tables'))
        stem = _strip_lead_no(stem, q.get('question_no'))
        # 2026-09-09 视觉转录卷：几何图只留占位（stem 里是裸 ⟦IMG⟧，没有资产），卡片上显示成「（图略）」，别把内部标记露给老师
        stem = stem.replace('⟦IMG⟧', '（图略）')
        stem, _n = _unicode_scripts(stem); n_scr += _n
        ims = []
        for aid in IMG_RE.findall(stem):
            fn = index.get(aid)
            if not fn:
                continue
            key = 'xb/' + aid
            fig_refs[key] = os.path.join(asset_dir, fn)
            ims.append(key)
        ans = str(q.get('answer_markdown') or '').strip()
        if len(ans) >= 4 and ans in stem:
            leak += 1
        data.append({
            'id': q['external_id'], 'tb': TB,
            'gg': g, 'lv': 'XB' + g,
            'cs': cs_of[pid], 'cd': label[pid],
            'src': kind, 's': kind,
            'n': q.get('question_no') or '', 'q': stem,
            't': _qtype(q.get('question_type_raw')), 'd': 0,
            'f': '', 'mf': [], 'p': [], 'k': [], 'im': ims,
            'pid': pid,
        })

    course_sub = OrderedDict()
    for g in GRADE_ORDER:
        nodes = OrderedDict()
        for x in data:
            if x['gg'] != g:
                continue
            nodes.setdefault(x['cs'], [x['cd'], 0, '%d学年' % papers[x['pid']]['year']])
            nodes[x['cs']][1] += 1
        if nodes:
            course_sub[g] = [[s, nd[0], nd[1], nd[2]] for s, nd in sorted(nodes.items())]

    cache = os.path.join(HERE, '_wmf_cache')
    broken = []
    for key, fp in sorted(fig_refs.items()):
        if not fp.lower().endswith('.wmf'):
            continue
        try:
            with Image.open(fp) as im:
                im.load()
        except Exception:
            broken.append((key[3:], fp))
    if broken:
        _rescue_wmf(broken, cache)
    FIG, skipped, rescued = {}, 0, 0
    for key, fp in sorted(fig_refs.items()):
        alt = os.path.join(cache, key[3:] + '.png')
        try:
            d64, W, H = _fig_b64(fp); FIG[key] = {'d': d64, 'w': W, 'h': H}
        except Exception:
            try:
                d64, W, H = _fig_b64(alt); FIG[key] = {'d': d64, 'w': W, 'h': H}; rescued += 1
            except Exception:
                skipped += 1
    n_hole = 0
    for x in data:
        miss = [k for k in x['im'] if k not in FIG]
        if not miss:
            continue
        for k in miss:
            x['q'] = re.sub(r'!\[[^\]]*\]\(asset://%s\)' % re.escape(k[3:]), '⟦图缺失⟧', x['q'])
        x['im'] = [k for k in x['im'] if k in FIG]
        n_hole += 1
    if verbose:
        print('  [校本] 批次 %d 个 · 卷 %d / 题 %d（同题面去重 -%d，小问不单独成卡 -%d，解析污染剔除 -%d）；年级 %s'
              % (len(BS), len(live), len(data), n_dup, n_sub, len(polluted), {g: len(v) for g, v in course_sub.items()}))
        print('  [校本] 配图 %d 张（WMF 救回 %d，失败 %d → %d 题标⟦图缺失⟧），FIG %.1f MB；答案泄漏自查 %d 条；裸上下标转 %d 处'
              % (len(FIG), rescued, skipped, n_hole, len(json.dumps(FIG)) / 1e6, leak, n_scr))
    return data, course_sub, FIG


if __name__ == '__main__':
    d, c, f = load(os.path.dirname(os.path.abspath(__file__)))
    print('自测：', len(d), '题', {g: len(v) for g, v in c.items()}, len(f), '图')
    if d:
        print('样例：', json.dumps({k: v for k, v in d[0].items() if k != 'q'}, ensure_ascii=False))
        print('节点：', c.get('八上', [])[:3])
