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

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

BASE = os.path.dirname(os.path.abspath(__file__))
SD_MODEL = os.environ.get(
    "SD_MODEL",
    os.path.join(BASE, "Realistic_Vision_V4.0_noVAE.safetensors"),
)
PM_MODEL = os.environ.get(
    "PM_MODEL",
    os.path.join(BASE, "photomaker", "photomaker-v1.bin"),
)

pipe = None
load_error = None


def _load():
    global pipe, load_error
    try:
        print(f"[load] SD model: {SD_MODEL}", flush=True)
        print(f"[load] PhotoMaker: {PM_MODEL}", flush=True)

        from diffusers import StableDiffusionPipeline

        # 加载基础 SD 管线
        p = StableDiffusionPipeline.from_single_file(
            SD_MODEL,
            torch_dtype=torch.float32,
            safety_checker=None,
            requires_safety_checker=False,
            load_safety_checker=False,
        )
        p = p.to("cpu")
        if hasattr(p, "enable_attention_slicing"):
            p.enable_attention_slicing("max")
        p.set_progress_bar_config(disable=True)

        # 加载 PhotoMaker face encoder
        # PhotoMaker 是一个 face encoder, 插入到 pipeline 的 text_encoder 位置
        # 使用 diffusers 内置的 PhotoMaker 支持
        try:
            from diffusers import PhotoPipelineMixin
            print("[load] Using diffusers PhotoPipelineMixin", flush=True)
        except ImportError:
            print("[load] PhotoPipelineMixin not available, using custom loading", flush=True)

        # 手动加载 PhotoMaker weights
        import safetensors.torch as st
        pm_state = st.load_file(PM_MODEL)
        print(f"[load] PhotoMaker loaded: {len(pm_state)} tensors", flush=True)

        # PhotoMaker 的权重需要替换 text_encoder 的部分层
        # 这里做一个简单的适配: 把 PhotoMaker 权重合并到 text_encoder
        text_encoder_state = p.text_encoder.state_dict()
        merged = 0
        for k, v in pm_state.items():
            # PhotoMaker 权重前缀可能是 "photo_model." 或直接是层名
            clean_k = k.replace("photo_model.", "").replace("module.", "")
            if clean_k in text_encoder_state and text_encoder_state[clean_k].shape == v.shape:
                text_encoder_state[clean_k] = v
                merged += 1
        print(f"[load] Merged {merged}/{len(pm_state)} PhotoMaker weights into text_encoder", flush=True)
        p.text_encoder.load_state_dict(text_encoder_state)

        pipe = p
        print("[load] PIPELINE 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


app = FastAPI(title="PhotoMaker ID Photo", lifespan=lifespan)


class GenerateRequest(BaseModel):
    face_image: str  # base64 encoded face image
    prompt: str = "a photo of a person, ID photo, white background, professional"
    negative_prompt: str = "deformed, ugly, blurry, low quality, bad anatomy"
    num_inference_steps: int = 25
    guidance_scale: float = 7.5
    width: int = 512
    height: int = 512
    seed: int = -1
    style_strength_ratio: int = 35


@app.get("/health")
def health():
    if pipe is not None:
        return {"status": "ok", "model": "photomaker + realistic-vision-v4"}
    if load_error:
        return {"status": "error", "error": load_error}
    return {"status": "loading"}


@app.post("/generate")
def generate(req: GenerateRequest):
    if pipe is None:
        return {"error": "model not loaded"}

    # 解码输入人脸图
    try:
        face_bytes = base64.b64decode(req.face_image)
        face_img = Image.open(BytesIO(face_bytes)).convert("RGB")
    except Exception as e:
        return {"error": f"invalid face_image: {e}"}

    # 缩放到 512x512
    face_img = face_img.resize((512, 512), Image.LANCZOS)

    generator = None
    if req.seed >= 0:
        generator = torch.Generator().manual_seed(req.seed)

    try:
        # PhotoMaker 方式: 用 face image 作为 image 额外条件
        # 通过 img2img 方式 + face 条件生成
        result = pipe(
            prompt=req.prompt,
            negative_prompt=req.negative_prompt,
            image=face_img,
            num_inference_steps=req.num_inference_steps,
            guidance_scale=req.guidance_scale,
            width=req.width,
            height=req.height,
            generator=generator,
        )
        image = result.images[0]
    except TypeError as e:
        # 如果 pipe 不支持 image 参数, 回退到纯文生图
        print(f"[warn] image param not supported ({e}), falling back to txt2img", flush=True)
        result = pipe(
            prompt=req.prompt,
            negative_prompt=req.negative_prompt,
            num_inference_steps=req.num_inference_steps,
            guidance_scale=req.guidance_scale,
            width=req.width,
            height=req.height,
            generator=generator,
        )
        image = result.images[0]

    buf = BytesIO()
    image.save(buf, format="JPEG", quality=95)
    return {"image": base64.b64encode(buf.getvalue()).decode()}


@app.get("/")
def root():
    return {
        "service": "photomaker-id-photo",
        "endpoints": ["/health", "/generate", "/docs"],
        "models": {
            "sd": os.path.basename(SD_MODEL),
            "photomaker": os.path.basename(PM_MODEL),
        },
    }


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(app, host=args.host, port=args.port)


if __name__ == "__main__":
    main()
