"""
BiSeNet 人脸部件分割 (CelebAMask-HQ, 19 类)
ONNX 推理：input 512x512 RGB，output (1,19,H,W) logits
返回：人脸 + 头发 mask（用于证件照换底）
"""
import os
import numpy as np
import cv2
import onnxruntime as ort

MODEL_PATH = os.path.join(os.path.dirname(__file__), "models", "bisenet", "bisenet_resnet34.onnx")
INPUT_SIZE = 512

# yakhyo/face-parsing 实测索引（与 CelebAMask-HQ 标准顺序不一致）
# 实测：14=hair(被模型当成头发，训练时合并了 hat/hair)
#      16=cloth(肩部衣服，被模型当成 necklace/neck 区域)
#      17=neck
PARTS = {
    0:  "background",
    1:  "skin",
    2:  "nose",
    3:  "eye_g",
    4:  "l_eye",
    5:  "r_eye",
    6:  "l_brow",
    7:  "r_brow",
    8:  "l_ear",
    9:  "r_ear",
    10: "mouth",
    11: "u_lip",
    12: "l_lip",
    13: "hair",        # 残留少量
    14: "hair_hat",    # 头发/帽（合并）
    15: "ear_r",
    16: "cloth",       # 衣服/肩部（被合并到这里）
    17: "neck",
    18: "unused",
}

# 证件照保留区域：皮肤 + 五官 + 嘴 + 头发 + 衣领/肩 + 脖子
KEEP_LABELS = {1, 2, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 16, 17}

_session = None


def _get_session():
    global _session
    if _session is None:
        if not os.path.exists(MODEL_PATH):
            raise FileNotFoundError(
                f"BiSeNet 模型不存在: {MODEL_PATH}\n"
                "下载: https://github.com/yakhyo/face-parsing/releases/download/weights/resnet34.onnx"
            )
        _session = ort.InferenceSession(MODEL_PATH, providers=["CPUExecutionProvider"])
    return _session


def parse_face(image_bgr, only_keep_labels=None, face_bbox=None):
    """输入 BGR 图像，返回 mask（0=背景，255=前景），与原图同尺寸。
    face_bbox: (x1,y1,x2,y2) 必填，用于向下/向外扩展肩膀区域（模型只识别到颈部）。"""
    if only_keep_labels is None:
        only_keep_labels = KEEP_LABELS

    h, w = image_bgr.shape[:2]
    img = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
    img = cv2.resize(img, (INPUT_SIZE, INPUT_SIZE), interpolation=cv2.INTER_LINEAR)

    mean = np.array([0.485, 0.456, 0.406], dtype=np.float32)
    std = np.array([0.229, 0.224, 0.225], dtype=np.float32)
    x = (img.astype(np.float32) / 255.0 - mean) / std
    x = x.transpose(2, 0, 1)[None, ...].astype(np.float32)

    sess = _get_session()
    outputs = sess.run(None, {"input": x})[0]
    label_map = outputs[0].argmax(axis=0).astype(np.uint8)

    keep_mask = np.zeros((INPUT_SIZE, INPUT_SIZE), dtype=np.uint8)
    for lbl in only_keep_labels:
        keep_mask[label_map == lbl] = 255

    keep_mask = cv2.resize(keep_mask, (w, h), interpolation=cv2.INTER_LINEAR)
    _, keep_mask = cv2.threshold(keep_mask, 127, 255, cv2.THRESH_BINARY)

    # —— 关键：用 face_bbox 向外扩展到肩膀 —— #
    if face_bbox is not None:
        x1, y1, x2, y2 = face_bbox
        fh = y2 - y1
        fw = x2 - x1
        cx = (x1 + x2) // 2

        # 上方留 0.3 头顶
        top = max(0, y1 - int(fh * 0.3))
        # 下方覆盖完整肩膀 (4.0× 脸高)，限制到图内
        bot = min(h, y2 + int(fh * 4.0))
        # 左右各扩 1.8 倍脸宽（保证肩部覆盖）
        side = int(fw * 1.8)
        left = max(0, cx - side)
        right = min(w, cx + side)

        # 椭圆中心 = bbox 中心
        ecx = (left + right) // 2
        ecy = (top + bot) // 2
        axes_x = (right - left) // 2
        axes_y = (bot - top) // 2

        shoulder_mask = np.zeros((h, w), dtype=np.uint8)
        cv2.ellipse(shoulder_mask,
                    (ecx, ecy),                  # center
                    (axes_x, axes_y),             # axes
                    0, 0, 360,                    # angle, start, end
                    255, -1)                      # filled

        # 与 segmentation mask 求并集
        keep_mask = cv2.bitwise_or(keep_mask, shoulder_mask)
        _, keep_mask = cv2.threshold(keep_mask, 127, 255, cv2.THRESH_BINARY)

    # 形态学：填洞 + 平滑边缘
    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5))
    keep_mask = cv2.morphologyEx(keep_mask, cv2.MORPH_CLOSE, kernel)
    keep_mask = cv2.GaussianBlur(keep_mask, (5, 5), 0)
    _, keep_mask = cv2.threshold(keep_mask, 127, 255, cv2.THRESH_BINARY)
    return keep_mask


def composite_on_background(image_bgr, mask, bg_color=(255, 255, 255)):
    """用 mask 把人物抠出，合成到纯色背景上。mask 0=背景，255=人物。"""
    mask_f = mask.astype(np.float32) / 255.0
    mask_f = mask_f[..., None]
    bg = np.array(bg_color, dtype=np.float32)
    out = image_bgr.astype(np.float32) * mask_f + bg * (1.0 - mask_f)
    return out.astype(np.uint8)


def change_background(image_bgr, bg_color=(255, 255, 255), face_bbox=None):
    """一步完成：分割 → 换底。返回 BGR 图像（与原图同尺寸）。"""
    mask = parse_face(image_bgr, face_bbox=face_bbox)
    return composite_on_background(image_bgr, mask, bg_color=bg_color)