Enhance SQL generation and streaming capabilities in the API server. Introduce optional parameters for streaming throttle and SQL stream granularity in NLChatRequest. Implement new functions for iterating SQL generation content pieces and adjusting streaming behavior based on user-defined settings. Update prompts for few-shot SQL adaptation and improve logging for SQL generation processes. Refactor orchestrator methods to support streaming responses and integrate few-shot SQL conditions. Update impact analysis documentation to reflect these changes.

This commit is contained in:
陈辅元
2026-04-16 13:48:44 +08:00
parent 695356a496
commit bff5f85d60
16 changed files with 586 additions and 103 deletions
+102 -29
View File
@@ -9,7 +9,6 @@ from __future__ import annotations
import os
import sys
import json
import html as html_lib
import asyncio
import logging
from pathlib import Path
@@ -151,7 +150,15 @@ class NLChatRequest(BaseModel):
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")
streaming_throttle: Optional[int] = Field(None, description="流式节流参数")
streaming_throttle: Optional[int] = Field(
None,
description="流式节流:相邻 content 分片之间的间隔毫秒数(0/None 表示不延迟)",
)
sql_stream_granularity: Optional[str] = Field(
None,
validation_alias=AliasChoices("sql_stream_granularity", "sqlStreamGranularity"),
description="SQL 生成流式分片:delta | char;不传则使用环境变量 SSE_SQL_GEN_SPLIT(默认 char)",
)
class IntentPayload(BaseModel):
@@ -247,6 +254,23 @@ def _sse_data(obj: Dict[str, Any]) -> bytes:
return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n".encode("utf-8")
def _iter_sql_gen_content_pieces(text: str, mode: Optional[str] = None) -> List[str]:
"""
将 LLM 流式片段再拆成前端期望的多条 SSE(与 chatStore onDelta 累加一致)。
每条均为:{"stage": "sql_gen", "stream_kind": "content", "content": "..."}
mode / 环境变量 SSE_SQL_GEN_SPLIT:
- char(请求默认):按 Unicode 标量逐字符发送(与「用户」「问题」逐条 data 一致)
- delta:与上游 LLM 每次 delta 一致(块更大、事件更少)
"""
if not text:
return []
raw = (mode or os.getenv("SSE_SQL_GEN_SPLIT", "char") or "char").strip().lower()
if raw in ("delta", "none", "0", "false"):
return [text]
return list(text)
# 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加
_DEFAULT_SSE_CHUNK_CHARS = int(os.getenv("SSE_STREAM_CHUNK_CHARS", "64"))
@@ -256,16 +280,21 @@ async def _sse_stream_text_chunks(
content: str,
*,
chunk_size: Optional[int] = None,
throttle_ms: Optional[int] = None,
) -> AsyncIterator[bytes]:
"""将长文本拆成多段 SSE,便于浏览器逐段渲染(流式)。"""
if not content:
return
size = max(8, chunk_size or _DEFAULT_SSE_CHUNK_CHARS)
delay = (throttle_ms or 0) / 1000.0 if throttle_ms and throttle_ms > 0 else 0.0
for i in range(0, len(content), size):
yield _sse_data(
{"stage": stage, "stream_kind": "content", "content": content[i : i + size]}
)
await asyncio.sleep(0)
if delay:
await asyncio.sleep(delay)
else:
await asyncio.sleep(0)
def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
@@ -584,29 +613,14 @@ def _nl_dict_from_generation(result: GenerationResult) -> Dict[str, Any]:
qz = result.metadata.get("question_zh_normalized")
if qo and qz:
payload["query_normalization"] = {"original": qo, "zh": qz}
if result.metadata.get("fewshot_golden_reuse"):
payload["fewshot_golden"] = {
"qid": result.metadata.get("fewshot_golden_qid"),
"score": result.metadata.get("fewshot_golden_score"),
}
return payload
def _sql_gen_stream_html(result: GenerationResult) -> str:
sql = (result.sql or "").strip()
inner = json.dumps({"sql": sql}, ensure_ascii=False)
parts = [f"<data>{inner}</data>"]
explain = ""
if result.metadata.get("sql_delivery_message"):
parts.append(
f'<div class="sql-delivery-body">{html_lib.escape(str(result.metadata["sql_delivery_message"]))}</div>'
)
if result.metadata.get("db_empty_feedback"):
explain = str(result.metadata["db_empty_feedback"])
elif result.warnings:
explain = str(result.warnings[0])
elif result.errors:
explain = "; ".join(str(e) for e in result.errors[:5])
if explain.strip():
parts.append(f'<div class="sql-explain-body">{html_lib.escape(explain)}</div>')
return "".join(parts)
async def _run_generate(
question: str,
top_k: int = 20,
@@ -684,7 +698,9 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
reply[:500] + ("…" if len(reply) > 500 else ""),
)
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "CHAT"})
async for pkt in _sse_stream_text_chunks("chat", reply):
async for pkt in _sse_stream_text_chunks(
"chat", reply, throttle_ms=request.streaming_throttle
):
yield pkt
data_dict = _conversation_nl_dict(reply)
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
@@ -692,8 +708,65 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
return
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "DATA_QUERY"})
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
dc_ctx = (dialog_block or "").strip() or None
top_k = 20
loop = asyncio.get_running_loop()
logger.info(
"[GEN/API/stream] dialect=%s top_k=%s dialog_context_chars=%s question_len=%s preview=%r",
dialect,
top_k,
len(dc_ctx) if dc_ctx else 0,
len(text),
text[:300] + ("…" if len(text) > 300 else ""),
)
try:
result = await _run_generate(text, dialog_context=dialog_block or None, request=request)
async with _maybe_override_orch_llm(orch, request) as o:
chunk_queue: asyncio.Queue[Optional[str]] = asyncio.Queue()
holder: Dict[str, Any] = {}
def _run_generate_sync() -> None:
try:
holder["result"] = o.generate(
question=text.strip(),
dialect=dialect,
top_k_candidates=top_k,
dialog_context=dc_ctx,
sql_stream_callback=lambda c: loop.call_soon_threadsafe(
chunk_queue.put_nowait, c
),
)
except Exception as e:
holder["error"] = e
finally:
loop.call_soon_threadsafe(chunk_queue.put_nowait, None)
gen_task = asyncio.create_task(asyncio.to_thread(_run_generate_sync))
throttle_ms = request.streaming_throttle or 0
sql_chunk_delay = throttle_ms / 1000.0 if throttle_ms > 0 else 0.0
gran = (
(request.sql_stream_granularity or "").strip()
or os.getenv("SSE_SQL_GEN_SPLIT", "char")
or "char"
).lower()
while True:
piece = await chunk_queue.get()
if piece is None:
break
for frag in _iter_sql_gen_content_pieces(piece, mode=gran):
yield _sse_data(
{
"stage": "sql_gen",
"stream_kind": "content",
"content": frag,
}
)
if sql_chunk_delay:
await asyncio.sleep(sql_chunk_delay)
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({"code": 500, "msg": str(e), "data": None})
@@ -727,9 +800,6 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
except Exception:
pass
html = _sql_gen_stream_html(result)
async for pkt in _sse_stream_text_chunks("sql_gen", html):
yield pkt
data_dict = _nl_dict_from_generation(result)
msg = "success" if result.valid else "partial"
yield _sse_data({"code": 200, "msg": msg, "data": data_dict})
@@ -856,7 +926,10 @@ async def nl_chat(request: NLChatRequest):
async def nl_chat_stream(request: NLChatRequest):
"""
自然语言对话流式接口(SSE)。
事件体为 JSON:分片 delta `{stage, stream_kind, content}` 或结束包 `{code, msg, data}`。
事件体为 JSON:分片 `data: {"stage","stream_kind","content"}`(例:sql_gen 时
`stream_kind` 为 `content`)或结束包 `{code, msg, data}`。
可选:`sql_stream_granularity` / `sqlStreamGranularity`(delta|char,未传则 `SSE_SQL_GEN_SPLIT`,默认 char)、
`streaming_throttle`(相邻 content 分片间隔毫秒)。
"""
return StreamingResponse(
_chat_stream_events(request),