# -*- coding: utf-8 -*-
"""源 PPTX → 逐页转写所需数据（文字块含位置、媒体全部导出、备注）。

输出：
  _data/slides.json     每页 {index, title, texts[{text,left,top,w,h,size,bold,level}], pics[{file,w,h}], tables, notes}
  assets/media/*        内嵌媒体按内容去重导出（供页面素材区下载/复用）

坐标统一换算成"幻灯片百分比"，组合形状会做子坐标系到页面坐标系的变换。
"""
import hashlib
import json
import os
import sys

from pptx import Presentation
from pptx.enum.shapes import MSO_SHAPE_TYPE
from pptx.oxml.ns import qn

EMU = 914400.0


def parse_xfrm(shape):
    """返回 (left, top, w, h) 英寸；取不到返回 None。"""
    try:
        return (shape.left / EMU, shape.top / EMU, shape.width / EMU, shape.height / EMU)
    except Exception:
        return None


def group_map(group):
    """组合形状的子坐标系 → 页面坐标系的仿射参数 (ox, oy, sx, sy)。"""
    try:
        xfrm = group._element.find(qn('p:grpSpPr')).find(qn('a:xfrm'))
        off = xfrm.find(qn('a:off'))
        ext = xfrm.find(qn('a:ext'))
        ch_off = xfrm.find(qn('a:chOff'))
        ch_ext = xfrm.find(qn('a:chExt'))
        gx, gy = int(off.get('x')), int(off.get('y'))
        gw, gh = int(ext.get('cx')), int(ext.get('cy'))
        cx, cy = int(ch_off.get('x')), int(ch_off.get('y'))
        cw, ch = int(ch_ext.get('cx')), int(ch_ext.get('cy'))
        sx = gw / cw if cw else 1.0
        sy = gh / ch if ch else 1.0
        return gx / EMU, gy / EMU, cx / EMU, cy / EMU, sx, sy
    except Exception:
        return None


def iter_shapes(shapes, tf=None):
    """递归遍历（含组合），把子坐标映射到页面坐标。

    tf = (ox, oy, cox, coy, sx, sy)：页面 = ox + (child - cox) * sx
    """
    for s in shapes:
        if s.shape_type == MSO_SHAPE_TYPE.GROUP:
            sub = group_map(s)
            if sub is None:
                yield from iter_shapes(s.shapes, tf)
                continue
            ox, oy, cox, coy, gsx, gsy = sub
            if tf is None:
                ntf = (ox, oy, cox, coy, gsx, gsy)
            else:
                pox, poy, pcox, pcoy, psx, psy = tf
                ntf = (pox + (ox - pcox) * psx, poy + (oy - pcoy) * psy,
                       cox, coy, psx * gsx, psy * gsy)
            yield from iter_shapes(s.shapes, ntf)
        else:
            yield s, tf


def box(s, tf):
    """形状在页面坐标系中的位置（英寸）。"""
    try:
        l, t, w, h = s.left / EMU, s.top / EMU, s.width / EMU, s.height / EMU
    except Exception:
        return None
    if tf:
        ox, oy, cox, coy, sx, sy = tf
        l = ox + (l - cox) * sx
        t = oy + (t - coy) * sy
        w, h = w * sx, h * sy
    return l, t, w, h


def extract(src, outdir):
    prs = Presentation(src)
    SW, SH = prs.slide_width / EMU, prs.slide_height / EMU
    media_dir = os.path.join(outdir, 'assets', 'media')
    os.makedirs(media_dir, exist_ok=True)
    seen = {}
    slides = []
    for idx, slide in enumerate(prs.slides, 1):
        texts, pics, tables = [], [], []
        title = ''
        for s, tf in iter_shapes(slide.shapes):
            b = box(s, tf)
            if s.has_text_frame:
                t = s.text_frame.text.strip()
                if not t:
                    continue
                sizes, bolds = [], []
                for p in s.text_frame.paragraphs:
                    for r in p.runs:
                        if r.font.size:
                            sizes.append(r.font.size.pt)
                        if r.font.bold:
                            bolds.append(True)
                maxsz = max(sizes) if sizes else 18.0
                body_lines = [p.text.strip() for p in s.text_frame.paragraphs if p.text.strip()]
                texts.append({'text': t, 'lines': body_lines, 'left': b[0] if b else 0,
                              'top': b[1] if b else 0, 'w': b[2] if b else SW,
                              'h': b[3] if b else 0.5, 'size': maxsz,
                              'bold': bool(bolds), 'is_title': False})
            if s.shape_type == MSO_SHAPE_TYPE.PICTURE:
                try:
                    blob = s.image.blob
                except Exception:
                    blob = None
                if blob:
                    h = hashlib.sha1(blob).hexdigest()[:10]
                    ext = (s.image.ext or 'png').lower()
                    if h not in seen:
                        name = 'img-%s.%s' % (h, ext)
                        with open(os.path.join(media_dir, name), 'wb') as f:
                            f.write(blob)
                        seen[h] = name
                    pics.append({'file': seen[h], 'left': b[0] if b else 0, 'top': b[1] if b else 0,
                                 'w': b[2] if b else SW, 'h': b[3] if b else 1})
            if getattr(s, 'has_table', False) and s.has_table:
                rows = [[c.text.strip() for c in row.cells] for row in s.table.rows]
                tables.append({'rows': rows, 'left': b[0] if b else 0, 'top': b[1] if b else 0})
        # 标题：优先 slide.shapes.title，其次最上方/最大字号文本
        try:
            if slide.shapes.title is not None and slide.shapes.title.text.strip():
                title = slide.shapes.title.text.strip()
        except Exception:
            pass
        if not title and texts:
            cand = sorted(texts, key=lambda x: (-x['size'], x['top']))[0]
            title = cand['text'].split('\n')[0][:60]
        notes = ''
        if slide.has_notes_slide and slide.notes_slide.notes_text_frame is not None:
            notes = slide.notes_slide.notes_text_frame.text.strip()
        slides.append({'index': idx, 'title': title,
                       'texts': texts, 'pics': pics, 'tables': tables, 'notes': notes})
    data = {'source': os.path.basename(src), 'slide_w': SW, 'slide_h': SH,
            'count': len(slides), 'media': sorted(set(seen.values())), 'slides': slides}
    os.makedirs(os.path.join(outdir, '_data'), exist_ok=True)
    with open(os.path.join(outdir, '_data', 'slides.json'), 'w', encoding='utf-8') as f:
        json.dump(data, f, ensure_ascii=False)
    return data


if __name__ == '__main__':
    src = sys.argv[1]
    out = sys.argv[2] if len(sys.argv) > 2 else '.'
    d = extract(src, out)
    n_txt = sum(len(s['texts']) for s in d['slides'])
    n_pic = sum(len(s['pics']) for s in d['slides'])
    print('页数 %d ｜ 文字块 %d ｜ 图片位 %d ｜ 去重媒体 %d'
          % (d['count'], n_txt, n_pic, len(d['media'])))
