bench_ext.py 17 KB
"""扩展基准:多步轨迹 + 多轮对话用例。

bench50_v2.py 只测「单发话语 → 首个工具选择」;本脚本补它测不到的:
- 多步 ReAct 轨迹(喂回固定工具结果,跑到模型给出最终文字为止)
- 多轮接续(翻页、代词回指、澄清后接续)
- 回答夹带新需求
- 「改审核人」类消歧(改带"审核"字样的字段 ≠ 审核操作)
- 长会话回指(30 轮后引用最早实体)
- 跨轮流程接续(弹表单后隔轮补字段,值须逐字保留)
- 提议卡后的状态纪律(不重复提议、不声称已执行)

用法:
  uv run --no-project bench_ext.py --arch old        # 现产线:统一 prompt + 6 工具
  uv run --no-project bench_ext.py --arch new        # skill-ReAct(Phase 2 落地后从仓库资源读 prompt/skill)
  uv run --no-project bench_ext.py --arch old --only M5
"""

import json
import sys
import time
import urllib.request

URL = "http://112.82.245.194:41434/v1/chat/completions"
MODEL = "qwen3.6-27b-iq3:latest"
MAX_STEPS = 8

# ---------------- 与 bench50_v2 相同的生产复刻 prompt / 工具定义 ----------------
from bench50_v2 import UNIFIED_PROMPT, ALL6  # noqa: E402


def build_context(arch):
    """返回 (system_prompt, tools)。new 架构落地后从仓库资源文件读取,避免复刻漂移。"""
    if arch == "old":
        return UNIFIED_PROMPT, ALL6
    if arch == "new":
        import pathlib
        root = pathlib.Path(__file__).resolve().parent.parent
        sp = (root / "src/main/resources/prompts/system.txt").read_text(encoding="utf-8")
        tools_json = (root / "src/main/resources/prompts/tools-bench.json").read_text(encoding="utf-8")
        return sp, json.loads(tools_json)
    raise SystemExit(f"未知 arch: {arch}")


# ---------------- 固定工具结果(fixtures) ----------------
CUSTOMERS = {
    "必胜客": {"客户名称": "必胜客", "联系电话": "13912345678", "地址": "上海市静安区南京西路1266号", "销售员": "王五"},
    "上海创远包装": {"客户名称": "上海创远包装", "联系电话": "021-58991234", "地址": "浦东新区川沙路500号", "销售员": "李四"},
    "苏州华为": {"客户名称": "苏州华为", "联系电话": "0512-67771234", "地址": "苏州工业园区", "销售员": "赵六"},
    "杭州大华印务": {"客户名称": "杭州大华印务", "联系电话": "0571-87654321", "地址": "杭州市余杭区", "销售员": "钱七"},
}

FORMS = {
    "报价": [{"formName": "报价单", "formId": "FQ01", "moduleId": "M-BJ", "source": "quoquotationmaster"}],
    "客户": [{"formName": "客户资料", "formId": "FC01", "moduleId": "M-JC", "source": "custmaster"}],
    "送货": [{"formName": "送货单", "formId": "FD01", "moduleId": "M-XS", "source": "deliverymaster"}],
    "销售": [{"formName": "销售订单", "formId": "FS01", "moduleId": "M-XS", "source": "salesorder"}],
}

READFORM_PAGES = {
    ("FC01", 1): "共 6 条记录,第 1 页:必胜客、上海创远包装、苏州华为、杭州大华印务",
    ("FC01", 2): "共 6 条记录,第 2 页:宁波天海包装、无锡礼盒厂",
    ("FD01", 1): "共 2 条记录,第 1 页:DH202607005(客户 必胜客)、DH202607008(客户 苏州华为)",
}


