#!/usr/bin/env python3
"""
InsightFace 证件照换头 API 服务
启动: python3 headswap_server.py [--port 7870]
"""
import os
import sys
import base64
import argparse
import time
import threading
from io import BytesIO
from contextlib import asynccontextmanager

import cv2
import numpy as np
from PIL import Image
from fastapi import FastAPI
from pydantic import BaseModel

BASE = os.path.dirname(os.path.abspath(__file__))
INSIGHTFACE_ROOT = os.path.join(BASE, "models", "insightface")

app_analyzer = None
load_error = None


def _load():
    global app_analyzer, load_error
    try:
        print("[load] initializing InsightFace...", flush=True)
        from insightface.app import FaceAnalysis
        app_analyzer = FaceAnalysis(
            name="buffalo_l",
            root=INSIGHTFACE_ROOT,
            providers=["CPUExecutionProvider"],
            allowed_modules=["detection", "landmark_2d_106"],
        )
        app_analyzer.prepare(ctx_id=0, det_size=(640, 640))
        print("[load] READY", flush=True)
    except Exception as e:
        import traceback
        traceback.print_exc()
        load_error = f"{type(e).__name__}: {e}"
        print(f"[load] FAILED: {load_error}", flush=True)


@asynccontextmanager
async def lifespan(_app):
    t = threading.Thread(target=_load, daemon=True)
    t.start()
    yield


api = FastAPI(title="InsightFace Head Swap", lifespan=lifespan)


class HeadSwapRequest(BaseModel):
    src_image: str  # base64 source (提供头的人)
    dst_image: str  # base64 target (要换头的照片)
    expand: float = 2.0
    quality: int = 95


def _decode_b64(b64_str: str) -> np.ndarray:
    data = base64.b64decode(b64_str)
    img = Image.open(BytesIO(data)).convert("RGB")
    return cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR)


def _get_affine(src_kps, dst_kps):
    src = src_kps.astype(np.float32)
    dst = dst_kps.astype(np.float32)
    m, _ = cv2.estimateAffinePartial2D(src, dst, method=cv2.LMEDS)
    return m


def _make_head_mask(kps, shape, scale=2.0):
    h, w = shape[:2]
    cx = kps[:, 0].mean()
    cy = kps[:, 1].mean()
    fw = kps[:, 0].max() - kps[:, 0].min()
    fh = kps[:, 1].max() - kps[:, 1].min()
    rx = fw * scale / 2
    ry = fh * scale * 1.1 / 2
    cy_shift = ry * 0.08
    mask = np.zeros((h, w), dtype=np.uint8)
    center = (int(cx), int(cy - cy_shift))
    axes = (int(rx), int(ry))
    cv2.ellipse(mask, center, axes, 0, 0, 360, 255, -1)
    blur_k = max(31, int(min(rx, ry) * 0.4) | 1)
    mask = cv2.GaussianBlur(mask, (blur_k, blur_k), blur_k // 3)
    return mask


def _swap(src: np.ndarray, dst: np.ndarray, expand: float) -> np.ndarray:
    src_faces = app_analyzer.get(src)
    dst_faces = app_analyzer.get(dst)
    if not src_faces:
        raise ValueError("source: no face detected")
    if not dst_faces:
        raise ValueError("target: no face detected")

    src_kps = src_faces[0].kps
    dst_kps = dst_faces[0].kps

    M = _get_affine(src_kps, dst_kps)
    if M is None:
        raise ValueError("affine transform failed")

    src_head_mask = _make_head_mask(src_kps, src.shape, scale=expand)
    h, w = dst.shape[:2]
    src_warped = cv2.warpAffine(src, M, (w, h),
                                 flags=cv2.INTER_LINEAR,
                                 borderMode=cv2.BORDER_CONSTANT,
                                 borderValue=(0, 0, 0))
    mask_warped = cv2.warpAffine(src_head_mask, M, (w, h),
                                  flags=cv2.INTER_LINEAR,
                                  borderMode=cv2.BORDER_CONSTANT,
                                  borderValue=0)

    erase_mask = (mask_warped > 20).astype(np.uint8) * 255
    kernel = np.ones((7, 7), np.uint8)
    erase_mask = cv2.dilate(erase_mask, kernel, iterations=2)
    erase_binary = (erase_mask > 128).astype(np.uint8) * 255
    dst_erased = cv2.inpaint(dst, erase_binary, inpaintRadius=10,
                              flags=cv2.INPAINT_TELEA)

    alpha = mask_warped.astype(np.float32) / 255.0
    alpha = alpha[:, :, np.newaxis]

    outer_ring = cv2.dilate(erase_mask, np.ones((20, 20), np.uint8)) - erase_mask
    color_shift = np.zeros(3, dtype=np.float32)
    if outer_ring.sum() > 0:
        outer_px = dst[outer_ring > 0].astype(np.float32).mean(axis=0)
        inner_mask = (mask_warped > 128)
        if inner_mask.sum() > 0:
            inner_px = src_warped[inner_mask].astype(np.float32).mean(axis=0)
            color_shift = (outer_px - inner_px) * 0.35

    src_adjusted = np.clip(src_warped.astype(np.float32) + color_shift, 0, 255)
    result = dst_erased.astype(np.float32) * (1 - alpha) + src_adjusted * alpha
    return np.clip(result, 0, 255).astype(np.uint8)


@api.get("/health")
def health():
    if app_analyzer is not None:
        return {"status": "ok", "model": "buffalo_l", "engine": "insightface"}
    if load_error:
        return {"status": "error", "error": load_error}
    return {"status": "loading"}


@api.post("/swap")
def swap(req: HeadSwapRequest):
    if app_analyzer is None:
        return {"error": "model not loaded"}
    t0 = time.time()
    try:
        src = _decode_b64(req.src_image)
        dst = _decode_b64(req.dst_image)
        result = _swap(src, dst, req.expand)
    except Exception as e:
        return {"error": str(e)}

    ok, enc = cv2.imencode(".jpg", result, [cv2.IMWRITE_JPEG_QUALITY, req.quality])
    if not ok:
        return {"error": "encode failed"}
    b64 = base64.b64encode(enc.tobytes()).decode()
    elapsed = time.time() - t0
    return {"image": b64, "time_ms": int(elapsed * 1000)}


@api.post("/generate")
def generate_alias(req: HeadSwapRequest):
    """兼容 OpenAI images API 格式"""
    return swap(req)


@api.get("/")
def root():
    return {
        "service": "insightface-headswap",
        "endpoints": ["/health", "/swap", "/docs"],
        "usage": "POST /swap with {src_image, dst_image} as base64",
    }


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--port", type=int, default=7870)
    parser.add_argument("--host", default="127.0.0.1")
    args = parser.parse_args()
    import uvicorn
    uvicorn.run(api, host=args.host, port=args.port)


if __name__ == "__main__":
    main()
