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:
+241
-39
@@ -15,9 +15,10 @@ from pathlib import Path
|
||||
from typing import Optional, List, Dict, Any, AsyncIterator
|
||||
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.responses import StreamingResponse, JSONResponse
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field
|
||||
from dotenv import load_dotenv
|
||||
|
||||
@@ -146,7 +147,6 @@ class NLChatRequest(BaseModel):
|
||||
validation_alias=AliasChoices("lang_code", "langCode"),
|
||||
description="语言:zh=简体中文,tc=繁体中文,en=英语;auto=自动",
|
||||
)
|
||||
taskId: Optional[str] = Field(None, description="任务ID")
|
||||
session_id: Optional[str] = Field(None, description="会话ID")
|
||||
visitor_biz_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):
|
||||
intent: IntentPayload
|
||||
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_after: Optional[str] = Field(None, description="流式叙述(SQL块下方说明)")
|
||||
|
||||
@@ -219,15 +224,6 @@ class SessionMessagePatchBody(BaseModel):
|
||||
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):
|
||||
fav_type: str = Field(..., description="sql | function | report")
|
||||
name: str = ""
|
||||
@@ -248,6 +244,25 @@ def _sse_data(obj: Dict[str, Any]) -> bytes:
|
||||
# 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加
|
||||
_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(
|
||||
stage: str,
|
||||
@@ -608,6 +623,12 @@ def _nl_dict_from_generation(result: GenerationResult) -> Dict[str, Any]:
|
||||
"intent": {"intent": "DATA_QUERY", "confidence": conf, "reason": reason},
|
||||
"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")
|
||||
qz = result.metadata.get("question_zh_normalized")
|
||||
if qo and qz:
|
||||
@@ -656,12 +677,57 @@ async def _run_generate(
|
||||
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]:
|
||||
if not request.message or not request.message.strip():
|
||||
lang = _normalize_lang_code(request.lang_code)
|
||||
yield _sse_data(
|
||||
{"stage": "sql_gen", "stream_kind": "content", "content": _localized_empty_input_reply(lang)}
|
||||
)
|
||||
msg = _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
|
||||
text = request.message.strip()
|
||||
orch = get_orchestrator()
|
||||
@@ -698,11 +764,34 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
|
||||
len(reply),
|
||||
reply[:500] + ("…" if len(reply) > 500 else ""),
|
||||
)
|
||||
async for pkt in _sse_stream_text_chunks("sql_gen", reply):
|
||||
yield pkt
|
||||
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):
|
||||
yield pkt
|
||||
data_dict = _conversation_nl_dict(reply)
|
||||
data_dict["llm"] = _effective_llm_route_from_request(request)
|
||||
await _append_session_if_needed(request, text, data_dict)
|
||||
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
|
||||
return
|
||||
|
||||
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:
|
||||
chunk_queue: asyncio.Queue[Optional[str]] = asyncio.Queue()
|
||||
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:
|
||||
try:
|
||||
@@ -739,26 +835,91 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
|
||||
loop.call_soon_threadsafe(chunk_queue.put_nowait, None)
|
||||
|
||||
gen_task = asyncio.create_task(asyncio.to_thread(_run_generate_sync))
|
||||
stream_acc = ""
|
||||
while True:
|
||||
piece = await chunk_queue.get()
|
||||
if piece is None:
|
||||
break
|
||||
# 不做二次切分:直接按 LLM 原生 delta 逐条透传给前端
|
||||
if piece:
|
||||
yield _sse_data(
|
||||
{
|
||||
"stage": "sql_gen",
|
||||
"stream_kind": "content",
|
||||
"content": 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(
|
||||
{"stage": "sql_gen", "stream_kind": "content", "content": chunk}
|
||||
)
|
||||
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
|
||||
if holder.get("error"):
|
||||
raise holder["error"]
|
||||
result = holder["result"]
|
||||
except Exception as e:
|
||||
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
|
||||
|
||||
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)
|
||||
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
|
||||
async def lifespan(app: FastAPI):
|
||||
async def lifespan(_app: FastAPI):
|
||||
logger.info("=" * 60)
|
||||
logger.info("Text2SQL API Server 启动中...")
|
||||
logger.info("=" * 60)
|
||||
@@ -819,6 +999,31 @@ app = FastAPI(
|
||||
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(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
@@ -916,8 +1121,15 @@ async def nl_chat(request: NLChatRequest):
|
||||
async def nl_chat_stream(request: NLChatRequest):
|
||||
"""
|
||||
自然语言对话流式接口(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(
|
||||
_chat_stream_events(request),
|
||||
@@ -998,7 +1210,7 @@ async def nl_delete_session(
|
||||
|
||||
|
||||
@app.post("/g3sb/api/nl/sql/execute")
|
||||
async def nl_sql_execute(_body: SqlExecuteBody):
|
||||
async def nl_sql_execute():
|
||||
return {
|
||||
"code": 501,
|
||||
"msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。",
|
||||
@@ -1007,23 +1219,15 @@ async def nl_sql_execute(_body: SqlExecuteBody):
|
||||
|
||||
|
||||
@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}
|
||||
|
||||
|
||||
@app.get("/g3sb/api/nl/operation-logs")
|
||||
async def nl_list_operation_logs(
|
||||
user_id: Optional[str] = None,
|
||||
visitor_biz_id: Optional[str] = None,
|
||||
limit: int = 50,
|
||||
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}}
|
||||
|
||||
|
||||
@@ -1070,11 +1274,9 @@ async def nl_delete_favorite(
|
||||
|
||||
@app.get("/g3sb/api/nl/knowledge-docs/uploads")
|
||||
async def nl_list_knowledge_uploads(
|
||||
doc_type: Optional[str] = None,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
):
|
||||
_ = doc_type
|
||||
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user