diff --git a/api_server.py b/api_server.py
index 39b62af..f718875 100644
--- a/api_server.py
+++ b/api_server.py
@@ -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 ... payload in the
+ # streamed text. This API already appends a canonical {"sql":...}
+ # at the end for the frontend to parse, so we filter any streamed block
+ # to avoid duplicate SQL / duplicate showing up client-side.
+ in_data_block = False
+ carry = ""
+ carry_len = 12 # enough to cover '' and '' 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("", 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("")
+ continue
+ end = buf.find("", i)
+ if end == -1:
+ # still inside ... 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("")
+
+ out_text = "".join(out_parts)
+ if out_text:
+ # Keep a small tail to detect a split '' 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- 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 直接拼接展示(且保留 ),
+ # 如果这里再补发 {"sql":...},就会导致“同一次流里出现两段 SQL/两段可见内容”。
+ # 因此前端不做过滤时,流式接口默认不再补发 结构化片段。
+ #
+ # 若后续需要为“仅消费结构化 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"{json.dumps({'sql': sql}, ensure_ascii=False)}"
+ 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}}