Enhance SQL streaming and response handling in api_server.py by introducing new optional fields in NLChatSuccessData for SQL and explanation. Refactor _chat_stream_events to support multiple streaming shapes and improve response structure. Remove deprecated SqlExecuteBody class and streamline SSE data generation for better clarity and performance.

This commit is contained in:
陈辅元
2026-04-17 09:07:51 +08:00
parent 4ee0ad2eec
commit 161b380e3a
+237 -35
View File
@@ -15,9 +15,10 @@ from pathlib import Path
from typing import Optional, List, Dict, Any, AsyncIterator from typing import Optional, List, Dict, Any, AsyncIterator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException, Body, Path as FPath from fastapi import FastAPI, HTTPException, Body, Path as FPath, Request
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse, JSONResponse from fastapi.responses import StreamingResponse, JSONResponse
from fastapi.exceptions import RequestValidationError
from pydantic import AliasChoices, BaseModel, ConfigDict, Field from pydantic import AliasChoices, BaseModel, ConfigDict, Field
from dotenv import load_dotenv from dotenv import load_dotenv
@@ -146,7 +147,6 @@ class NLChatRequest(BaseModel):
validation_alias=AliasChoices("lang_code", "langCode"), validation_alias=AliasChoices("lang_code", "langCode"),
description="语言:zh=简体中文,tc=繁体中文,en=英语;auto=自动", description="语言:zh=简体中文,tc=繁体中文,en=英语;auto=自动",
) )
taskId: Optional[str] = Field(None, description="任务ID")
session_id: Optional[str] = Field(None, description="会话ID") session_id: Optional[str] = Field(None, description="会话ID")
visitor_biz_id: Optional[str] = Field(None, description="访客业务ID") visitor_biz_id: Optional[str] = Field(None, description="访客业务ID")
user_id: Optional[str] = Field(None, description="用户ID") user_id: Optional[str] = Field(None, description="用户ID")
@@ -185,6 +185,11 @@ class DataQueryResult(BaseModel):
class NLChatSuccessData(BaseModel): class NLChatSuccessData(BaseModel):
intent: IntentPayload intent: IntentPayload
branch_result: DataQueryResult branch_result: DataQueryResult
sql: Optional[str] = Field(None, description="便捷字段:等价于 branch_result.sql(可包含 -- 注释)")
explanation: Optional[str] = Field(
None,
description="SQL 自然语言解释/交付说明(便捷字段;通常来自 sql_delivery_message 或 sql_explain)",
)
stream_narrative: Optional[str] = Field(None, description="流式叙述(SQL块上方说明)") stream_narrative: Optional[str] = Field(None, description="流式叙述(SQL块上方说明)")
stream_narrative_after: Optional[str] = Field(None, description="流式叙述(SQL块下方说明)") stream_narrative_after: Optional[str] = Field(None, description="流式叙述(SQL块下方说明)")
@@ -219,15 +224,6 @@ class SessionMessagePatchBody(BaseModel):
visitor_biz_id: Optional[str] = None visitor_biz_id: Optional[str] = None
class SqlExecuteBody(BaseModel):
sql: str
max_rows: Optional[int] = None
user_id: Optional[str] = None
visitor_biz_id: Optional[str] = None
chat_session_id: Optional[str] = None
chat_message_id: Optional[int] = None
class FavoriteCreateBody(BaseModel): class FavoriteCreateBody(BaseModel):
fav_type: str = Field(..., description="sql | function | report") fav_type: str = Field(..., description="sql | function | report")
name: str = "" name: str = ""
@@ -248,6 +244,25 @@ def _sse_data(obj: Dict[str, Any]) -> bytes:
# 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加 # 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加
_DEFAULT_SSE_CHUNK_CHARS = int(os.getenv("SSE_STREAM_CHUNK_CHARS", "64")) _DEFAULT_SSE_CHUNK_CHARS = int(os.getenv("SSE_STREAM_CHUNK_CHARS", "64"))
# 流式输出形态:
# - dual(默认):同时发送 UI 所需的 {stage, stream_kind, content} delta,
# 并同步发送 ApiEnvelope(data 内携带逐步增长的字段),兼顾 chatStore 打字效果与“data 内流式”。
# - envelope:仅 ApiEnvelope(适合非本仓库前端/自定义消费端;本仓库 Vue 侧会看起来不“打字”)。
# - delta:仅 delta + 末尾 envelope(旧行为)
_SSE_STREAM_SHAPE = (os.getenv("SSE_STREAM_SHAPE", "dual") or "dual").strip().lower()
def _sse_stream_shape_dual() -> bool:
return _SSE_STREAM_SHAPE in ("dual", "both", "1", "true", "yes", "on", "default")
def _sse_stream_shape_envelope_only() -> bool:
return _SSE_STREAM_SHAPE in ("envelope", "env", "data-only", "data_only")
def _sse_stream_shape_delta_only() -> bool:
return _SSE_STREAM_SHAPE in ("delta", "legacy", "0", "false", "no", "off")
async def _sse_stream_text_chunks( async def _sse_stream_text_chunks(
stage: str, stage: str,
@@ -608,6 +623,12 @@ def _nl_dict_from_generation(result: GenerationResult) -> Dict[str, Any]:
"intent": {"intent": "DATA_QUERY", "confidence": conf, "reason": reason}, "intent": {"intent": "DATA_QUERY", "confidence": conf, "reason": reason},
"branch_result": branch_result, "branch_result": branch_result,
} }
payload["sql"] = sql
payload["explanation"] = (
str(sdm).strip()
if sdm
else (str(explain).strip() if explain else None)
)
qo = result.metadata.get("question_original") qo = result.metadata.get("question_original")
qz = result.metadata.get("question_zh_normalized") qz = result.metadata.get("question_zh_normalized")
if qo and qz: if qo and qz:
@@ -656,12 +677,57 @@ async def _run_generate(
return await _call_with_orch(o) return await _call_with_orch(o)
async def _stream_business_manual_answer_in_data(reply: str) -> AsyncIterator[bytes]:
"""将寒暄等回复以 ApiEnvelope 流式输出:data.branch_result.answer 逐步增长。"""
r = reply or ""
size = max(8, _DEFAULT_SSE_CHUNK_CHARS)
acc = ""
for i in range(0, len(r), size):
acc += r[i : i + size]
yield _sse_data(
{
"code": 200,
"msg": "streaming",
"data": {
"intent": {"intent": "BUSINESS_MANUAL", "confidence": 1.0, "reason": "conversation"},
"branch_result": {"answer": acc},
},
}
)
await asyncio.sleep(0)
async def _yield_data_query_stream_narrative_envelope(narrative: str) -> AsyncIterator[bytes]:
"""发送一条 DATA_QUERY 的中间 envelope:模型输出进入 data.stream_narrative(SQL 仍为空)。"""
yield _sse_data(
{
"code": 200,
"msg": "streaming",
"data": {
"intent": {"intent": "DATA_QUERY", "confidence": 0.0, "reason": "generating"},
"branch_result": {
"sql": "",
"columns": [],
"rows": [],
"row_count": 0,
"truncated": False,
},
"stream_narrative": narrative,
},
}
)
await asyncio.sleep(0)
async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]: async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
if not request.message or not request.message.strip(): if not request.message or not request.message.strip():
lang = _normalize_lang_code(request.lang_code) lang = _normalize_lang_code(request.lang_code)
yield _sse_data( msg = _localized_empty_input_reply(lang)
{"stage": "sql_gen", "stream_kind": "content", "content": _localized_empty_input_reply(lang)} if _sse_stream_shape_envelope_only():
) yield _sse_data({"code": 400, "msg": msg, "data": None})
else:
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": msg})
yield _sse_data({"code": 400, "msg": msg, "data": None})
return return
text = request.message.strip() text = request.message.strip()
orch = get_orchestrator() orch = get_orchestrator()
@@ -698,11 +764,34 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
len(reply), len(reply),
reply[:500] + ("…" if len(reply) > 500 else ""), reply[:500] + ("…" if len(reply) > 500 else ""),
) )
if _sse_stream_shape_envelope_only():
async for pkt in _stream_business_manual_answer_in_data(reply):
yield pkt
elif _sse_stream_shape_dual():
size = max(8, _DEFAULT_SSE_CHUNK_CHARS)
acc = ""
for i in range(0, len(reply or ""), size):
piece = (reply or "")[i : i + size]
acc += piece
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": piece})
yield _sse_data(
{
"code": 200,
"msg": "streaming",
"data": {
"intent": {"intent": "BUSINESS_MANUAL", "confidence": 1.0, "reason": "conversation"},
"branch_result": {"answer": acc},
},
}
)
await asyncio.sleep(0)
else:
async for pkt in _sse_stream_text_chunks("sql_gen", reply): async for pkt in _sse_stream_text_chunks("sql_gen", reply):
yield pkt yield pkt
data_dict = _conversation_nl_dict(reply) data_dict = _conversation_nl_dict(reply)
data_dict["llm"] = _effective_llm_route_from_request(request) data_dict["llm"] = _effective_llm_route_from_request(request)
await _append_session_if_needed(request, text, data_dict) await _append_session_if_needed(request, text, data_dict)
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
return return
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver")) dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
@@ -721,6 +810,13 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
async with _maybe_override_orch_llm(orch, request) as o: async with _maybe_override_orch_llm(orch, request) as o:
chunk_queue: asyncio.Queue[Optional[str]] = asyncio.Queue() chunk_queue: asyncio.Queue[Optional[str]] = asyncio.Queue()
holder: Dict[str, Any] = {} holder: Dict[str, Any] = {}
# Some LLMs (or prompts) may emit a structured <data>...</data> payload in the
# streamed text. This API already appends a canonical <data>{"sql":...}</data>
# at the end for the frontend to parse, so we filter any streamed <data> block
# to avoid duplicate SQL / duplicate <data> showing up client-side.
in_data_block = False
carry = ""
carry_len = 12 # enough to cover '<data>' and '</data>' splits
def _run_generate_sync() -> None: def _run_generate_sync() -> None:
try: try:
@@ -739,26 +835,91 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
loop.call_soon_threadsafe(chunk_queue.put_nowait, None) loop.call_soon_threadsafe(chunk_queue.put_nowait, None)
gen_task = asyncio.create_task(asyncio.to_thread(_run_generate_sync)) gen_task = asyncio.create_task(asyncio.to_thread(_run_generate_sync))
stream_acc = ""
while True: while True:
piece = await chunk_queue.get() piece = await chunk_queue.get()
if piece is None: if piece is None:
break break
# 不做二次切分:直接按 LLM 原生 delta 逐条透传给前端 # 不做二次切分:直接按 LLM 原生 delta 逐条透传给前端
if piece: if piece:
buf = carry + piece
carry = ""
out_parts: List[str] = []
i = 0
while i < len(buf):
if not in_data_block:
start = buf.find("<data>", i)
if start == -1:
out_parts.append(buf[i:])
break
if start > i:
out_parts.append(buf[i:start])
in_data_block = True
i = start + len("<data>")
continue
end = buf.find("</data>", i)
if end == -1:
# still inside <data>... keep a small tail to handle boundary splits
if len(buf) - i > carry_len:
carry = buf[-carry_len:]
else:
carry = buf[i:]
buf = ""
break
in_data_block = False
i = end + len("</data>")
out_text = "".join(out_parts)
if out_text:
# Keep a small tail to detect a split '<data>' across chunk boundaries.
if not in_data_block and len(out_text) > carry_len:
chunk = out_text[:-carry_len]
carry = out_text[-carry_len:] + carry
elif not in_data_block and not carry:
chunk = out_text
else:
chunk = ""
if chunk:
if _sse_stream_shape_envelope_only():
stream_acc += chunk
async for pkt in _yield_data_query_stream_narrative_envelope(stream_acc):
yield pkt
elif _sse_stream_shape_dual():
stream_acc += chunk
yield _sse_data( yield _sse_data(
{ {"stage": "sql_gen", "stream_kind": "content", "content": chunk}
"stage": "sql_gen",
"stream_kind": "content",
"content": piece,
}
) )
async for pkt in _yield_data_query_stream_narrative_envelope(stream_acc):
yield pkt
else:
yield _sse_data(
{"stage": "sql_gen", "stream_kind": "content", "content": chunk}
)
# Flush any buffered non-<data> tail once the stream ends.
if carry and not in_data_block:
if _sse_stream_shape_envelope_only():
stream_acc += carry
async for pkt in _yield_data_query_stream_narrative_envelope(stream_acc):
yield pkt
elif _sse_stream_shape_dual():
stream_acc += carry
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": carry})
async for pkt in _yield_data_query_stream_narrative_envelope(stream_acc):
yield pkt
else:
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": carry})
carry = ""
await gen_task await gen_task
if holder.get("error"): if holder.get("error"):
raise holder["error"] raise holder["error"]
result = holder["result"] result = holder["result"]
except Exception as e: except Exception as e:
logger.error(f"[API/stream] 生成异常: {e}", exc_info=True) logger.error(f"[API/stream] 生成异常: {e}", exc_info=True)
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": str(e)}) err = str(e) or "stream generate error"
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": err})
# 结束包:供前端稳定拿到错误 envelope(不参与展示)
yield _sse_data({"code": 500, "msg": err, "data": None})
return return
sql_out = (result.sql or "").strip() sql_out = (result.sql or "").strip()
@@ -793,9 +954,28 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
data_dict["llm"] = _effective_llm_route_from_request(request) data_dict["llm"] = _effective_llm_route_from_request(request)
await _append_session_if_needed(request, text, data_dict) await _append_session_if_needed(request, text, data_dict)
# NOTE:
# 前端当前会把所有 delta 的 content 直接拼接展示(且保留 <data>),
# 如果这里再补发 <data>{"sql":...}</data>,就会导致“同一次流里出现两段 SQL/两段可见内容”。
# 因此前端不做过滤时,流式接口默认不再补发 <data> 结构化片段。
#
# 若后续需要为“仅消费结构化 SQL”的客户端启用,可通过环境变量开关恢复该能力。
if os.getenv("SSE_APPEND_SQL_DATA_TAG", "").strip().lower() in {"1", "true", "yes", "on"}:
try:
sql = (data_dict.get("sql") or "").strip()
data_tag = f"<data>{json.dumps({'sql': sql}, ensure_ascii=False)}</data>"
async for pkt in _sse_stream_text_chunks("sql_gen", f"\n{data_tag}\n"):
yield pkt
except Exception:
# 不影响主流程:仅为前端解析增强
pass
# 结束包:前端解析到 ApiEnvelope 后返回最终结构(该 JSON 不含 stage/stream_kind/content,不会展示到 UI)
yield _sse_data({"code": 200, "msg": "success" if result.valid else "partial", "data": data_dict})
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(_app: FastAPI):
logger.info("=" * 60) logger.info("=" * 60)
logger.info("Text2SQL API Server 启动中...") logger.info("Text2SQL API Server 启动中...")
logger.info("=" * 60) logger.info("=" * 60)
@@ -819,6 +999,31 @@ app = FastAPI(
lifespan=lifespan lifespan=lifespan
) )
def _envelope_error(code: int, msg: str, data: Any = None) -> Dict[str, Any]:
return {"code": int(code), "msg": msg or "", "data": data}
@app.exception_handler(HTTPException)
async def _http_exception_handler(_: Request, exc: HTTPException):
# 前端统一按 ApiEnvelope 解析;这里用 HTTP 200 + code 字段表达业务错误
msg = exc.detail if isinstance(exc.detail, str) else str(exc.detail)
return JSONResponse(status_code=200, content=_envelope_error(exc.status_code, msg, None))
@app.exception_handler(RequestValidationError)
async def _validation_exception_handler(_: Request, exc: RequestValidationError):
# Pydantic 校验失败时也包装成 ApiEnvelope,避免前端收到 {"detail":[...]} 结构
return JSONResponse(
status_code=200,
content=_envelope_error(422, "request validation error", {"detail": exc.errors()}),
)
@app.exception_handler(Exception)
async def _unhandled_exception_handler(_: Request, exc: Exception):
logger.error("[API] unhandled exception: %s", exc, exc_info=True)
return JSONResponse(status_code=200, content=_envelope_error(500, str(exc) or "internal error", None))
app.add_middleware( app.add_middleware(
CORSMiddleware, CORSMiddleware,
allow_origins=["*"], allow_origins=["*"],
@@ -916,8 +1121,15 @@ async def nl_chat(request: NLChatRequest):
async def nl_chat_stream(request: NLChatRequest): async def nl_chat_stream(request: NLChatRequest):
""" """
自然语言对话流式接口(SSE)。 自然语言对话流式接口(SSE)。
事件体为 JSON:仅发送分片 `data: {"stage":"sql_gen","stream_kind":"content","content":"..."}`,
流结束表示完成;不再发送额外的结束包。 默认(`SSE_STREAM_SHAPE=dual`)同时发送两类事件(适配本仓库前端 chatStore 的增量渲染):
- UI 增量:`data: {"stage":"sql_gen","stream_kind":"content","content":"..."}`(用于打字/拼接 sqlGenHtml)
- 结构化增量:`data: {"code":200,"msg":"streaming","data":{...}}`(`DATA_QUERY` 时增长 `data.stream_narrative`)
- 结束:再发送最终 `{"code":200,"msg":"success|partial","data":{...}}`(含真实 SQL 等)
其它模式:
- `SSE_STREAM_SHAPE=envelope`:仅 ApiEnvelope(自定义消费端可用;本仓库 UI 可能不显示“打字”)
- `SSE_STREAM_SHAPE=delta`:仅 delta + 末尾 envelope(旧行为)
""" """
return StreamingResponse( return StreamingResponse(
_chat_stream_events(request), _chat_stream_events(request),
@@ -998,7 +1210,7 @@ async def nl_delete_session(
@app.post("/g3sb/api/nl/sql/execute") @app.post("/g3sb/api/nl/sql/execute")
async def nl_sql_execute(_body: SqlExecuteBody): async def nl_sql_execute():
return { return {
"code": 501, "code": 501,
"msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。", "msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。",
@@ -1007,23 +1219,15 @@ async def nl_sql_execute(_body: SqlExecuteBody):
@app.post("/g3sb/api/nl/operation-logs") @app.post("/g3sb/api/nl/operation-logs")
async def nl_post_operation_log(_body: Dict[str, Any] = Body(...)): async def nl_post_operation_log():
return {"code": 200, "msg": "success", "data": None} return {"code": 200, "msg": "success", "data": None}
@app.get("/g3sb/api/nl/operation-logs") @app.get("/g3sb/api/nl/operation-logs")
async def nl_list_operation_logs( async def nl_list_operation_logs(
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
limit: int = 50, limit: int = 50,
offset: int = 0, offset: int = 0,
q: Optional[str] = None,
op_type: Optional[str] = None,
result: Optional[str] = None,
date_from: Optional[str] = None,
date_to: Optional[str] = None,
): ):
_ = (user_id, visitor_biz_id, q, op_type, result, date_from, date_to)
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}} return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}
@@ -1070,11 +1274,9 @@ async def nl_delete_favorite(
@app.get("/g3sb/api/nl/knowledge-docs/uploads") @app.get("/g3sb/api/nl/knowledge-docs/uploads")
async def nl_list_knowledge_uploads( async def nl_list_knowledge_uploads(
doc_type: Optional[str] = None,
limit: int = 50, limit: int = 50,
offset: int = 0, offset: int = 0,
): ):
_ = doc_type
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}} return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}