#!/usr/bin/env python3
"""
MiniCPM-V 2.6 int4 量化版 → OpenAI 兼容 HTTP API（FastAPI）。
部署在 NVIDIA A30 24GB 上，显存占用 ~6GB，单张推理 2-5 秒。

兼容接口:
    POST /v1/chat/completions   (OpenAI 格式, 支持图片 base64 / URL)
    GET  /v1/models
    GET  /health

启动:
    # 先下载模型: python3 -c "from modelscope import snapshot_download; snapshot_download('openbmb/MiniCPM-V-2_6-int4', cache_dir='/data/volc_minicpm/models/.ms_cache', local_dir='/data/volc_minicpm/models/MiniCPM-V-2_6-int4')"
    # 启动服务:
    uvicorn minicpmv_server:app --host 0.0.0.0 --port 30000 --log-level info
"""
import io
import os
import re
import json
import time
import base64
import httpx
import inspect
from pathlib import Path
from typing import List, Optional, Literal, Any, Dict

import torch
from PIL import Image
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel, Field
from transformers import AutoModel, AutoTokenizer


# ============================================================
# 1. 加载模型（启动时一次，后面 0 开销）
# ============================================================
MODEL_DIR = os.environ.get("MINICPMV_MODEL_DIR", "/data/volc_minicpm/models/MiniCPM-V-2_6-int4")
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
LOAD_DTYPE = os.environ.get("MINICPMV_DTYPE", "bf16")  # bf16 / fp16 / fp32
MAX_NEW_TOKENS = int(os.environ.get("MINICPMV_MAX_NEW", "1536"))
TEMPERATURE = float(os.environ.get("MINICPMV_TEMP", "0.2"))

_dtype_map = {
    "bf16": torch.bfloat16,
    "fp16": torch.float16,
    "fp32": torch.float32,
    "auto": "auto",
}

_model_loaded = False
_model: Optional[AutoModel] = None
_tokenizer: Optional[AutoTokenizer] = None


def load_model_once():
    global _model, _tokenizer, _model_loaded
    if _model_loaded:
        return
    t0 = time.time()
    print(f"[MiniCPM-V] Loading model from: {MODEL_DIR} (device={DEVICE}, dtype={LOAD_DTYPE}) ...", flush=True)
    if not Path(MODEL_DIR).exists() or len(list(Path(MODEL_DIR).glob("*.safetensors"))) == 0 and len(list(Path(MODEL_DIR).glob("*.bin"))) == 0:
        raise RuntimeError(
            f"找不到模型权重在: {MODEL_DIR}\n"
            f"请先下载: python3 -c \"from modelscope import snapshot_download; "
            f"snapshot_download('openbmb/MiniCPM-V-2_6-int4', local_dir='{MODEL_DIR}')\""
        )

    kwargs = dict(trust_remote_code=True, local_files_only=True)
    dtype = _dtype_map.get(LOAD_DTYPE, torch.bfloat16)
    if isinstance(dtype, torch.dtype):
        # bf16 需要 Ampere+ (A30 是 SM80, 支持 ✅)
        if dtype == torch.bfloat16 and torch.cuda.is_available() and not torch.cuda.is_bf16_supported():
            print("[MiniCPM-V] WARN: GPU 不支持 bf16, 降级 fp16", flush=True)
            dtype = torch.float16
        kwargs["torch_dtype"] = dtype
    # 使用官方推荐 SDPA attention (A30 完美支持)
    kwargs["attn_implementation"] = "sdpa"

    _tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR, trust_remote_code=True, local_files_only=True)
    _model = AutoModel.from_pretrained(MODEL_DIR, **kwargs)
    _model = _model.eval()
    if DEVICE == "cuda":
        _model = _model.cuda()

    # 预热: 跑一次空 prompt 让 CUDA kernel 编译好, 避免首张图慢
    try:
        with torch.no_grad():
            msgs = [{"role": "user", "content": ["hello"]}]
            # 没有图片的纯文本预热
            _ = _model.chat(image=None, msgs=msgs, tokenizer=_tokenizer, max_new_tokens=4)
    except Exception as e:
        print(f"[MiniCPM-V] 预热跳过 (无图模式不可用? 没关系): {e}", flush=True)

    dt = time.time() - t0
    mem = ""
    if torch.cuda.is_available():
        mem = f", 显存占用={torch.cuda.memory_allocated()/1024**3:.1f}GB / {torch.cuda.get_device_properties(0).total_memory/1024**3:.1f}GB"
    print(f"[MiniCPM-V] 模型加载完成, 耗时 {dt:.1f}s{mem}", flush=True)
    _model_loaded = True


