#!/usr/bin/env python3
"""
大模型 Flask API 服务器

这个脚本实现了一个 REST API 服务，提供文本生成和聊天补全功能。
API 兼容 OpenAI API 的部分接口规范，便于集成到现有的 AI 应用中。

启动方式：
  python flask_api_server.py
  
API 端点：
  - GET  /health              - 健康检查
  - POST /v1/chat/completions - 聊天补全（对话模式）
  - POST /v1/completions      - 文本补全（生成模式）
"""

from flask import Flask, request, jsonify
from transformers import AutoTokenizer, AutoModelForCausalLM
import torch
import json

# ============================================================================
# Flask 应用初始化与模型加载
# ============================================================================

app = Flask(__name__)

# 全局加载模型和分词器（在应用启动时执行）
# 这样做的优势：
# 1. 避免每个请求都重新加载模型，这会非常耗时
# 2. 模型在内存中保持加载状态，减少延迟
# 缺点：
# 1. 应用启动时间较长
# 2. 内存占用始终保持较高水位
# 3. 不支持模型热更新（需要重启应用）
print("Loading model...")
model_name = "Qwen/Qwen2.5-1.5B-Instruct"
tokenizer = AutoTokenizer.from_pretrained(model_name)

# 在实际部署中，建议改进为：
# - 使用 bfloat16 替代 float16，提升稳定性
# - 添加 device_map="auto" 和 gradient_checkpointing_enable()
# - 设置 max_memory 限制，防止 OOM
model = AutoModelForCausalLM.from_pretrained(model_name, dtype=torch.float16)
model.eval()  # 启用评估模式
print("Model loaded!")

# ============================================================================
# 端点 1: 健康检查
# ============================================================================

@app.route('/health', methods=['GET'])
def health():
    """
    健康检查端点
    
    返回：
      {
        "status": "ok"
      }
    
    用途：
    - 负载均衡器可定期调用此端点检测服务是否可用
    - 容器编排系统（K8s）可用此判断是否需要重启 Pod
    - 监控系统可跟踪服务可用性
    """
    return jsonify({"status": "ok"})

# ============================================================================
# 端点 2: 聊天补全（OpenAI 兼容）
# ============================================================================

@app.route('/v1/chat/completions', methods=['POST'])
def chat_completions():
    """
    聊天补全端点 - 对话模式 API
    
    兼容 OpenAI Chat Completions API：
    https://platform.openai.com/docs/api-reference/chat/create
    
    请求体格式：
    {
      "messages": [
        {"role": "user", "content": "你好"},
        {"role": "assistant", "content": "你好！有什么我可以帮你的吗？"},
        {"role": "user", "content": "今天天气怎么样？"}
      ],
      "temperature": 0.7,           # 可选，默认 0.7
      "max_tokens": 256             # 可选，默认 256
    }
    
    返回格式（OpenAI 兼容）：
    {
      "choices": [
        {
          "message": {
            "role": "assistant",
            "content": "我是一个 AI 助手，无法获取实时天气信息..."
          },
          "index": 0,
          "finish_reason": "length"
        }
      ]
    }
    """
    # 解析请求 JSON 数据
    data = request.json
    
    # 从请求中提取参数，使用默认值
    # messages: 对话历史，格式为 [{"role": "user"/"assistant", "content": "..."}, ...]
    messages = data.get('messages', [])
    
    # max_tokens: 最多生成多少个新 token
    max_tokens = data.get('max_tokens', 256)
    
    # temperature: 生成的随机性（0.0-2.0）
    temperature = data.get('temperature', 0.7)
    
    # 步骤 1：使用 chat template 格式化消息
    # 这会将 messages 列表转换为模型期望的提示格式
    # 包含系统角色、对话历史和生成提示
    text = tokenizer.apply_chat_template(
        messages,
        tokenize=False,  # 返回字符串而不是 token ID
        add_generation_prompt=True  # 添加特殊 token 指示开始生成
    )
    
    # 步骤 2：分词化输入文本
    model_inputs = tokenizer([text], return_tensors="pt")
    
    # 步骤 3：执行推理生成文本
    # torch.no_grad() 禁用梯度计算以加速推理和节省显存
    with torch.no_grad():
        generated_ids = model.generate(
            model_inputs.input_ids,
            attention_mask=model_inputs.get("attention_mask"),  # 明确传递注意力掩码
            max_new_tokens=max_tokens,
            temperature=temperature,
            top_p=0.95,
            do_sample=True
        )
    
    # 步骤 4：解码生成的 token 序列为文本
    full_response = tokenizer.decode(generated_ids[0], skip_special_tokens=True)
    
    # 步骤 5：提取仅 assistant 部分的响应
    # 完整的 full_response 包含了输入的 chat template 和生成的文本
    # 我们只需要返回 assistant 的回复部分
    response = full_response.split("assistant\n")[-1] if "assistant" in full_response else full_response
    
    # 返回 OpenAI 兼容的响应格式
    return jsonify({
        "choices": [
            {
                "message": {
                    "role": "assistant",
                    "content": response
                },
                "index": 0,
                "finish_reason": "length"
            }
        ]
    })

