#!/usr/bin/env python3
"""
Automated Certificate Photo & Text Replacement Pipeline

Stage 1: Portrait Removal & Background Restoration
Stage 2: Portrait Foreground Extraction & Compositing
Stage 3: Targeted Text Erasure & Reconstruction
"""

import os
import sys
import argparse
import logging
import hashlib
import json
from pathlib import Path
from typing import Dict, List, Tuple, Optional, Any, Union
from dataclasses import dataclass
from enum import Enum

import cv2
import numpy as np
from PIL import Image, ImageDraw, ImageFont, ImageFilter
import torch
import onnxruntime as ort
import easyocr

try:
    from insightface.app import FaceAnalysis
except ImportError:
    FaceAnalysis = None

try:
    import rembg
except ImportError:
    rembg = None


logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)


class ModelType(Enum):
    BIREFNET = "birefnet"
    RMBG = "rmbg"
    LAMA = "lama"
    INSWAPPER = "inswapper"
    INSGENDERAGE = "insgenderage"


@dataclass
class PipelineConfig:
    model_dir: Path = Path("/data/onlyface/models")
    checkpoint_dir: Path = Path("/data/onlyface/checkpoints")
    device: str = "cuda" if torch.cuda.is_available() else "cpu"
    birefnet_model: str = "BiRefNet-general-epoch_244.onnx"
    rmbg_model: str = "rmbg-2.0.onnx"
    lama_model: str = "big-lama.pt"
    inswapper_model: str = "inswapper_128.onnx"
    insgenderage_model: str = "genderage.onnx"
    dilation_kernel_size: int = 7
    text_mask_expansion: int = 2
    gaussian_blur_sigma: Tuple[float, float] = (0.5, 1.0)
    noise_sigma: float = 2.0
    font_size_range: Tuple[int, int] = (12, 24)
    landmark_model: str = "buffalo_l"
    face_det_size: Tuple[int, int] = (640, 640)


@dataclass
class TextField:
    key: str
    value: str
    bbox: Optional[Tuple[int, int, int, int]] = None
    font_info: Optional[Dict] = None