def fixture_result(name, args):
    args = args or {}
    if name == "findForms":
        kw = args.get("keyword", "")
        for k, forms in FORMS.items():
            if k in kw:
                return json.dumps(forms, ensure_ascii=False)
        return json.dumps([{"formName": kw + "列表", "formId": "FX99", "moduleId": "M-QT", "source": "misc"}],
                          ensure_ascii=False)
    if name == "readFormData":
        page = int(args.get("page") or 1)
        key = (args.get("formId", ""), page)
        if key in READFORM_PAGES:
            return READFORM_PAGES[key]
        kw = args.get("keyword") or ""
        for cname, rec in CUSTOMERS.items():
            if cname in kw:
                return f"共 1 条记录:{json.dumps(rec, ensure_ascii=False)}"
        return f"共 0 条记录(formId={args.get('formId','')} page={page} keyword={kw})"
    if name == "lookupRecord":
        rk = args.get("recordKeyword", "")
        for cname, rec in CUSTOMERS.items():
            if cname in rk or rk in cname:
                return json.dumps(rec, ensure_ascii=False)
        return f"未找到「{rk}」对应的记录"
    if name == "collectForm":
        return (f"[已弹出「{args.get('entityKeyword','')}」表单,预填={args.get('knownFieldsJson') or '{}'}。"
                "等待用户在表单里补齐并点提交,本轮到此为止。]")
    if name == "proposeWrite":
        return (f"[已生成待确认提议 OP-TEST-1:action={args.get('action','')},实体={args.get('entityKeyword','')},"
                f"记录={args.get('recordKeyword','')}。仅提议,未执行;请用户点【确认】。]")
    if name == "askUser":
        return "[问题已发给用户,等待回答,本轮到此为止。]"
    if name == "useSkill":
        return "[skill 文本已载入上下文]"
    return "OK"


# ---------------- 轨迹模拟器 ----------------
def call(body, timeout=200):
    req = urllib.request.Request(
        URL, data=json.dumps(body).encode(),
        headers={"Content-Type": "application/json", "Authorization": "Bearer ollama"})
    with urllib.request.urlopen(req, timeout=timeout) as r:
        return json.loads(r.read())


def run_trajectory(system_prompt, tools, history, utterance):
    """history: [(role, text)] 已渲染的往轮。返回 (calls, final_text, steps, err)。"""
    messages = [{"role": "system", "content": system_prompt}]
    for role, text in history:
        messages.append({"role": role, "content": text})
    messages.append({"role": "user", "content": utterance})
    calls = []
    for step in range(MAX_STEPS):
        body = {"model": MODEL, "stream": False, "temperature": 0.1, "top_p": 0.9,
                "reasoning_effort": "none", "tools": tools, "messages": messages}
        try:
            resp = call(body)
        except Exception as ex:
            return calls, None, step, str(ex)
        msg = resp["choices"][0]["message"]
        tcs = msg.get("tool_calls") or []
        if not tcs:
            return calls, (msg.get("content") or "").strip(), step + 1, None
        messages.append({"role": "assistant", "content": msg.get("content") or "",
                         "tool_calls": tcs})
        for tc in tcs:
            name = tc["function"]["name"]
            try:
                args = json.loads(tc["function"].get("arguments") or "{}")
            except Exception:
                args = {}
            calls.append((name, args))
            messages.append({"role": "tool", "tool_call_id": tc.get("id", "tc"),
                            "content": fixture_result(name, args)})
    return calls, None, MAX_STEPS, "step-cap"


# ---------------- 判分辅助 ----------------
def argstr(args):
    return json.dumps(args or {}, ensure_ascii=False)


def find(calls, name, pred=None):
    for n, a in calls:
        if n == name and (pred is None or pred(a)):
            return a
    return None


def any_read(calls, needle=None):
    for n, a in calls:
        if n in ("findForms", "readFormData", "lookupRecord"):
            if needle is None or needle in argstr(a):
                return a
    return None


def asked(calls, final_text):
    return find(calls, "askUser") is not None or (final_text and ("?" in final_text or "?" in final_text))


# ---------------- 用例 ----------------
# 每条: id, cat, history, utter, judge(calls, text) -> ("PASS"|"PARTIAL"|"FAIL", reason)

def j_shenheren_update(target_val):
    def judge(calls, text):
        p = find(calls, "proposeWrite")
        if p and p.get("action") == "update":
            if target_val in argstr(p):
                return "PASS", ""
            return "PARTIAL", f"update 但参数缺 {target_val}"
        if p and p.get("action") in ("examine", "cancelExamine"):
            return "FAIL", f"误判为审核操作 action={p.get('action')}"
        if asked(calls, text):
            return "PARTIAL", "反问(信息已足够)"
        return "FAIL", "未产生 update 提议"
    return judge


def mk(id_, cat, history, utter, judge, desc):
    return {"id": id_, "cat": cat, "history": history, "utter": utter, "judge": judge, "desc": desc}


