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}}