class ModelManager:
    def __init__(self, config: PipelineConfig):
        self.config = config
        self._models: Dict[str, Any] = {}
        self._session_options = ort.SessionOptions()
        self._session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
        if config.device == "cuda":
            self._providers = ["CUDAExecutionProvider", "CPUExecutionProvider"]
        else:
            self._providers = ["CPUExecutionProvider"]

    def get_birefnet(self) -> ort.InferenceSession:
        if "birefnet" not in self._models:
            model_path = self.config.model_dir / self.config.birefnet_model
            if not model_path.exists():
                self._download_birefnet(model_path)
            self._models["birefnet"] = ort.InferenceSession(
                str(model_path), sess_options=self._session_options, providers=self._providers
            )
            logger.info(f"Loaded BiRefNet from {model_path}")
        return self._models["birefnet"]

    def get_rmbg(self) -> ort.InferenceSession:
        if "rmbg" not in self._models:
            model_path = self.config.model_dir / self.config.rmbg_model
            if not model_path.exists():
                self._download_rmbg(model_path)
            self._models["rmbg"] = ort.InferenceSession(
                str(model_path), sess_options=self._session_options, providers=self._providers
            )
            logger.info(f"Loaded RMBG from {model_path}")
        return self._models["rmbg"]

    def get_lama(self) -> torch.nn.Module:
        if "lama" not in self._models:
            model_path = self.config.model_dir / self.config.lama_model
            if not model_path.exists():
                self._download_lama(model_path)
            self._models["lama"] = self._load_lama(model_path)
            logger.info(f"Loaded LaMa from {model_path}")
        return self._models["lama"]

    def get_inswapper(self) -> ort.InferenceSession:
        if "inswapper" not in self._models:
            model_path = self.config.model_dir / self.config.inswapper_model
            if not model_path.exists():
                self._download_inswapper(model_path)
            self._models["inswapper"] = ort.InferenceSession(
                str(model_path), sess_options=self._session_options, providers=self._providers
            )
            logger.info(f"Loaded InSwapper from {model_path}")
        return self._models["inswapper"]

    def get_insgenderage(self) -> ort.InferenceSession:
        if "insgenderage" not in self._models:
            model_path = self.config.model_dir / self.config.insgenderage_model
            if not model_path.exists():
                self._download_insgenderage(model_path)
            self._models["insgenderage"] = ort.InferenceSession(
                str(model_path), sess_options=self._session_options, providers=self._providers
            )
            logger.info(f"Loaded InGenderAge from {model_path}")
        return self._models["insgenderage"]

    def get_face_analyzer(self) -> "FaceAnalysis":
        if "face_analyzer" not in self._models:
            if FaceAnalysis is None:
                raise ImportError("insightface not installed. pip install insightface")
            app = FaceAnalysis(name=self.config.landmark_model, root=str(self.config.model_dir))
            app.prepare(ctx_id=0 if self.config.device == "cuda" else -1, det_size=self.config.face_det_size)
            self._models["face_analyzer"] = app
            logger.info(f"Loaded FaceAnalysis ({self.config.landmark_model})")
        return self._models["face_analyzer"]

    def get_easyocr(self):
        if "easyocr" not in self._models:
            self._models["easyocr"] = easyocr.Reader(['ch_sim', 'en'], gpu=(self.config.device == "cuda"))
            logger.info("Loaded EasyOCR")
        return self._models["easyocr"]

    def get_rembg_session(self):
        if rembg is None:
            raise ImportError("rembg not installed. pip install rembg")
        if "rembg" not in self._models:
            # Use isnet-general-use for CPU (fast, 179MB), birefnet-general for GPU
            if self.config.device == "cuda":
                self._models["rembg"] = rembg.new_session("birefnet-general")
            else:
                self._models["rembg"] = rembg.new_session("isnet-general-use")
        return self._models["rembg"]

    def _download_birefnet(self, path: Path):
        logger.info(f"Downloading BiRefNet to {path}")
        import urllib.request
        url = "https://huggingface.co/zhengpeng7/BiRefNet/resolve/main/BiRefNet-general-epoch_244.onnx"
        urllib.request.urlretrieve(url, path)

    def _download_rmbg(self, path: Path):
        logger.info(f"Downloading RMBG-2.0 to {path}")
        import urllib.request
        url = "https://huggingface.co/briaai/RMBG-2.0/resolve/main/rmbg-2.0.onnx"
        urllib.request.urlretrieve(url, path)

    def _download_lama(self, path: Path):
        logger.info(f"Downloading LaMa to {path}")
        import urllib.request
        url = "https://github.com/Sanster/models/releases/download/add_big_lama/big-lama.pt"
        urllib.request.urlretrieve(url, path)

    def _download_inswapper(self, path: Path):
        logger.info(f"Downloading InSwapper to {path}")
        import urllib.request
        url = "https://huggingface.co/deepinsight/inswapper/resolve/main/inswapper_128.onnx"
        urllib.request.urlretrieve(url, path)

    def _download_insgenderage(self, path: Path):
        logger.info(f"Downloading InGenderAge to {path}")
        import urllib.request
        url = "https://huggingface.co/deepinsight/genderage/resolve/main/genderage.onnx"
        urllib.request.urlretrieve(url, path)

    def _load_lama(self, path: Path) -> torch.nn.Module:
        model = torch.jit.load(path, map_location=self.config.device)
        model.eval()
        return model