# ============================================================================
# 端点 3: 文本补全（OpenAI 兼容）
# ============================================================================

@app.route('/v1/completions', methods=['POST'])
def completions():
    """
    文本补全端点 - 生成模式 API
    
    兼容 OpenAI Completions API：
    https://platform.openai.com/docs/api-reference/completions/create
    
    请求体格式：
    {
      "prompt": "今天天气很",
      "temperature": 0.7,           # 可选，默认 0.7
      "max_tokens": 256             # 可选，默认 256
    }
    
    返回格式（OpenAI 兼容）：
    {
      "choices": [
        {
          "text": "晴朗，适合出游。",
          "index": 0,
          "finish_reason": "length"
        }
      ]
    }
    
    与 chat/completions 的区别：
    - chat/completions: 用于对话场景，接收多轮对话消息
    - completions: 用于文本补全，给定一个开头文本，模型继续生成
    """
    # 解析请求 JSON 数据
    data = request.json
    
    # prompt: 输入的文本开头，模型会基于此生成后续内容
    prompt = data.get('prompt', '')
    
    # max_tokens: 最多生成多少个新 token
    max_tokens = data.get('max_tokens', 256)
    
    # temperature: 生成的随机性（0.0-2.0）
    temperature = data.get('temperature', 0.7)
    
    # 步骤 1：直接分词化输入 prompt（不需要 chat template）
    # 因为这是纯文本补全，不涉及聊天格式
    model_inputs = tokenizer([prompt], return_tensors="pt")
    
    # 步骤 2：执行推理生成文本
    with torch.no_grad():
        generated_ids = model.generate(
            model_inputs.input_ids,
            attention_mask=model_inputs.get("attention_mask"),  # 明确传递注意力掩码
            max_new_tokens=max_tokens,
            temperature=temperature,
            top_p=0.95,
            do_sample=True
        )
    
    # 步骤 3：解码生成的 token 序列为文本
    completion = tokenizer.decode(generated_ids[0], skip_special_tokens=True)
    
    # 步骤 4：提取仅生成部分（去掉输入的 prompt）
    # generated_ids 包含 [input_ids + generated_ids]
    # 所以解码后的文本包含了原始 prompt 和生成内容
    # 我们只返回生成的新部分：completion[len(prompt):]
    
    # 返回 OpenAI 兼容的响应格式
    return jsonify({
        "choices": [
            {
                "text": completion[len(prompt):],  # 仅返回新生成的部分
                "index": 0,
                "finish_reason": "length"
            }
        ]
    })

# ============================================================================
# 应用启动
# ============================================================================

if __name__ == '__main__':
    """
    Flask 应用入口点
    
    启动 Flask 开发服务器监听所有网络接口
    
    参数说明：
    - host='0.0.0.0': 监听所有网络接口（0.0.0.0 代表所有 IPv4 地址）
      * 生产环境建议使用 127.0.0.1 限制本地访问，再通过反向代理（nginx）暴露
    - port=8000: 监听端口号
    - debug=False: 关闭调试模式（生产环境必须关闭，开启会自动重载代码）
    - threaded=False: 不启用多线程
      * PyTorch 模型加载和推理不是线程安全的
      * 多个线程同时执行模型推理会导致错误
      * 需要使用队列或进程池来处理并发请求
    
    生产部署建议：
    1. 使用 gunicorn 替代 Flask 内置服务器
    2. 配置多个 worker 进程处理请求
    3. 前置反向代理（nginx）进行负载均衡
    4. 使用 Redis 缓存常见请求
    5. 添加请求队列管理系统（如 Celery）处理长时间运行的请求
    6. 启用 request timeout 防止无限挂起
    
    启动命令：
      # 开发模式
      python flask_api_server.py
      
      # 生产模式（需先 pip install gunicorn）
      gunicorn -w 4 -b 0.0.0.0:8000 flask_api_server:app
    """
    print("Starting API server on http://0.0.0.0:8000")
    app.run(host='0.0.0.0', port=8000, debug=False, threaded=False)
