"""OpenAI 兼容端点:/v1/models 与 /v1/chat/completions(SSE 流式 + 普通 JSON)。 生成是阻塞调用,放线程池执行;流式用 asyncio.Queue 把后端回调桥接成 SSE 生成器。 客户端断开时通过 threading.Event 让 backend 的 on_text 返回 True 中断生成。 """ from __future__ import annotations import asyncio import json import os import re import threading import time import uuid from fastapi import APIRouter, Request from fastapi.responses import JSONResponse, StreamingResponse from mlx_streaming import config from mlx_streaming.server.state import EngineManager def _sse(obj) -> str: return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n" def _delta_of(prev: str, cur: str) -> str: """后端上报的是累计全文,这里取出本次新增。 按真实公共前缀切,而不是 `cur[len(prev):]`:后者假设 prev 一定是 cur 的前缀,一旦后端 改写了已发出的尾部(历史上流式解码到半个汉字会先发 U+FFFD、下一步再换成真字符), 长度相同就切出空串,那个字符永久丢失。此处宁可重发一小段,也不丢字。 """ if cur.startswith(prev): return cur[len(prev):] n = 0 for a, b in zip(prev, cur): if a != b: break n += 1 return cur[n:] _MASK_RE = re.compile(r"\[(NAME|ID|TEAM|STATION|CAR|TRAIN):(emp_[a-z0-9]+)\]", re.I) def _mask_notice(content: str) -> str: """工具结果里带脱敏占位符时,就地追加一句原样透传的约束。 实测(2026-08-02):工具返回 `[TEAM:emp_ypsl004]`,模型答"乘务三组"——把占位符替换成 了自己编的名字,用户拿到错数据。系统提示里已有同样的规则,但它在 4000 token 之前, 而占位符就在眼前;贴着数据再说一次显著更硬,且模型无关。 只在真的检出占位符时加,不给普通工具结果增加 token。 """ if not content or not _MASK_RE.search(content): return content return (f"{content}\n\n[脱敏提示] 上面形如 [TEAM:emp_xxx]、[NAME:emp_xxx]、[ID:emp_xxx] 的是" "脱敏占位符,界面会自动还原成真实值。引用时**必须连方括号整体原样复制**," "严禁替换成「乘务一组」这类具体名称(你看不到真实值,写具体名称一定是编的)," "也不要拆掉方括号或向用户解释占位符。") def _normalize_messages(raw) -> list[dict]: """OpenAI messages → chat template 可用消息。 关键:assistant 的 tool_calls 与 role=tool 的消息**原样透传**——Qwen chat template 会把它们渲染成模型训练时学过的格式( 包裹)。 此前把工具返回改写成普通 user 文本,模型不认为是权威工具结果,会无视真实数据 自行编造报告(2026-07 实测复现)。只做 content 分段列表的拍平。 """ out = [] for m in raw: role = m.get("role", "user") content = m.get("content") if isinstance(content, list): # OpenAI 分段 content:[{"type":"text","text":...}] content = "".join( p.get("text", "") for p in content if isinstance(p, dict) and p.get("type") == "text") if role == "tool": # 原生透传:模板渲染为 ,模型按训练格式 ground 到工具数据 item = {"role": "tool", "content": _mask_notice(content or "")} if m.get("tool_call_id"): item["tool_call_id"] = m["tool_call_id"] if m.get("name"): item["name"] = m["name"] out.append(item) continue if role == "assistant" and m.get("tool_calls"): # 原生透传:模板渲染为 。arguments 必须是 JSON 字符串(OpenAI 规范) tcs = [] for tc in m["tool_calls"]: fn = tc.get("function", {}) args = fn.get("arguments") if args is not None and not isinstance(args, str): args = json.dumps(args, ensure_ascii=False) tcs.append({"id": tc.get("id"), "type": "function", "function": {"name": fn.get("name"), "arguments": args}}) out.append({"role": "assistant", "content": content or None, "tool_calls": tcs}) continue out.append({"role": role, "content": content or ""}) return out def _tool_name(t) -> str: if not isinstance(t, dict): return "" fn = t.get("function") if isinstance(fn, dict) and fn.get("name"): return str(fn["name"]) return str(t.get("name") or "") _tools_filter_seen: set = set() def _filter_tools(tools): """按 config.tools_allow() 白名单裁剪工具,全被裁掉时返回 None(本轮不暴露工具)。 见 config.tools_allow 的说明:这不只是省 token,主要是把 prompt 头部钉稳,让前缀快照 能命中——调用方的工具发现是动态的,工具集一变头部就变,每个新会话都要全量 prefill。 首次遇到某个「保留/丢弃」组合时打一行日志,便于确认线上到底发来了什么。 """ allow = config.tools_allow() if not allow or not tools: return tools allow_set = set(allow) kept, dropped = [], [] for t in tools: (kept if _tool_name(t) in allow_set else dropped).append(t) if dropped: sig = (tuple(_tool_name(t) for t in kept), tuple(_tool_name(t) for t in dropped)) if sig not in _tools_filter_seen: _tools_filter_seen.add(sig) print(f"[TOOLS_ALLOW] 保留 {len(kept)} 个 {list(sig[0])}," f"丢弃 {len(dropped)} 个 {list(sig[1])}。" f"目的是让 prompt 头部固定、前缀快照命中;设 TOOLS_ALLOW= 可关闭。", flush=True) return kept or None def make_openai_router(mgr: EngineManager) -> APIRouter: router = APIRouter() def _model_id() -> str: return os.path.basename(str(mgr.args.model).rstrip("/")) or "sparkle" def _auth(req: Request): """Bearer 任意非空 token 即放行;设了 SPARKLE_API_KEY 则必须匹配。""" auth = req.headers.get("authorization", "") tok = auth[7:].strip() if auth.lower().startswith("bearer ") else "" want = os.environ.get("SPARKLE_API_KEY", "") ok = (tok == want) if want else bool(tok) if ok: return None return JSONResponse(status_code=401, content={ "error": {"message": "unauthorized", "type": "authentication_error"}}) def _unavailable(): msg = "engine reloading" if mgr.reloading else "engine not loaded" return JSONResponse(status_code=503, content={"error": msg}) def _prompt_tokens(messages, tools=None) -> int: """prompt token 数:真后端用 tokenizer 精确编码;fake 退化为字符数估算。 必须带上 tools:工具定义是由 chat template 渲染进 prompt 的,实测 4 个工具就占 1349 token(占总 prompt 三分之一)。漏掉它上报的 usage 会显著偏小,而调用方 (中台)按 usage 判断要不要压缩上下文,少算就会算漏。 """ tok = getattr(mgr.backend, "_tok", None) if tok is not None: try: from mlx_streaming.cli import _encode_chat return len(_encode_chat(tok, messages, tools=tools)) except Exception: # noqa: BLE001 估算失败不阻塞生成 pass return sum(len(m.get("content") or "") for m in messages) def _usage(prompt_tokens: int, completion_tokens: int) -> dict: return {"prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, "total_tokens": prompt_tokens + completion_tokens} def _chunk(model_id, cmpl_id, created, delta=None, role=False, finish=None, usage=None): if usage is not None: # OpenAI 风格:usage chunk 的 choices 为空数组 return {"id": cmpl_id, "object": "chat.completion.chunk", "created": created, "model": model_id, "choices": [], "usage": usage} d = {} if role: d["role"] = "assistant" if delta: d["content"] = delta return {"id": cmpl_id, "object": "chat.completion.chunk", "created": created, "model": model_id, "choices": [{"index": 0, "delta": d, "finish_reason": finish}]} @router.get("/v1/models") async def list_models(req: Request): err = _auth(req) if err is not None: return err return {"object": "list", "data": [{ "id": _model_id(), "object": "model", "created": int(time.time()), "owned_by": "sparkle"}]} @router.post("/v1/chat/completions") async def chat_completions(req: Request): err = _auth(req) if err is not None: return err if not mgr.loaded or mgr.reloading: return _unavailable() body = await req.json() # chat_template_kwargs / enable_thinking:接受并忽略。 # tools:渲染进 chat template(Qwen 模板把工具定义注入提示词),模型据此输出 # 文本,客户端兜底解析——server 侧不做原生 tool parser。 raw_tools = body.get("tools") or None tools = _filter_tools(raw_tools) # 「模型没调工具就编数据」有两种成因:调用方压根没发工具,或发了而模型无视。 # 两者在 server 侧看不出区别(都只是一次普通请求),故每轮记一行收到的工具名。 print(f"[REQ] 收到工具 {len(raw_tools or [])} 个 " f"{[_tool_name(t) for t in (raw_tools or [])]} → 渲染 " f"{len(tools or [])} 个 {[_tool_name(t) for t in (tools or [])]}", flush=True) # 采样策略(关键):模型 generation_config 要求 do_sample(temp 0.7/top_p 0.8/top_k 20), # 贪心会诱发幻觉。默认走采样模式(质量优先,非投机);仅显式 temperature=0 走贪心+MTP。 req_temp = body.get("temperature") sampling = None if req_temp is None or float(req_temp) > 0: gc = getattr(mgr.backend, "_gen_cfg", None) or {} sampling = { "temperature": float(req_temp) if req_temp is not None else float(gc.get("temperature", 0.7)), "top_p": body.get("top_p") or gc.get("top_p", 0.8), "top_k": body.get("top_k") or gc.get("top_k", 20), } messages = _normalize_messages(body.get("messages") or []) stream = bool(body.get("stream")) include_usage = bool((body.get("stream_options") or {}).get("include_usage")) req_max = body.get("max_tokens") or body.get("max_completion_tokens") model_id = _model_id() cmpl_id = "chatcmpl-" + uuid.uuid4().hex[:24] created = int(time.time()) loop = asyncio.get_running_loop() disconnect = threading.Event() # 客户端断开 → on_text 返回 True 中断生成 async def _watch_disconnect(): while not disconnect.is_set(): if await req.is_disconnected(): disconnect.set() return await asyncio.sleep(0.2) if not stream: async with mgr.lock: # 引擎单对话:并发请求排队等锁 if not mgr.loaded or mgr.reloading: # 等锁期间可能开始重建 return _unavailable() mgr.apply_gen_params(req_max) prompt_tokens = _prompt_tokens(messages, tools) watcher = asyncio.create_task(_watch_disconnect()) # 客户端断开也中断,防僵尸生成堵锁 try: res = await loop.run_in_executor( None, lambda: mgr.backend.generate( messages, lambda _f, _n: disconnect.is_set(), tools=tools, sampling=sampling)) except Exception as e: # noqa: BLE001 return JSONResponse(status_code=500, content={"error": str(e)}) finally: disconnect.set() watcher.cancel() mgr.last_tok_per_s = res.tok_per_s return {"id": cmpl_id, "object": "chat.completion", "created": created, "model": model_id, "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": res.text}}], "usage": _usage(prompt_tokens, res.n_tokens)} async def event_stream(): async with mgr.lock: if not mgr.loaded or mgr.reloading: yield _sse({"error": "engine reloading"}) yield "data: [DONE]\n\n" return mgr.apply_gen_params(req_max) prompt_tokens = _prompt_tokens(messages, tools) q = asyncio.Queue() def on_text(full, _n): if disconnect.is_set(): return True loop.call_soon_threadsafe(q.put_nowait, full) return False def _run(): # 阻塞生成,放线程池;结束/异常都经由 queue 通知 try: res = mgr.backend.generate(messages, on_text, tools=tools, sampling=sampling) loop.call_soon_threadsafe(q.put_nowait, ("__done__", res)) except Exception as e: # noqa: BLE001 loop.call_soon_threadsafe(q.put_nowait, ("__error__", e)) watcher = asyncio.create_task(_watch_disconnect()) loop.run_in_executor(None, _run) prev = "" result = None yield _sse(_chunk(model_id, cmpl_id, created, role=True)) try: while True: item = await q.get() if isinstance(item, tuple): if item[0] == "__done__": result = item[1] break raise item[1] delta = _delta_of(prev, item) prev = item if delta: yield _sse(_chunk(model_id, cmpl_id, created, delta=delta)) finally: disconnect.set() # 客户端断开/生成结束都确保后端被中断 watcher.cancel() if result is not None: mgr.last_tok_per_s = result.tok_per_s yield _sse(_chunk(model_id, cmpl_id, created, finish="stop")) if include_usage: n = result.n_tokens if result is not None else 0 yield _sse(_chunk(model_id, cmpl_id, created, usage=_usage(prompt_tokens, n))) yield "data: [DONE]\n\n" return StreamingResponse(event_stream(), media_type="text/event-stream") return router