def long_history():
    pairs = [
        ("查一下杭州大华印务的资料", "杭州大华印务:印刷类客户,地址杭州市余杭区,销售员钱七。"),
        ("必胜客上个月下了几单?", "必胜客上个月共 3 张销售订单。"),
        ("铜版纸库存还有多少?", "铜版纸当前库存 12,000 张。"),
        ("上海创远包装的地址?", "浦东新区川沙路500号。"),
        ("这个月开了几张送货单?", "本月共 18 张送货单。"),
        ("苏州华为的销售员是谁?", "赵六。"),
        ("白卡纸多少库存?", "300g 白卡纸库存 8,500 张。"),
        ("最近有新采购订单吗?", "最近一周有 2 张采购订单。"),
        ("应收账款总额?", "当前应收账款合计 156 万元。"),
        ("宁波天海包装是我们客户吗?", "是,客户资料里有宁波天海包装。"),
        ("今天的生产任务?", "今日 4 个生产任务,均在排程中。"),
        ("金田纸业的电话?", "0571-88881234。"),
        ("上季度销售额最高的客户?", "必胜客,合计 89 万元。"),
        ("画册类产品最近报价几张?", "最近一个月 5 张画册报价单。"),
        ("无锡礼盒厂有欠款吗?", "无锡礼盒厂当前无逾期欠款。"),
    ]
    h = []
    for q, a in pairs:
        h.append(("user", q))
        h.append(("assistant", a))
    return h


def j_m4(calls, text):
    if any_read(calls, "杭州大华印务"):
        return "PASS", ""
    if "钱七" in (text or ""):
        return "PASS", "从历史正确回忆(未重查)"
    if asked(calls, text):
        return "PARTIAL", "反问是哪个客户(历史里可查)"
    return "FAIL", "既没读 杭州大华印务 也没答对销售员"


def j_m5(calls, text):
    c = find(calls, "collectForm")
    if c:
        s = argstr(c)
        if "大16开" in s and "5000" in s:
            return "PASS", ""
        if "16开" in s and "大16开" not in s:
            return "FAIL", "尺寸值被改写(大16开→16开),违反逐字照抄"
        return "PARTIAL", f"collectForm 但预填不全: {s[:80]}"
    p = find(calls, "proposeWrite", lambda a: a.get("action") == "create")
    if p and "大16开" in argstr(p):
        return "PASS", "直接 create 且值逐字"
    return "FAIL", "未回到报价表单流程"


def j_m7(calls, text):
    if find(calls, "proposeWrite"):
        return "FAIL", "重复生成提议"
    done_words = [w for w in ("已修改", "已完成", "修改成功") if w in (text or "")]
    if done_words:
        return "FAIL", f"声称已执行: {done_words}"
    if any_read(calls, "必胜客"):
        return "PASS", ""
    return "FAIL", "未查询必胜客地址"


def j_m3(calls, text):
    upd = find(calls, "proposeWrite",
               lambda a: a.get("action") == "update" and "13800138000" in argstr(a) and "必胜客" in argstr(a))
    qry = any_read(calls, "上海")
    if upd and qry:
        return "PASS", ""
    if upd:
        return "PARTIAL", "只完成修改提议,丢了夹带的查询"
    if qry:
        return "PARTIAL", "只做了查询,丢了修改"
    return "FAIL", "两个请求都没完成"