# ============================================================
# 2. 图片解析: base64 / URL / 本地路径
# ============================================================
def _to_pil(img_url_or_b64: str) -> Image.Image:
    """把各种图片输入格式转成 PIL.Image(RGB)."""
    s = img_url_or_b64.strip()
    # case 1: 本地文件路径
    p = Path(s)
    if p.exists() and p.is_file():
        return Image.open(p).convert("RGB")
    # case 2: http(s) URL
    if s.startswith("http://") or s.startswith("https://"):
        resp = httpx.get(s, timeout=30, follow_redirects=True)
        resp.raise_for_status()
        return Image.open(io.BytesIO(resp.content)).convert("RGB")
    # case 3: data URI  base64
    m = re.match(r"^data:image/(png|jpeg|jpg|webp|gif|bmp);base64,(.+)$", s, re.I)
    if m:
        s = m.group(2)
    try:
        raw = base64.b64decode(s, validate=False)
        return Image.open(io.BytesIO(raw)).convert("RGB")
    except Exception:
        raise ValueError(f"图片格式无法解析: 长度={len(s)}, 前50字符={s[:50]!r}")


# ============================================================
# 3. OpenAI 兼容数据模型
# ============================================================
class MessagePart(BaseModel):
    type: Literal["text", "image_url"]
    text: Optional[str] = None
    image_url: Optional[Dict[str, Any]] = None  # {"url": "xxx"}


class ChatMessage(BaseModel):
    role: Literal["system", "user", "assistant"]
    # 两种格式: str 或者 List[MessagePart]
    content: Any = Field(..., description="字符串或多模态消息数组")


class ChatRequest(BaseModel):
    model: str = "minicpmv"
    messages: List[ChatMessage]
    max_tokens: Optional[int] = None
    temperature: Optional[float] = None
    top_p: Optional[float] = None
    stream: Optional[bool] = False
    response_format: Optional[Dict[str, Any]] = None  # {"type": "json_object"} 时我们会强制输出 JSON


# ============================================================
# 4. FastAPI App
# ============================================================
app = FastAPI(title="MiniCPM-V Inference API", version="1.0.0")


@app.on_event("startup")
def _startup():
    load_model_once()


@app.get("/health")
def health():
    return {
        "status": "ok" if _model_loaded else "loading",
        "device": DEVICE,
        "dtype": LOAD_DTYPE,
        "model_dir": MODEL_DIR,
        "cuda_available": torch.cuda.is_available(),
        "cuda_mem_allocated_gb": round(torch.cuda.memory_allocated()/1024**3, 2) if torch.cuda.is_available() else 0,
    }


@app.get("/v1/models")
def list_models():
    return {"object": "list", "data": [{"id": "minicpmv", "object": "model", "owned_by": "local"}]}


def _extract_messages(messages: List[ChatMessage]):
    """把 OpenAI message 格式转成 MiniCPM-V 需要的:
       (PIL_Image 或 None, chat_msglist)
       如果 message 里只有 text 没有 image, image=None
       image 只取第一条用户消息里的 image_url (单图模型)
    """
    image: Optional[Image.Image] = None
    mcpm_msgs = []
    for m in messages:
        role = m.role
        if role == "system":
            # MiniCPM-V 2.6 目前用 msgs 里 role=system 也行, 但最好塞到 user 前面
            # 简化: 把 system 合并成第一条 user 消息的前缀
            mcpm_msgs.append({"role": "user", "content": [str(m.content)]})
            mcpm_msgs.append({"role": "assistant", "content": "好的，我会遵守。"})
            continue
        content = m.content
        parts: List[Any] = []
        if isinstance(content, str):
            parts = [content]
        elif isinstance(content, list):
            for p in content:
                if isinstance(p, MessagePart):
                    if p.type == "text":
                        parts.append(p.text or "")
                    elif p.type == "image_url":
                        url = (p.image_url or {}).get("url", "")
                        if url:
                            try:
                                image = _to_pil(url)
                            except Exception as e:
                                raise HTTPException(400, f"图片加载失败: {e}")
                elif isinstance(p, dict):
                    if p.get("type") == "text":
                        parts.append(p.get("text", ""))
                    elif p.get("type") == "image_url":
                        url = (p.get("image_url") or {}).get("url", "")
                        if url:
                            try:
                                image = _to_pil(url)
                            except Exception as e:
                                raise HTTPException(400, f"图片加载失败: {e}")
        # 过滤空字符串
        parts = [p for p in parts if not (isinstance(p, str) and p == "")]
        if not parts:
            continue
        # 官方格式: content = [PIL_Image, "text"]，但我们如果有 image 就放在 user 消息的第一个元素前面
        if role == "user" and image is not None:
            merged = [image] + list(parts)
            mcpm_msgs.append({"role": "user", "content": merged})
            image = None  # 只消费一次
        else:
            # MiniCPM-V 的 msgs 里 content 可以是 list[str] 或 str
            text = "\n".join(str(p) for p in parts if isinstance(p, str))
            mcpm_msgs.append({"role": role, "content": text})
    return mcpm_msgs


