sparkle/mlx_streaming/server/openai.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

339 lines
15 KiB
Python

"""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
会把它们渲染成模型训练时学过的格式(<tool_call> 与 <tool_response> 包裹)。
此前把工具返回改写成普通 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":
# 原生透传:模板渲染为 <tool_response>,模型按训练格式 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"):
# 原生透传:模板渲染为 <tool_call>。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 模板把工具定义注入提示词),模型据此输出
# <tool_call> 文本,客户端兜底解析——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