class Stage1_PortraitRemoval:
    def __init__(self, model_manager: ModelManager, config: PipelineConfig):
        self.mm = model_manager
        self.config = config

    def _preprocess(self, image: np.ndarray, target_size: Tuple[int, int] = (1024, 1024)) -> Tuple[np.ndarray, Tuple[float, float]]:
        h, w = image.shape[:2]
        scale = min(target_size[0] / w, target_size[1] / h)
        new_w, new_h = int(w * scale), int(h * scale)
        resized = cv2.resize(image, (new_w, new_h), interpolation=cv2.INTER_LINEAR)
        
        canvas = np.zeros((target_size[1], target_size[0], 3), dtype=np.uint8)
        canvas[:new_h, :new_w] = resized
        return canvas, (scale, scale)

    def _postprocess_mask(self, mask: np.ndarray, original_shape: Tuple[int, int], scale: Tuple[float, float]) -> np.ndarray:
        h, w = original_shape[:2]
        mask = mask[:int(h * scale[1]), :int(w * scale[0])]
        mask = cv2.resize(mask, (w, h), interpolation=cv2.INTER_LINEAR)
        return mask

    def _dilate_mask(self, mask: np.ndarray) -> np.ndarray:
        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (self.config.dilation_kernel_size, self.config.dilation_kernel_size))
        return cv2.dilate(mask, kernel, iterations=1)

    def run_birefnet(self, image: np.ndarray) -> np.ndarray:
        session = self.mm.get_birefnet()
        input_tensor, scale = self._preprocess(image)
        input_tensor = input_tensor.astype(np.float32) / 255.0
        input_tensor = np.transpose(input_tensor, (2, 0, 1))[np.newaxis, ...]
        
        ort_inputs = {session.get_inputs()[0].name: input_tensor}
        ort_outs = session.run(None, ort_inputs)
        mask = ort_outs[0][0, 0]
        mask = 1 / (1 + np.exp(-mask))
        mask = (mask > 0.5).astype(np.uint8) * 255
        
        return self._postprocess_mask(mask, image.shape, scale)

    def run_rmbg(self, image: np.ndarray) -> np.ndarray:
        session = self.mm.get_rmbg()
        input_tensor, scale = self._preprocess(image)
        input_tensor = input_tensor.astype(np.float32) / 255.0
        mean = np.array([0.485, 0.456, 0.406], dtype=np.float32).reshape(1, 1, 3)
        std = np.array([0.229, 0.224, 0.225], dtype=np.float32).reshape(1, 1, 3)
        input_tensor = (input_tensor - mean) / std
        input_tensor = np.transpose(input_tensor, (2, 0, 1))[np.newaxis, ...]
        
        ort_inputs = {session.get_inputs()[0].name: input_tensor}
        ort_outs = session.run(None, ort_inputs)
        mask = ort_outs[0][0, 0]
        mask = 1 / (1 + np.exp(-mask))
        mask = (mask > 0.5).astype(np.uint8) * 255
        
        return self._postprocess_mask(mask, image.shape, scale)

    def run_rembg(self, image: np.ndarray) -> np.ndarray:
        session = self.mm.get_rembg_session()
        pil_img = Image.fromarray(cv2.cvtColor(image, 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)
        mask = np.array(output)[:, :, 3]
        return mask

    def run_lama_inpaint(self, image: np.ndarray, mask: np.ndarray) -> np.ndarray:
        model = self.mm.get_lama()
        h, w = image.shape[:2]
        
        img_tensor = torch.from_numpy(image.astype(np.float32) / 255.0).permute(2, 0, 1).unsqueeze(0).to(self.config.device)
        mask_tensor = torch.from_numpy(mask.astype(np.float32) / 255.0).unsqueeze(0).unsqueeze(0).to(self.config.device)
        
        with torch.no_grad():
            result = model(img_tensor, mask_tensor)
        
        result = result.squeeze(0).permute(1, 2, 0).cpu().numpy()
        result = (np.clip(result, 0, 1) * 255).astype(np.uint8)
        return result

    def execute(self, image: np.ndarray, use_birefnet: bool = True) -> Tuple[np.ndarray, np.ndarray]:
        logger.info("Stage 1: Portrait Removal & Background Restoration")
        
        # Use rembg for segmentation (handles both GPU and CPU efficiently)
        mask = self.run_rembg(image)
        
        mask = self._dilate_mask(mask)
        
        logger.info("Running LaMa inpainting...")
        bg_clean = self.run_lama_inpaint(image, mask)
        
        logger.info("Stage 1 completed")
        return bg_clean, mask


class Stage2_PortraitCompositing:
    def __init__(self, model_manager: ModelManager, config: PipelineConfig):
        self.mm = model_manager
        self.config = config

    def _get_landmarks(self, image: np.ndarray) -> Optional[np.ndarray]:
        app = self.mm.get_face_analyzer()
        faces = app.get(image)
        if not faces:
            return None
        face = max(faces, key=lambda f: f.det_score)
        return face.landmark_2d_106 if hasattr(face, 'landmark_2d_106') else face.landmark_2d_68

    def _compute_affine(self, src_pts: np.ndarray, dst_pts: np.ndarray) -> np.ndarray:
        src_center = np.mean(src_pts, axis=0)
        dst_center = np.mean(dst_pts, axis=0)
        
        src_norm = src_pts - src_center
        dst_norm = dst_pts - dst_center
        
        scale = np.linalg.norm(dst_norm) / (np.linalg.norm(src_norm) + 1e-8)
        
        cos_theta = np.sum(src_norm * dst_norm) / (np.linalg.norm(src_norm) * np.linalg.norm(dst_norm) + 1e-8)
        cos_theta = np.clip(cos_theta, -1.0, 1.0)
        theta = np.arccos(cos_theta)
        
        cross = np.cross(src_norm[0], dst_norm[0])
        if cross < 0:
            theta = -theta
        
        M = cv2.getRotationMatrix2D(tuple(dst_center), np.degrees(theta), scale)
        M[:, 2] += dst_center - src_center * scale
        
        return M

    def _extract_foreground(self, image: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
        session = self.mm.get_rembg_session()
        pil_img = Image.fromarray(cv2.cvtColor(image, 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)
        output_np = np.array(output)
        mask = output_np[:, :, 3]
        foreground = cv2.cvtColor(output_np[:, :, :3], cv2.COLOR_RGBA2BGR)
        foreground = cv2.bitwise_and(foreground, foreground, mask=mask)
        return foreground, mask

    def _multiband_blend(self, fg: np.ndarray, bg: np.ndarray, mask: np.ndarray, levels: int = 5) -> np.ndarray:
        fg = fg.astype(np.float32) / 255.0
        bg = bg.astype(np.float32) / 255.0
        mask = mask.astype(np.float32) / 255.0
        
        gaussian_pyr_fg = [fg]
        gaussian_pyr_bg = [bg]
        gaussian_pyr_mask = [mask]
        
        for i in range(levels):
            gaussian_pyr_fg.append(cv2.pyrDown(gaussian_pyr_fg[-1]))
            gaussian_pyr_bg.append(cv2.pyrDown(gaussian_pyr_bg[-1]))
            gaussian_pyr_mask.append(cv2.pyrDown(gaussian_pyr_mask[-1]))
        
        laplacian_pyr_fg = []
        laplacian_pyr_bg = []
        
        for i in range(levels):
            up_fg = cv2.pyrUp(gaussian_pyr_fg[i+1], dstsize=(gaussian_pyr_fg[i].shape[1], gaussian_pyr_fg[i].shape[0]))
            up_bg = cv2.pyrUp(gaussian_pyr_bg[i+1], dstsize=(gaussian_pyr_bg[i].shape[1], gaussian_pyr_bg[i].shape[0]))
            laplacian_pyr_fg.append(gaussian_pyr_fg[i] - up_fg)
            laplacian_pyr_bg.append(gaussian_pyr_bg[i] - up_bg)
        
        laplacian_pyr_fg.append(gaussian_pyr_fg[-1])
        laplacian_pyr_bg.append(gaussian_pyr_bg[-1])
        
        blended_pyr = []
        for i in range(levels + 1):
            m = gaussian_pyr_mask[i] if i < len(gaussian_pyr_mask) else gaussian_pyr_mask[-1]
            if len(m.shape) == 2:
                m = m[:, :, np.newaxis]
            blended = laplacian_pyr_fg[i] * m + laplacian_pyr_bg[i] * (1 - m)
            blended_pyr.append(blended)
        
        result = blended_pyr[-1]
        for i in range(levels - 1, -1, -1):
            result = cv2.pyrUp(result, dstsize=(blended_pyr[i].shape[1], blended_pyr[i].shape[0]))
            result += blended_pyr[i]
        
        result = np.clip(result, 0, 1)
        return (result * 255).astype(np.uint8)

    def execute(self, target_portrait: np.ndarray, bg_clean: np.ndarray, orig_mask: np.ndarray) -> np.ndarray:
        logger.info("Stage 2: Portrait Foreground Extraction & Compositing")
        
        orig_landmarks = self._get_landmarks(cv2.bitwise_and(bg_clean, bg_clean, mask=cv2.bitwise_not(orig_mask)))
        target_landmarks = self._get_landmarks(target_portrait)
        
        if orig_landmarks is None or target_landmarks is None:
            logger.warning("Face landmarks not detected, using center alignment")
            h, w = bg_clean.shape[:2]
            th, tw = target_portrait.shape[:2]
            M = np.eye(2, 3, dtype=np.float32)
            M[0, 2] = (w - tw) / 2
            M[1, 2] = (h - th) / 2
        else:
            key_indices = [30, 36, 45, 48, 54] if len(orig_landmarks) >= 68 else list(range(min(5, len(orig_landmarks))))
            src_pts = target_landmarks[key_indices].astype(np.float32)
            dst_pts = orig_landmarks[key_indices].astype(np.float32)
            M = self._compute_affine(src_pts, dst_pts)
        
        fg, alpha = self._extract_foreground(target_portrait)
        
        warped_fg = cv2.warpAffine(fg, M, (bg_clean.shape[1], bg_clean.shape[0]), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT)
        warped_alpha = cv2.warpAffine(alpha, M, (bg_clean.shape[1], bg_clean.shape[0]), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT)
        
        warped_alpha = cv2.GaussianBlur(warped_alpha, (5, 5), 0)
        
        result = self._multiband_blend(warped_fg, bg_clean, warped_alpha)
        
        logger.info("Stage 2 completed")
        return result


class Stage3_TextReplacement:
    def __init__(self, model_manager: ModelManager, config: PipelineConfig):
        self.mm = model_manager
        self.config = config

    def _detect_text_fields(self, image: np.ndarray, target_fields: Dict[str, str]) -> List[TextField]:
        reader = self.mm.get_easyocr()
        results = reader.readtext(image)
        
        fields = []
        for bbox, text, conf in results:
            for key, expected_value in target_fields.items():
                if self._text_matches(text, expected_value, key):
                    x_coords = [p[0] for p in bbox]
                    y_coords = [p[1] for p in bbox]
                    x1, x2 = int(min(x_coords)), int(max(x_coords))
                    y1, y2 = int(min(y_coords)), int(max(y_coords))
                    
                    font_info = self._analyze_font(image, (x1, y1, x2, y2), text)
                    
                    fields.append(TextField(
                        key=key,
                        value=target_fields[key],
                        bbox=(x1, y1, x2, y2),
                        font_info=font_info
                    ))
                    break
        
        return fields

    def _text_matches(self, detected: str, expected: str, key: str) -> bool:
        detected_clean = detected.strip().upper()
        expected_clean = expected.strip().upper()
        
        if key.upper() in detected_clean or detected_clean in expected_clean:
            return True
        
        if key.lower() == "sex" and detected_clean in ["M", "F", "男", "女"]:
            return True
        if key.lower() in ["dob", "date of birth", "出生日期"] and any(c.isdigit() for c in detected):
            return True
        if key.lower() in ["doi", "date of issue", "签发日期"] and any(c.isdigit() for c in detected):
            return True
        
        return False

    def _analyze_font(self, image: np.ndarray, bbox: Tuple[int, int, int, int], text: str) -> Dict:
        x1, y1, x2, y2 = bbox
        roi = image[y1:y2, x1:x2]
        
        if roi.size == 0:
            return {"size": 16, "color": (0, 0, 0), "font": "DejaVuSans.ttf"}
        
        hsv = cv2.cvtColor(roi, cv2.COLOR_BGR2HSV)
        mask = cv2.inRange(hsv, np.array([0, 0, 0]), np.array([180, 255, 120]))
        text_pixels = roi[mask > 0]
        
        if len(text_pixels) > 0:
            color = tuple(map(int, np.median(text_pixels, axis=0)))
        else:
            color = (0, 0, 0)
        
        font_size = max(10, min(y2 - y1, 32))
        
        return {"size": font_size, "color": color, "font": "DejaVuSans.ttf"}

    def _create_text_mask(self, image_shape: Tuple[int, int], fields: List[TextField]) -> np.ndarray:
        mask = np.zeros(image_shape[:2], dtype=np.uint8)
        for field in fields:
            if field.bbox:
                x1, y1, x2, y2 = field.bbox
                exp = self.config.text_mask_expansion
                x1 = max(0, x1 - exp)
                y1 = max(0, y1 - exp)
                x2 = min(image_shape[1], x2 + exp)
                y2 = min(image_shape[0], y2 + exp)
                cv2.rectangle(mask, (x1, y1), (x2, y2), 255, -1)
        return mask

    def _inpaint_text(self, image: np.ndarray, mask: np.ndarray) -> np.ndarray:
        model = self.mm.get_lama()
        h, w = image.shape[:2]
        
        img_tensor = torch.from_numpy(image.astype(np.float32) / 255.0).permute(2, 0, 1).unsqueeze(0).to(self.config.device)
        mask_tensor = torch.from_numpy(mask.astype(np.float32) / 255.0).unsqueeze(0).unsqueeze(0).to(self.config.device)
        
        with torch.no_grad():
            result = model(img_tensor, mask_tensor)
        
        result = result.squeeze(0).permute(1, 2, 0).cpu().numpy()
        result = (np.clip(result, 0, 1) * 255).astype(np.uint8)
        return result

    def _render_text_overlay(self, image: np.ndarray, fields: List[TextField]) -> np.ndarray:
        pil_img = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB)).convert("RGBA")
        overlay = Image.new("RGBA", pil_img.size, (0, 0, 0, 0))
        draw = ImageDraw.Draw(overlay)
        
        for field in fields:
            if not field.bbox or not field.font_info:
                continue
            
            x1, y1, x2, y2 = field.bbox
            font_size = field.font_info.get("size", 16)
            color = field.font_info.get("color", (0, 0, 0))
            font_path = field.font_info.get("font", "DejaVuSans.ttf")
            
            try:
                font = ImageFont.truetype(font_path, font_size)
            except:
                font = ImageFont.load_default()
            
            # Use textbbox instead of deprecated textsize
            bbox = draw.textbbox((0, 0), field.value, font=font)
            text_w = bbox[2] - bbox[0]
            text_h = bbox[3] - bbox[1]
            tx = x1 + (x2 - x1 - text_w) // 2
            ty = y1 + (y2 - y1 - text_h) // 2
            
            draw.text((tx, ty), field.value, font=font, fill=(*color, 255))
        
        overlay_np = np.array(overlay)
        
        # Apply Gaussian blur only to RGB channels, preserve alpha
        sigma = np.random.uniform(*self.config.gaussian_blur_sigma)
        rgb = cv2.GaussianBlur(overlay_np[:, :, :3], (0, 0), sigma)
        alpha = overlay_np[:, :, 3]
        overlay_np = np.dstack([rgb, alpha])
        
        # Add noise only to RGB channels
        noise = np.random.normal(0, self.config.noise_sigma, rgb.shape).astype(np.float32)
        rgb = np.clip(rgb.astype(np.float32) + noise, 0, 255).astype(np.uint8)
        overlay_np = np.dstack([rgb, alpha])
        
        return overlay_np

    def _composite_text(self, base: np.ndarray, overlay: np.ndarray) -> np.ndarray:
        # Ensure base is RGBA
        if base.shape[2] == 3:
            base_rgba = cv2.cvtColor(base, cv2.COLOR_BGR2RGBA)
        else:
            base_rgba = base
        
        # Ensure overlay is RGBA
        if overlay.shape[2] == 3:
            overlay_rgba = cv2.cvtColor(overlay, cv2.COLOR_BGR2RGBA)
        else:
            overlay_rgba = overlay
        
        # Composite RGB channels with alpha blending
        alpha = overlay_rgba[:, :, 3:4] / 255.0
        base_rgb = base_rgba[:, :, :3]
        overlay_rgb = overlay_rgba[:, :, :3]
        
        # Alpha blend: result = base * (1 - alpha) + overlay * alpha
        result_rgb = base_rgb * (1 - alpha) + overlay_rgb * alpha
        result_rgb = np.clip(result_rgb, 0, 255).astype(np.uint8)
        
        # Keep original alpha from base (or use overlay alpha)
        result_alpha = base_rgba[:, :, 3:4]
        
        result_rgba = np.concatenate([result_rgb, result_alpha], axis=2)
        return cv2.cvtColor(result_rgba, cv2.COLOR_RGBA2BGR)

    def execute(self, image: np.ndarray, text_fields: Dict[str, str]) -> np.ndarray:
        logger.info("Stage 3: Targeted Text Erasure & Reconstruction")
        
        detected_fields = self._detect_text_fields(image, text_fields)
        
        if not detected_fields:
            logger.warning("No target text fields detected")
            return image
        
        logger.info(f"Detected {len(detected_fields)} text fields: {[f.key for f in detected_fields]}")
        
        text_mask = self._create_text_mask(image.shape, detected_fields)
        
        logger.info("Inpainting text regions...")
        image_no_text = self._inpaint_text(image, text_mask)
        
        logger.info("Rendering new text overlay...")
        overlay = self._render_text_overlay(image_no_text, detected_fields)
        
        logger.info("Compositing final result...")
        result = self._composite_text(image_no_text, overlay)
        
        logger.info("Stage 3 completed")
        return result


class CertificatePipeline:
    def __init__(self, config: Optional[PipelineConfig] = None):
        self.config = config or PipelineConfig()
        self.mm = ModelManager(self.config)
        self.stage1 = Stage1_PortraitRemoval(self.mm, self.config)
        self.stage2 = Stage2_PortraitCompositing(self.mm, self.config)
        self.stage3 = Stage3_TextReplacement(self.mm, self.config)

    def process(self, 
                original_image: Union[str, Path, np.ndarray],
                target_portrait: Union[str, Path, np.ndarray],
                text_fields: Dict[str, str],
                output_path: Union[str, Path],
                use_birefnet: bool = True) -> np.ndarray:
        
        if isinstance(original_image, (str, Path)):
            orig_img = cv2.imread(str(original_image))
            if orig_img is None:
                raise ValueError(f"Cannot read original image: {original_image}")
        else:
            orig_img = original_image.copy()
        
        if isinstance(target_portrait, (str, Path)):
            target_img = cv2.imread(str(target_portrait))
            if target_img is None:
                raise ValueError(f"Cannot read target portrait: {target_portrait}")
        else:
            target_img = target_portrait.copy()
        
        logger.info(f"Processing: {original_image} -> {output_path}")
        
        bg_clean, portrait_mask = self.stage1.execute(orig_img, use_birefnet)
        
        portrait_done = self.stage2.execute(target_img, bg_clean, portrait_mask)
        
        final_result = self.stage3.execute(portrait_done, text_fields)
        
        cv2.imwrite(str(output_path), final_result)
        logger.info(f"Saved result to {output_path}")
        
        return final_result


def main():
    parser = argparse.ArgumentParser(description="Certificate Photo & Text Replacement Pipeline")
    parser.add_argument("--original", "-o", required=True, help="Path to original certificate image")
    parser.add_argument("--portrait", "-p", required=True, help="Path to new portrait image")
    parser.add_argument("--output", "-O", required=True, help="Path to output image")
    parser.add_argument("--fields", "-f", required=True, help="JSON string of text fields to replace, e.g. '{\"Sex\": \"M\", \"Date of Issue\": \"12 DEC 2026\"}'")
    parser.add_argument("--model-dir", default="/data/onlyface/models", help="Model directory")
    parser.add_argument("--use-rmbg", action="store_true", help="Use RMBG instead of BiRefNet for segmentation")
    parser.add_argument("--device", choices=["cuda", "cpu"], default="cuda" if torch.cuda.is_available() else "cpu")
    
    args = parser.parse_args()
    
    text_fields = json.loads(args.fields)
    
    config = PipelineConfig(
        model_dir=Path(args.model_dir),
        device=args.device
    )
    
    pipeline = CertificatePipeline(config)
    
    pipeline.process(
        original_image=args.original,
        target_portrait=args.portrait,
        text_fields=text_fields,
        output_path=args.output,
        use_birefnet=not args.use_rmbg
    )
    
    logger.info("Pipeline completed successfully!")


if __name__ == "__main__":
    main()