# -*- coding: utf-8 -*-
"""把出版社课件里"别的学校校徽"的**图片本体**替换成广东金融学院校徽。

与上一版（逐页贴图 + 白色底板覆盖）的区别：
  - 直接替换母版/版式里那张图片的二进制内容，因此校徽真正位于**母版体系内**，
    所有继承该版式的页面自动生效，页面上不再新增任何形状（没有白板、没有多余图层）；
  - 替换后按**等比 contain 适配**调整图片框：只在原框内缩放，宽高比严格保持，
    绝不出现单轴拉伸/压缩；位置保持原框中心不变。

识别方式：按图片内容的 sha1 前缀匹配（跨 16 份课件实测）
  - 9d5d9f75a2 → 同济大学校徽（jpg，151 处引用）
  - c255d8f3a6 → 武汉学院校徽（png，44 处）
  - c031d176c7 → 另一枚方形机构校徽（png，14 处，方形框用方形校徽图）
用法：python _replace_logo.py <输入目录或文件...> [--out 输出目录]
"""
import hashlib
import os
import sys

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

HERE = os.path.dirname(os.path.abspath(__file__))
BRAND = os.path.join(os.path.dirname(HERE), '课件成品-PPT', 'assets', 'brand')
ASSETS = {
    '9d5d9f75a2': os.path.join(BRAND, 'logo-master.jpg'),   # 同济大学（jpg 位）
    'c255d8f3a6': os.path.join(BRAND, 'logo-master-white.png'),  # 武汉学院位：位于深蓝装饰带上，用白色版
    'c031d176c7': os.path.join(BRAND, 'logo-emblem.png'),   # 方形校徽位
}
EMU = 914400.0


def sha(blob):
    return hashlib.sha1(blob).hexdigest()[:10]


def aspect(path):
    with Image.open(path) as im:
        return im.width / float(im.height)


def contain(box_w, box_h, cx, cy, ar):
    """等比 contain：在原框内按目标宽高比缩放，返回 (left, top, w, h)，中心不变。"""
    if box_w / box_h > ar:
        h = box_h
        w = box_h * ar
    else:
        w = box_w
        h = box_w / ar
    return cx - w / 2.0, cy - h / 2.0, w, h


def replace(prs, stats):
    """两遍处理：先按原始哈希登记要替换的图片位，再统一替换本体并逐框等比适配。

    注意：同一张图片在包内是**共享 part**（多个版式引用同一个 part），
    若边遍历边替换，后续版式读到的哈希已变成新图而被漏掉，因此必须先登记。
    """
    containers = [('master', m) for m in prs.slide_masters]
    containers += [('layout', l) for l in prs.slide_layouts]
    containers += [('slide', s) for s in prs.slides]

    def walk(shapes):
        """递归进组合形状：有些课件的校徽被放在组合里，只扫顶层会漏。"""
        for sh in shapes:
            if sh.shape_type == MSO_SHAPE_TYPE.GROUP:
                yield from walk(sh.shapes)
            else:
                yield sh

    todo_shapes = []          # (kind, shape, asset)
    part_asset = {}           # partname -> asset
    for kind, c in containers:
        for shape in walk(c.shapes):
            if shape.shape_type != MSO_SHAPE_TYPE.PICTURE:
                continue
            try:
                key = sha(shape.image.blob)
            except Exception:
                continue
            asset = ASSETS.get(key)
            if not asset:
                continue
            rId = shape._element.blipFill.blip.get(qn('r:embed'))
            try:
                part = shape.part.related_part(rId)
            except Exception:
                continue
            part_asset[part.partname] = asset
            todo_shapes.append((kind, shape, asset))

    # 1) 替换图片本体（每个 part 只写一次）
    for partname, asset in part_asset.items():
        part = None
        for _k, shape, a in todo_shapes:
            if a is asset:
                rId = shape._element.blipFill.blip.get(qn('r:embed'))
                cand = shape.part.related_part(rId)
                if cand.partname == partname:
                    part = cand
                    break
        if part is None:
            continue
        new_blob = open(asset, 'rb').read()
        part._blob = new_blob
        if hasattr(part, '_image'):
            try:
                from pptx.parts.image import Image as PptxImage
                part._image = PptxImage.from_blob(new_blob)
            except Exception:
                pass
        stats['parts'] += 1

    # 2) 每个引用位等比 contain 适配（绝不单轴拉伸）
    for kind, shape, asset in todo_shapes:
        ar = aspect(asset)
        bw, bh = shape.width / EMU, shape.height / EMU
        cx, cy = shape.left / EMU + bw / 2.0, shape.top / EMU + bh / 2.0
        left, top, w, h = contain(bw, bh, cx, cy, ar)
        shape.left, shape.top = int(left * EMU), int(top * EMU)
        shape.width, shape.height = int(w * EMU), int(h * EMU)
        stats['shapes'] += 1
        stats['where'][kind] = stats['where'].get(kind, 0) + 1
    return stats


def main():
    argv = sys.argv[1:]
    outdir = None
    if '--out' in argv:
        i = argv.index('--out')
        outdir = argv[i + 1]
        del argv[i:i + 2]
    files = []
    for a in argv:
        if os.path.isdir(a):
            files += [os.path.join(a, f) for f in sorted(os.listdir(a))
                      if f.lower().endswith('.pptx') and not f.startswith('~$')]
        else:
            files.append(a)
    if not files:
        print(__doc__)
        return 1
    if outdir:
        os.makedirs(outdir, exist_ok=True)
    for f in files:
        prs = Presentation(f)
        stats = {'parts': 0, 'shapes': 0, 'where': {}}
        replace(prs, stats)
        if stats['parts'] == 0:
            print('%-44s 未发现已知校徽图片，跳过' % os.path.basename(f)[:42])
            continue
        out = os.path.join(outdir, os.path.basename(f)) if outdir else f
        prs.save(out)
        print('%-44s 替换图片 %d 个 / 适配框 %d 处 ｜ %s'
              % (os.path.basename(f)[:42], stats['parts'], stats['shapes'],
                 ' '.join('%s%d' % (k, v) for k, v in sorted(stats['where'].items()))))
    return 0


if __name__ == '__main__':
    sys.exit(main())
