339 lines
15 KiB
Python
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
|