CASES = [
    # --- 「审核」字样消歧:改带审核字样的字段是 update,不是 examine ---
    mk("S1", "消歧", [], "把BJD202607001的审核人改成张三",
       j_shenheren_update("张三"), "改审核人=update"),
    mk("S2", "消歧", [], "报价单BJD202607082的审核人换成李四",
       j_shenheren_update("李四"), "换审核人=update"),
    mk("S3", "消歧", [],
       "查一下BJD202607001是谁审核的",
       lambda calls, text: ("PASS", "") if any_read(calls) or (text and not any(ch.isdigit() for ch in text))
       else ("FAIL", "既没读也没答"), "查审核人=查询"),
    mk("S4", "消歧", [], "BJD202607001不用审了,帮我撤回来",
       lambda calls, text:
       ("PASS", "") if find(calls, "proposeWrite", lambda a: a.get("action") == "cancelExamine")
       else (("PARTIAL", "反问") if asked(calls, text) else ("FAIL", "未识别为销审")),
       "撤回审核=cancelExamine"),
    # 上下文无候选、记录指代不清 → 反问哪张与直接 cancelInvalid 提议同为正确行为
    mk("S5", "消歧", [], "上次作废的那张送货单帮我恢复一下",
       lambda calls, text:
       ("PASS", "") if find(calls, "proposeWrite", lambda a: a.get("action") == "cancelInvalid")
       or asked(calls, text) else ("FAIL", "未识别为复原"),
       "恢复作废=cancelInvalid或问哪张"),
    # --- 多轮接续 ---
    mk("M1", "多轮接续",
       [("user", "查一下客户列表"),
        ("assistant", "共 6 个客户,第 1 页:必胜客、上海创远包装、苏州华为、杭州大华印务。(表单:客户资料 formId=FC01 moduleId=M-JC)")],
       "下一页",
       lambda calls, text:
       ("PASS", "") if find(calls, "readFormData", lambda a: int(a.get("page") or 1) == 2)
       else ("FAIL", "没有翻到第 2 页"), "翻页接续"),
    mk("M2", "多轮接续",
       [("user", "查一下必胜客的资料"),
        ("assistant", "必胜客:地址 上海市静安区南京西路1266号,销售员 王五。")],
       "他的电话是多少",
       lambda calls, text:
       (("PASS", "") if "13912345678" in (text or "") else ("PARTIAL", "读了但答案未含固定电话"))
       if any_read(calls, "必胜客") else ("FAIL", "未解析代词「他」=必胜客"),
       "代词回指"),
    mk("M6", "多轮接续",
       [("user", "作废那张送货单"),
        ("assistant", "请问是哪一张送货单?候选:DH202607005(必胜客)、DH202607008(苏州华为)")],
       "第一张",
       lambda calls, text:
       ("PASS", "") if find(calls, "proposeWrite",
                            lambda a: a.get("action") in ("invalid",) and "DH202607005" in argstr(a))
       else (("PARTIAL", "再次反问") if asked(calls, text) else ("FAIL", "未把「第一张」解析为 DH202607005")),
       "澄清后接续"),
    # --- 夹带新需求 ---
    mk("M3", "夹带", [("user", "把必胜客的电话改一下"), ("assistant", "请问要改成什么号码?")],
       "13800138000,另外查下上海有哪些客户", j_m3, "回答+新查询"),
    # --- 长会话回指 ---
    mk("M4", "长会话回指", long_history(), "最早问的那个客户,帮我查下它的销售员", j_m4, "30 轮后回指最早实体"),
    # --- 跨轮表单接续 ---
    mk("M5", "跨轮表单",
       [("user", "报价纸盒"),
        ("assistant", "[已为您弹出「报价」表单,请填写客户、产品、数量、尺寸等字段后提交。]"),
        ("user", "先查下必胜客的电话"),
        ("assistant", "必胜客电话:13912345678。")],
       "报价那个,尺寸大16开,数量5000", j_m5, "隔轮补字段,值逐字"),
    # --- 提议卡后的状态纪律 ---
    mk("M7", "提议纪律",
       [("user", "把必胜客电话改成13800138000"),
        ("assistant", "[已生成待确认提议 OP123:修改 客户「必胜客」联系电话 → 13800138000],请点击【确认】执行。")],
       "顺便查下它的地址", j_m7, "提议后查询:不重提、不谎称已执行"),
]


def main():
    arch = "old"
    only = None
    argv = sys.argv[1:]
    if "--arch" in argv:
        arch = argv[argv.index("--arch") + 1]
    if "--only" in argv:
        only = argv[argv.index("--only") + 1]
    system_prompt, tools = build_context(arch)
    cases = [c for c in CASES if only is None or c["id"] == only]
    print(f"arch={arch}  model={MODEL}  cases={len(cases)}")
    results = []
    for c in cases:
        t0 = time.time()
        calls, text, steps, err = run_trajectory(system_prompt, tools, c["history"], c["utter"])
        dt = time.time() - t0
        if err:
            verdict, reason = "FAIL", f"ERROR {err}"
        else:
            verdict, reason = c["judge"](calls, text)
        results.append((c, verdict, reason))
        mark = {"PASS": "✓", "PARTIAL": "◐", "FAIL": "✗"}[verdict]
        traj = " → ".join(n for n, _ in calls) or "文字"
        print(f'{mark} [{c["cat"]}] {c["id"]} {c["desc"]}: {traj}'
              f'{"  |  " + reason if reason else ""}  ({steps}步 {dt:.0f}s)')
        if verdict != "PASS" and text:
            print(f'    最终答复: {text[:100]}')
    n = len(results)
    strict = sum(v == "PASS" for _, v, _ in results)
    lenient = sum(v in ("PASS", "PARTIAL") for _, v, _ in results)
    by_cat = {}
    for c, v, _ in results:
        s, l, t = by_cat.get(c["cat"], (0, 0, 0))
        by_cat[c["cat"]] = (s + (v == "PASS"), l + (v != "FAIL"), t + 1)
    print(f"\n[{arch}] 严格 {strict}/{n}  宽松 {lenient}/{n}")
    for cat, (s, l, t) in by_cat.items():
        print(f"  {cat}: 严格 {s}/{t} 宽松 {l}/{t}")


if __name__ == "__main__":
    main()