from fastapi import FastAPI, HTTPException
from pydantic import BaseModel, Field
from typing import List, Optional
import torch
from transformers import AutoTokenizer, AutoModelForSequenceClassification
from modelscope import snapshot_download
import logging
import os

# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

app = FastAPI(title="BGE Rerank Service", version="1.0.0")

# 请求和响应模型
class RerankRequest(BaseModel):
    query: str
    documents: List[str]
    top_n: Optional[int] = None
    score_threshold: Optional[float] = None

class RerankDocument(BaseModel):
    text: str
    score: float
    index: int
    relevance_score: float  # 兼容OpenAI格式，必填字段

class RerankResponse(BaseModel):
    results: List[RerankDocument]
    query: str
    total_documents: int

# 全局变量存储模型和分词器
tokenizer = None
model = None
device = None

def load_model():
    """使用ModelScope下载并加载BGE-rerank-v2-m3模型"""
    global tokenizer, model, device
    
    try:
        logger.info("开始下载BGE-rerank-v2-m3模型...")
        
        # 使用ModelScope下载模型
        model_dir = snapshot_download('BAAI/bge-reranker-v2-m3')
        logger.info(f"模型下载完成，路径: {model_dir}")
        
        # 检测设备
        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        logger.info(f"使用设备: {device}")
        
        # 加载分词器和模型
        logger.info("加载分词器...")
        tokenizer = AutoTokenizer.from_pretrained(model_dir)
        
        logger.info("加载模型...")
        model = AutoModelForSequenceClassification.from_pretrained(model_dir)
        model.to(device)
        model.eval()
        
        logger.info("模型加载完成!")
        
    except Exception as e:
        logger.error(f"模型加载失败: {str(e)}")
        raise e

@app.on_event("startup")
async def startup_event():
    """应用启动时加载模型"""
    load_model()

@app.get("/")
async def root():
    """健康检查端点"""
    return {
        "message": "BGE Rerank Service is running",
        "model": "BAAI/bge-reranker-v2-m3",
        "device": str(device) if device else "unknown"
    }

@app.get("/health")
async def health_check():
    """健康检查端点"""
    if model is None or tokenizer is None:
        raise HTTPException(status_code=503, detail="Model not loaded")
    
    return {
        "status": "healthy",
        "model_loaded": True,
        "device": str(device)
    }

@app.post("/v1/rerank", response_model=RerankResponse)
async def rerank_documents(request: RerankRequest):
    """重排序文档"""
    if model is None or tokenizer is None:
        raise HTTPException(status_code=503, detail="Model not loaded")
    
    if not request.documents:
        raise HTTPException(status_code=400, detail="Documents list cannot be empty")
    
    if not request.query.strip():
        raise HTTPException(status_code=400, detail="Query cannot be empty")
    
    try:
        # 构建查询-文档对
        pairs = [[request.query, doc] for doc in request.documents]
        
        # 分词和编码
        inputs = tokenizer(
            pairs, 
            padding=True, 
            truncation=True, 
            return_tensors='pt', 
            max_length=512
        )
        
        # 移动到设备
        inputs = {k: v.to(device) for k, v in inputs.items()}
        
        # 推理
        with torch.no_grad():
            outputs = model(**inputs)
            scores = torch.sigmoid(outputs.logits.squeeze(-1))
        
        # 转换为CPU并获取分数
        scores = scores.cpu().tolist()
        
        # 创建结果列表
        results = []
        for i, (doc, score) in enumerate(zip(request.documents, scores)):
            # 应用分数阈值过滤
            if request.score_threshold is None or score >= request.score_threshold:
                results.append(RerankDocument(
                    text=doc,
                    score=score,
                    index=i,
                    relevance_score=score  # 显式设置relevance_score
                ))
        
        # 按分数降序排序
        results.sort(key=lambda x: x.score, reverse=True)
        
        # 应用top_n限制
        if request.top_n is not None:
            results = results[:request.top_n]
        
        # Pydantic模型不允许动态添加字段，relevance_score已在模型定义中设置
        
        return RerankResponse(
            results=results,
            query=request.query,
            total_documents=len(request.documents)
        )
        
    except Exception as e:
        logger.error(f"重排序过程中出错: {str(e)}")
        raise HTTPException(status_code=500, detail=f"Reranking failed: {str(e)}")

@app.post("/api/rerank")
async def rerank_compatible(request: RerankRequest):
    """兼容Dify的rerank API端点"""
    response = await rerank_documents(request)
    
    # 转换为Dify期望的格式
    docs = []
    for result in response.results:
        docs.append({
            "index": result.index,
            "text": result.text,
            "score": result.score,
            "relevance_score": result.score  # 添加relevance_score字段以兼容OpenAI格式
        })
    
    return {
        "model": "bge-reranker-v2-m3",
        "docs": docs,
        "query": response.query,
        "total": response.total_documents
    }

@app.post("/rerank")
async def rerank_legacy(request: RerankRequest):
    """兼容性端点 - /rerank"""
    return await rerank_documents(request)

# 添加根路径的API信息
@app.get("/")
async def root():
    """API根路径信息"""
    return {
        "service": "BGE Rerank Service",
        "version": "1.0.0",
        "endpoints": {
            "health": "/health",
            "rerank_v1": "/v1/rerank",
            "rerank_legacy": "/rerank", 
            "rerank_api": "/api/rerank"
        },
        "model": "BAAI/bge-reranker-v2-m3",
        "status": "running"
    }

if __name__ == "__main__":
    import uvicorn
    
    # 启动服务
    uvicorn.run(
        "rerank_service:app",
        host="0.0.0.0",
        port=8000,
        reload=False,  # 生产环境建议设为False
        log_level="info"
    )
