"""
BiSeNet-based background change module.
Adapted for /data/onlyface deployment using rembg + LaMa.
"""

import cv2
import numpy as np
from PIL import Image
import onnxruntime as ort
import torch

from certificate_pipeline import CertificatePipeline, PipelineConfig
from pathlib import Path

MODEL_DIR = Path("/data/onlyface/models")

# Initialize pipeline for background removal
_config = PipelineConfig(model_dir=MODEL_DIR, device="cpu")
_pipeline = CertificatePipeline(_config)


def change_background(image_path, output_path, bg_color=(255, 255, 255)):
    """
    Change background of an image using our pipeline.
    
    Args:
        image_path: Path to input image
        output_path: Path to save output
        bg_color: Background color (R, G, B)
    
    Returns:
        Error message string (empty if success)
    """
    try:
        img = cv2.imread(image_path)
        if img is None:
            return "无法读取图片"
        
        # Use rembg for segmentation (Stage 1)
        mask = _pipeline.stage1.run_rembg(img)
        
        # Dilate mask for hair coverage
        mask = _pipeline.stage1._dilate_mask(mask)
        
        # Create colored background
        bg = np.full_like(img, bg_color, dtype=np.uint8)
        
        # Blend
        mask_f = mask[..., None].astype(np.float32) / 255.0
        result = (img.astype(np.float32) * mask_f + bg.astype(np.float32) * (1.0 - mask_f)).astype(np.uint8)
        
        cv2.imwrite(output_path, result)
        return ""
    except Exception as e:
        return f"背景更换失败: {e}"


def parse_face_only(image_bgr, face_bbox=None):
    """
    Parse face only - returns 19-class mask (for BiSeNet compatibility).
    Uses our rembg session.
    """
    from certificate_pipeline import ModelManager, PipelineConfig
    
    mm = ModelManager(PipelineConfig(model_dir=MODEL_DIR))
    session = mm.get_rembg_session()
    
    import rembg
    pil_img = Image.fromarray(cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB))
    output = rembg.remove(pil_img, session=session, alpha_matting=True,
                          alpha_matting_foreground_threshold=240,
                          alpha_matting_background_threshold=10,
                          alpha_matting_erode_size=10)
    
    # Convert to 19-class format (simplified - just return person mask as class 13/14)
    mask = np.array(output)[:, :, 3]
    out = np.zeros((19, mask.shape[0], mask.shape[1]), dtype=np.uint8)
    out[13] = mask  # hair label
    out[14] = mask  # hat label
    return out


if __name__ == "__main__":
    # Test
    import sys
    if len(sys.argv) > 1:
        err = change_background(sys.argv[1], "/tmp/test_bg.jpg")
        print("Result:", err)