@app.post("/v1/chat/completions")
def chat(req: ChatRequest):
    if not _model_loaded:
        load_model_once()
    assert _model is not None and _tokenizer is not None

    mcpm_msgs = _extract_messages(req.messages)
    if not mcpm_msgs:
        raise HTTPException(400, "messages 为空")

    # 检查最后一条是不是 user
    if mcpm_msgs[-1]["role"] != "user":
        mcpm_msgs.append({"role": "user", "content": "继续"})

    # 生成参数
    temperature = req.temperature if req.temperature is not None else TEMPERATURE
    max_new = req.max_tokens if req.max_tokens and req.max_tokens > 0 else MAX_NEW_TOKENS
    if temperature == 0:
        temperature = 0.01  # MiniCPM-V 不允许 0
    top_p = req.top_p if req.top_p is not None else 0.7

    # response_format=json_object 时, 强制输出 JSON (在 prompt 里注入指令)
    json_mode = (req.response_format or {}).get("type") == "json_object"
    if json_mode:
        last = mcpm_msgs[-1]
        extra = "\n\n请只输出合法的 JSON，不要输出任何其他内容、解释、markdown 代码块。"
        if isinstance(last["content"], list):
            last["content"].append(extra)
        else:
            last["content"] = str(last["content"]) + extra

    # 找到是否有图片需要传 image 参数:
    #   官方 API: model.chat(image=PIL, msgs=msgs, tokenizer=...)
    #   如果 msgs[-1] 里已经把 PIL 嵌到了 content list, 那 image 参数传 None 也行
    image_arg = None  # 默认 None
    # 找最后一条 user 消息里有没有 PIL 元素
    for m in reversed(mcpm_msgs):
        if m["role"] != "user":
            continue
        c = m["content"]
        if isinstance(c, list):
            pil_imgs = [x for x in c if isinstance(x, Image.Image)]
            if pil_imgs:
                image_arg = pil_imgs[0]
                # 注意：同时在 content 里把 PIL 保留, 一般没问题
                break

    # 调用模型
    t0 = time.time()
    try:
        with torch.no_grad():
            out_text = _model.chat(
                image=image_arg,
                msgs=mcpm_msgs,
                tokenizer=_tokenizer,
                max_new_tokens=max_new,
                temperature=temperature,
                top_p=top_p,
                do_sample=True if temperature > 0.01 else False,
            )
    except torch.cuda.OutOfMemoryError as e:
        raise HTTPException(500, f"CUDA OOM: {e}")
    except Exception as e:
        raise HTTPException(500, f"模型推理失败: {type(e).__name__}: {e}")
    dt = time.time() - t0

    # json mode: 清理一下可能的 markdown ```json ``` 包裹
    if json_mode:
        m = re.search(r"\{.*\}", out_text, re.S)
        if m:
            try:
                json.loads(m.group(0))
                out_text = m.group(0)
            except Exception:
                pass
        # 验证合法 JSON
        try:
            json.loads(out_text)
        except Exception as e:
            out_text = json.dumps({
                "error": f"模型输出不是合法 JSON: {e}",
                "raw": out_text,
            }, ensure_ascii=False)

    # OpenAI 兼容响应体
    resp = {
        "id": f"cmpl-{os.urandom(6).hex()}",
        "object": "chat.completion",
        "created": int(time.time()),
        "model": "minicpmv",
        "choices": [{
            "index": 0,
            "message": {"role": "assistant", "content": out_text},
            "finish_reason": "stop",
        }],
        "usage": {
            "prompt_tokens": 0,   # 懒得算, 不影响功能
            "completion_tokens": 0,
            "total_tokens": 0,
        },
        "metrics": {"elapsed_seconds": round(dt, 3), "device": DEVICE},
    }
    if req.stream:
        # 简化版 non-stream 响应, 客户端如果真的要 stream 我们退化成单次
        return resp
    return resp


# ============================================================
# 5. 命令行直接运行
# ============================================================
if __name__ == "__main__":
    import argparse
    ap = argparse.ArgumentParser()
    ap.add_argument("--host", default="0.0.0.0")
    ap.add_argument("--port", type=int, default=30000)
    ap.add_argument("--model-dir", default=None)
    args = ap.parse_args()
    if args.model_dir:
        os.environ["MINICPMV_MODEL_DIR"] = args.model_dir
    load_model_once()
    import uvicorn
    uvicorn.run(app, host=args.host, port=args.port, log_level="info")
