#!/usr/bin/env python3
# 在 GPU 机上对任意 OpenAI 兼容 llama-server 跑同一组中文提示词，落盘结果
# 用法: quality_test.py <端口> <模型名> <输出目录>
import json, sys, time, os, urllib.request, urllib.error

port = sys.argv[1]; model = sys.argv[2]; outdir = sys.argv[3]
os.makedirs(outdir, exist_ok=True)

PROMPTS = [
 ("A-中文创作",
  "写一个 300 字左右的短篇开头：一位退休的钟表匠在暴雨夜发现店里多了一座他没修过的钟。要求：第一人称，克制、有画面感，不要解释设定。", False),
 ("B-文档改写",
  "把下面这段改写成给非技术同事看的 5 条要点，每条不超过 25 字：三值量化把模型权重限制在 -1、0、+1 三个值上，每 128 个权重共享一个 FP16 缩放因子；权重在 Hadamard 旋转基下存储，推理时对激活做匹配变换；这样 27B 模型只需 5.95GB。", False),
 ("C-数学推理",
  "一个水池有两个进水管和一个出水管。甲管单独注满需 6 小时，乙管单独注满需 8 小时，出水管排空满池需 12 小时。三管同时打开，多久注满？请给出推理过程。", False),
 ("D-指令跟随",
  "请用恰好 5 条 bullet 总结 llama.cpp 本地推理相比云端 API 的优势。硬性要求：每条必须以动词开头、每条不超过 20 个字、且 5 条合计提到的优势不超过 3 个不同方面。", False),
 ("E-数学推理-带思考",
  "一个水池有两个进水管和一个出水管。甲管单独注满需 6 小时，乙管单独注满需 8 小时，出水管排空满池需 12 小时。三管同时打开，多久注满？请给出推理过程。", True),
]

url = "http://127.0.0.1:%s/v1/chat/completions" % port

def call(prompt, thinking, use_kwargs=True):
    payload = {
        "model": model,
        "messages": [{"role": "user", "content": prompt}],
        "temperature": 0.7, "top_p": 0.9, "top_k": 20,
        "max_tokens": 2048, "seed": 42,
    }
    if use_kwargs:
        payload["chat_template_kwargs"] = {"enable_thinking": bool(thinking)}
    req = urllib.request.Request(url, data=json.dumps(payload).encode(),
                                headers={"Content-Type": "application/json"})
    return json.load(urllib.request.urlopen(req, timeout=1800))

for name, prompt, thinking in PROMPTS:
    t0 = time.time()
    try:
        try:
            r = call(prompt, thinking)
        except urllib.error.HTTPError as e:
            print("[warn] %s: kwargs 被拒(%s)，回退无 kwargs" % (name, e.code))
            r = call(prompt, thinking, use_kwargs=False)
        dt = time.time() - t0
        msg = r["choices"][0]["message"]
        content = msg.get("content") or ""
        reasoning = msg.get("reasoning_content") or msg.get("reasoning") or ""
        usage = r.get("usage", {}) or {}
        tim = r.get("timings", {}) or {}
        keep = {k: tim.get(k) for k in ("prompt_ms", "predicted_ms",
                "prompt_per_second", "predicted_per_second")}
        with open(os.path.join(outdir, name + ".md"), "w", encoding="utf-8") as f:
            f.write("# %s\n\n**提示词**\n\n%s\n\n**思考模式**: %s\n\n---\n\n## 正文\n\n%s\n\n---\n\n## 思考链(reasoning_content, 前 800 字)\n\n%s\n\n---\n\n## 指标\n\n- 墙钟耗时: %.1fs\n- usage: `%s`\n- timings: `%s`\n"
                    % (name, prompt, thinking, content, reasoning[:800], dt,
                       json.dumps(usage, ensure_ascii=False), json.dumps(keep, ensure_ascii=False)))
        print("[OK] %-16s %6.1fs  completion=%s  %.1f tok/s"
              % (name, dt, usage.get("completion_tokens"), keep.get("predicted_per_second") or 0))
    except Exception as e:
        print("[FAIL] %s: %r" % (name, e))
