1315 lines
49 KiB
Python
1315 lines
49 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
FastAPI 服务入口 - Text2SQL NL Chat API
|
||
包装 Text2SQLOrchestrator 为 REST API 服务,供前端调用
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import sys
|
||
import json
|
||
import asyncio
|
||
import logging
|
||
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, 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
|
||
|
||
_REPO_DIR = Path(__file__).resolve().parent
|
||
_BACKEND_DIR = _REPO_DIR / "backend"
|
||
|
||
if str(_BACKEND_DIR) not in sys.path:
|
||
sys.path.insert(0, str(_BACKEND_DIR))
|
||
|
||
from utils.repo_logging import configure_text2sql_api_logging
|
||
|
||
configure_text2sql_api_logging(_REPO_DIR)
|
||
|
||
from bootstrap import setup_environment, load_schema, create_orchestrator, resolve_sql_dialect
|
||
from agents.orchestrator import GenerationResult
|
||
from nl_lite_store import lite_nl_store
|
||
from utils.dialog_classifier import DialogIntent, classify_dialog
|
||
from utils.dialog_context import (
|
||
last_assistant_was_data_query,
|
||
messages_to_text2sql_context,
|
||
is_likely_follow_up,
|
||
)
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# per-request 覆盖 LLM client 时,使用锁避免并发串改 orchestrator.deepseek
|
||
_ORCH_LLM_LOCK = asyncio.Lock()
|
||
|
||
# 加载 .env 文件(支持 PyInstaller 打包后的目录结构)
|
||
def _find_env_file() -> Path:
|
||
"""查找 .env 文件,支持多种运行环境"""
|
||
import sys
|
||
|
||
# 尝试1: 当前工作目录
|
||
cwd_env = Path.cwd() / ".env"
|
||
if cwd_env.exists():
|
||
return cwd_env
|
||
|
||
# 尝试2: PyInstaller 临时目录(单文件模式)
|
||
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
|
||
meipass_env = Path(sys._MEIPASS) / ".env"
|
||
if meipass_env.exists():
|
||
return meipass_env
|
||
|
||
# 尝试3: 脚本/可执行文件所在目录
|
||
if getattr(sys, 'frozen', False):
|
||
# PyInstaller 打包后
|
||
exe_dir = Path(sys.executable).parent
|
||
env_path = exe_dir / ".env"
|
||
if env_path.exists():
|
||
return env_path
|
||
else:
|
||
# 开发环境
|
||
script_dir = Path(__file__).resolve().parent
|
||
env_path = script_dir / ".env"
|
||
if env_path.exists():
|
||
return env_path
|
||
|
||
# 默认返回当前目录
|
||
return cwd_env
|
||
|
||
env_file = _find_env_file()
|
||
if env_file.exists():
|
||
load_dotenv(env_file)
|
||
logger.info(f"[OK] 已加载配置文件: {env_file}")
|
||
else:
|
||
logger.warning(f"[WARN] 未找到 .env 文件: {env_file}")
|
||
|
||
orchestrator = None
|
||
schema_manager = None
|
||
|
||
|
||
def get_orchestrator():
|
||
"""获取或初始化 orchestrator"""
|
||
global orchestrator, schema_manager
|
||
|
||
if orchestrator is not None:
|
||
return orchestrator
|
||
|
||
if not setup_environment():
|
||
raise RuntimeError(
|
||
"环境检查失败(请查看上方 WARNING/INFO,常见原因:"
|
||
"未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding;"
|
||
"或 Schema 文件路径不对、LLM Key 未设置)"
|
||
)
|
||
|
||
schema_path = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
|
||
schema_meta_path = os.getenv("SCHEMA_META_PATH", None)
|
||
|
||
schema_manager = load_schema(schema_path, schema_meta_path)
|
||
|
||
class Args:
|
||
# LLM 路由由 LLM_SERVICE_CODE + 对应 Key 决定;此处仅保留历史字段以兼容 create_orchestrator 签名
|
||
api_key = os.getenv("DEEPSEEK_API_KEY")
|
||
model = os.getenv("MODEL_PRIMARY", "deepseek-chat")
|
||
temperature = float(os.getenv("TEMPERATURE", "0.3"))
|
||
max_tokens = int(os.getenv("MAX_TOKENS", "4096"))
|
||
max_retry = int(os.getenv("MAX_RETRY", "2"))
|
||
vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma")
|
||
no_vector_search = False # 启用向量搜索(ChromaDB 已修复)
|
||
no_fewshot = not os.getenv("FEWSHOT_ENABLED", "true").lower() in ("true", "1", "yes")
|
||
fewshot_top_k = int(os.getenv("FEWSHOT_TOP_K", "3"))
|
||
fewshot_min_rating = int(os.getenv("FEWSHOT_MIN_RATING", "7"))
|
||
no_translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() in (
|
||
"0",
|
||
"false",
|
||
"no",
|
||
"off",
|
||
)
|
||
|
||
orchestrator = create_orchestrator(schema_manager, Args())
|
||
logger.info("[OK] Orchestrator 初始化完成")
|
||
|
||
return orchestrator
|
||
|
||
|
||
class NLChatRequest(BaseModel):
|
||
model_config = ConfigDict(populate_by_name=True)
|
||
|
||
message: str = Field(..., description="用户输入的自然语言问题")
|
||
service_code: Optional[str] = Field(None, description="模型路由: openai | deepseek")
|
||
model: Optional[str] = Field(None, description="模型名称(可选,覆盖默认模型)")
|
||
# 与前端约定:zh=简体中文,tc=繁体中文,en=英语;可选 auto。JSON 可同时使用 lang_code 或 langCode。
|
||
lang_code: Optional[str] = Field(
|
||
"auto",
|
||
validation_alias=AliasChoices("lang_code", "langCode"),
|
||
description="语言:zh=简体中文,tc=繁体中文,en=英语;auto=自动",
|
||
)
|
||
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")
|
||
|
||
|
||
class IntentPayload(BaseModel):
|
||
intent: str
|
||
confidence: Optional[float] = None
|
||
reason: Optional[str] = None
|
||
|
||
|
||
class DataQueryResult(BaseModel):
|
||
sql: str = Field(..., description="生成的SQL语句")
|
||
columns: List[str] = Field(default_factory=list, description="查询结果列名")
|
||
rows: List[Dict[str, Any]] = Field(default_factory=list, description="查询结果行数据")
|
||
row_count: int = Field(0, description="结果行数")
|
||
truncated: bool = Field(False, description="是否截断")
|
||
sql_explain: Optional[str] = Field(None, description="SQL自然语言说明")
|
||
can_export: Optional[bool] = Field(True, description="是否允许导出")
|
||
db_execution_status: Optional[int] = Field(
|
||
None,
|
||
description="库探针:1=有数据行,0=无行需追问补充,-1=执行失败(重试),None=未探针",
|
||
)
|
||
db_empty_feedback: Optional[str] = Field(
|
||
None, description="探针0时的无行说明与追问(含问题分析)"
|
||
)
|
||
follow_up_required: bool = Field(
|
||
False,
|
||
description="True 表示探针为0:需用户补充条件后重新提问以重新生成SQL",
|
||
)
|
||
sql_delivery_message: Optional[str] = Field(
|
||
None, description="探针为1时LLM生成的面向用户SQL交付说明"
|
||
)
|
||
|
||
|
||
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块下方说明)")
|
||
|
||
|
||
class ApiEnvelope(BaseModel):
|
||
code: int = Field(200, description="状态码,200表示成功")
|
||
msg: str = Field("", description="消息")
|
||
data: Optional[NLChatSuccessData] = Field(None, description="响应数据")
|
||
|
||
|
||
class ErrorResponse(BaseModel):
|
||
code: int
|
||
msg: str
|
||
data: Optional[Any] = None
|
||
|
||
|
||
class SessionCreateBody(BaseModel):
|
||
title: Optional[str] = None
|
||
user_id: Optional[str] = None
|
||
visitor_biz_id: Optional[str] = None
|
||
|
||
|
||
class SessionTitlePatchBody(BaseModel):
|
||
title: Optional[str] = None
|
||
user_id: Optional[str] = None
|
||
visitor_biz_id: Optional[str] = None
|
||
|
||
|
||
class SessionMessagePatchBody(BaseModel):
|
||
content: str
|
||
user_id: Optional[str] = None
|
||
visitor_biz_id: Optional[str] = None
|
||
|
||
|
||
class FavoriteCreateBody(BaseModel):
|
||
fav_type: str = Field(..., description="sql | function | report")
|
||
name: str = ""
|
||
desc: Optional[str] = None
|
||
user_id: Optional[str] = None
|
||
visitor_biz_id: Optional[str] = None
|
||
sql: Optional[str] = None
|
||
sql_explain: Optional[str] = None
|
||
path: Optional[str] = None
|
||
reportPath: Optional[str] = None
|
||
params: Optional[str] = None
|
||
|
||
|
||
def _sse_data(obj: Dict[str, Any]) -> bytes:
|
||
return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n".encode("utf-8")
|
||
|
||
|
||
# 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加
|
||
_DEFAULT_SSE_CHUNK_CHARS = int(os.getenv("SSE_STREAM_CHUNK_CHARS", "64"))
|
||
|
||
# 流式输出形态:
|
||
# - dual(默认):发送 {stage, stream_kind, content} delta(与 delta 一致);DATA_QUERY 生成阶段不再
|
||
# 每条 delta 重复一份完整 ApiEnvelope,避免与前端拼接重复、体积 O(n²)。需要“仅 envelope 增长”时用 envelope。
|
||
# - envelope:仅 ApiEnvelope,data.stream_narrative 逐步增长(无 stage delta;适合只解析 envelope 的客户端)。
|
||
# - delta:仅 delta + 末尾 envelope(与 dual 在 SQL 流式段落行为一致)。
|
||
_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,
|
||
content: str,
|
||
*,
|
||
chunk_size: Optional[int] = None,
|
||
) -> AsyncIterator[bytes]:
|
||
"""将长文本拆成多段 SSE,便于浏览器逐段渲染(流式)。"""
|
||
if not content:
|
||
return
|
||
size = max(8, chunk_size or _DEFAULT_SSE_CHUNK_CHARS)
|
||
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)
|
||
|
||
|
||
def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
|
||
"""寒暄/元问题等:与前端 BUSINESS_MANUAL + branch_result.answer 一致。"""
|
||
r = (reply or "").strip() or "您好。"
|
||
return {
|
||
"intent": {"intent": "BUSINESS_MANUAL", "confidence": 1.0, "reason": "conversation"},
|
||
"branch_result": {"answer": r},
|
||
}
|
||
|
||
|
||
def _effective_llm_route_from_request(request: NLChatRequest) -> Dict[str, Any]:
|
||
"""
|
||
计算“本次请求实际使用的 LLM 路由信息”(用于前端联调回显)。
|
||
注意:此处仅回显路由结果,不包含密钥等敏感信息。
|
||
"""
|
||
model = _normalize_model_label(request.model)
|
||
sc = _infer_service_code(request.service_code, model)
|
||
|
||
sc2 = sc
|
||
if sc2 == "deepseek" and not _has_deepseek_key():
|
||
if _has_openai_key():
|
||
sc2 = "openai"
|
||
if sc2 == "openai" and not _has_openai_key():
|
||
if _has_deepseek_key():
|
||
sc2 = "deepseek"
|
||
|
||
# 与 _maybe_override_orch_llm 保持一致:网关不兼容时丢弃 model 覆盖,让下游走默认模型
|
||
effective_model = model
|
||
if effective_model is not None:
|
||
if (sc2 or "").strip().lower() == "openai" and "deepseek" in effective_model.strip().lower():
|
||
effective_model = None
|
||
elif (sc2 or "").strip().lower() == "deepseek" and effective_model.strip().lower().startswith("gpt"):
|
||
effective_model = None
|
||
|
||
return {
|
||
"service_code": sc2,
|
||
"model": effective_model,
|
||
}
|
||
|
||
|
||
def _normalize_lang_code(code: Optional[str]) -> str:
|
||
raw = (code or "auto").strip()
|
||
if not raw:
|
||
return "auto"
|
||
# 兼容中文取值
|
||
if raw in ("简体中文", "简体", "中文(简体)"):
|
||
return "zh"
|
||
if raw in ("繁体中文", "繁體中文", "繁体", "繁體"):
|
||
return "tc"
|
||
if raw in ("英语", "英文", "英語"):
|
||
return "en"
|
||
|
||
c = raw.lower().replace("_", "-")
|
||
if c in ("zh", "tc", "en", "auto"):
|
||
return c
|
||
|
||
# 兼容前端可能传的语言标签
|
||
if c in ("english", "en-us", "en-gb"):
|
||
return "en"
|
||
if c in ("zh-cn", "zh-hans", "zh-sg"):
|
||
return "zh"
|
||
if c in ("zh-tw", "zh-hant", "zh-hk", "traditional", "traditional-chinese"):
|
||
return "tc"
|
||
return "auto"
|
||
|
||
|
||
def _normalize_model_label(model: Optional[str]) -> Optional[str]:
|
||
"""
|
||
前端可能传展示文案(如 'GPT-4o mini' / 'DeepSeek V3')。
|
||
这里统一映射到真实模型名(openai/deepseek 各自可用的 model id)。
|
||
"""
|
||
raw = (model or "").strip()
|
||
if not raw:
|
||
return None
|
||
m = raw.strip().lower()
|
||
m = m.replace("_", "-").replace(" ", "-")
|
||
while "--" in m:
|
||
m = m.replace("--", "-")
|
||
|
||
# OpenAI
|
||
if m in ("gpt-4o-mini", "gpt4o-mini", "gpt-4o-mini"):
|
||
return "gpt-4o-mini"
|
||
if m in ("gpt-4o", "gpt4o"):
|
||
return "gpt-4o"
|
||
|
||
# DeepSeek
|
||
if m in ("deepseek-v3", "deepseekv3", "deepseek-v3.0", "deepseekv3.0", "deepseek-v3"):
|
||
return "deepseek-chat"
|
||
if m in ("deepseek-chat", "deepseek-reasoner"):
|
||
return m
|
||
|
||
return raw # 未知时原样透传(交由网关/后端判定)
|
||
|
||
|
||
def _infer_service_code(service_code: Optional[str], model_name: Optional[str]) -> Optional[str]:
|
||
sc = (service_code or "").strip().lower()
|
||
if sc:
|
||
if sc in ("openai", "deepseek"):
|
||
return sc
|
||
if sc in ("gpt", "chatgpt"):
|
||
return "openai"
|
||
if "deepseek" in sc:
|
||
return "deepseek"
|
||
|
||
m = (model_name or "").strip().lower()
|
||
if m:
|
||
if "deepseek" in m:
|
||
return "deepseek"
|
||
if m.startswith("gpt"):
|
||
return "openai"
|
||
return None
|
||
|
||
|
||
def _has_openai_key() -> bool:
|
||
return bool((os.getenv("OPENAI_API_KEY") or "").strip())
|
||
|
||
|
||
def _has_deepseek_key() -> bool:
|
||
return bool((os.getenv("DEEPSEEK_API_KEY") or "").strip())
|
||
|
||
|
||
def _localized_conversation_reply(lang_code: str) -> str:
|
||
if lang_code == "en":
|
||
return (
|
||
"Hello, I'm the Text2SQL assistant.\n"
|
||
"Describe what you want to query or aggregate in natural language "
|
||
"(e.g., available balance of an account; summarize unsettled trades by broker).\n"
|
||
"Type quit or exit to leave."
|
||
)
|
||
if lang_code == "tc":
|
||
return (
|
||
"您好,我是業務庫 Text2SQL 助手。\n"
|
||
"請用自然語言描述要查詢或統計的內容(例如:查詢某帳戶可用餘額、按經紀商匯總未結算交易筆數)。\n"
|
||
"輸入 quit 或 exit 可退出。"
|
||
)
|
||
# zh / auto
|
||
return (
|
||
"您好,我是业务库 Text2SQL 助手。\n"
|
||
"请用自然语言描述要查询或统计的内容(例如:查询某账户可用余额、按经纪商汇总未结算交易笔数)。\n"
|
||
"输入 quit 或 exit 可退出。"
|
||
)
|
||
|
||
|
||
def _localized_empty_input_reply(lang_code: str) -> str:
|
||
if lang_code == "en":
|
||
return "Please enter a concrete business query, or type quit to exit."
|
||
if lang_code == "tc":
|
||
return "請輸入具體的業務查詢問題,或輸入 quit 退出。"
|
||
return "请输入具体的业务查询问题,或输入 quit 退出。"
|
||
|
||
|
||
def _lang_label(lang_code: str) -> str:
|
||
if lang_code == "en":
|
||
return "English"
|
||
if lang_code == "tc":
|
||
return "Traditional Chinese"
|
||
return "Simplified Chinese"
|
||
|
||
|
||
async def _translate_explain_text(
|
||
orch,
|
||
text: Optional[str],
|
||
lang_code: str,
|
||
) -> Optional[str]:
|
||
"""
|
||
将“说明类”文本翻译到 lang_code(仅用于 sql_delivery_message / db_empty_feedback / sql_explain)。
|
||
zh 直接返回;en/tc 使用 LLM 翻译(短输出,避免引入格式)。
|
||
"""
|
||
t = (text or "").strip()
|
||
if not t:
|
||
return text
|
||
if lang_code in ("auto", "zh"):
|
||
return text
|
||
target = _lang_label(lang_code)
|
||
|
||
messages = [
|
||
{
|
||
"role": "system",
|
||
"content": (
|
||
"You are a translation assistant.\n"
|
||
f"Translate the following text to {target}.\n"
|
||
"Rules:\n"
|
||
"- Keep the meaning identical; do not add new information.\n"
|
||
"- Keep SQL keywords/code unchanged if present.\n"
|
||
"- Output plain text only (no Markdown, no code fences, no JSON).\n"
|
||
),
|
||
},
|
||
{"role": "user", "content": t},
|
||
]
|
||
|
||
def _call() -> str:
|
||
msg = orch.deepseek.chat(messages, temperature=0.0, top_p=1.0, max_completion_tokens=320)
|
||
return (msg.content or "").strip()
|
||
|
||
try:
|
||
out = await asyncio.to_thread(_call)
|
||
return out or text
|
||
except Exception as e:
|
||
logger.warning("[API] explain translation skipped: %s", e)
|
||
return text
|
||
|
||
|
||
@asynccontextmanager
|
||
async def _maybe_override_orch_llm(orch, request: NLChatRequest):
|
||
"""
|
||
若请求传 service_code / model,则临时覆盖 orchestrator.deepseek;
|
||
使用锁保证同一时刻仅一个请求修改该引用。
|
||
"""
|
||
model = _normalize_model_label(request.model)
|
||
sc = _infer_service_code(request.service_code, model)
|
||
lang = _normalize_lang_code(request.lang_code)
|
||
needs_override = (sc is not None) or (model is not None) or (lang == "en")
|
||
if not needs_override:
|
||
yield orch
|
||
return
|
||
|
||
from llm.router import create_llm_client
|
||
|
||
tmp = None
|
||
if sc is not None or model is not None:
|
||
# 兜底:前端切什么就用什么;但若对应 key 未配置,则自动回退到另一家,保证可用
|
||
sc2 = sc
|
||
if sc2 == "deepseek" and not _has_deepseek_key():
|
||
if _has_openai_key():
|
||
logger.warning("[API] deepseek requested but key missing, fallback to openai")
|
||
sc2 = "openai"
|
||
if sc2 == "openai" and not _has_openai_key():
|
||
if _has_deepseek_key():
|
||
logger.warning("[API] openai requested but key missing, fallback to deepseek")
|
||
sc2 = "deepseek"
|
||
|
||
llm_kwargs: Dict[str, Any] = {}
|
||
if model is not None:
|
||
# 保护:openai 网关下 deepseek-chat 会 404;deepseek 官方下 gpt-* 也会失败
|
||
if (sc2 or "").strip().lower() == "openai" and "deepseek" in model.strip().lower():
|
||
pass
|
||
elif (sc2 or "").strip().lower() == "deepseek" and model.strip().lower().startswith("gpt"):
|
||
pass
|
||
else:
|
||
llm_kwargs["model_name"] = model
|
||
|
||
tmp = create_llm_client(sc2, **llm_kwargs)
|
||
|
||
async with _ORCH_LLM_LOCK:
|
||
old = getattr(orch, "deepseek", None)
|
||
old_translate = getattr(orch, "translate_english_to_zh", None)
|
||
if tmp is not None:
|
||
orch.deepseek = tmp
|
||
# 英文界面:不要把英文问句归一成中文,否则下游解释会倾向中文
|
||
if lang == "en" and hasattr(orch, "translate_english_to_zh"):
|
||
orch.translate_english_to_zh = False
|
||
try:
|
||
yield orch
|
||
finally:
|
||
if tmp is not None:
|
||
orch.deepseek = old
|
||
if old_translate is not None and hasattr(orch, "translate_english_to_zh"):
|
||
orch.translate_english_to_zh = old_translate
|
||
|
||
|
||
async def _load_session_text2sql_context(
|
||
request: NLChatRequest,
|
||
user_text: str,
|
||
) -> tuple[str, bool]:
|
||
"""
|
||
在写入本轮之前读取会话历史,构造 Text2SQL 上文,并判断上一轮助手是否为数据查询。
|
||
"""
|
||
sid = (request.session_id or "").strip()
|
||
if not sid:
|
||
return "", False
|
||
data = await lite_nl_store.get_messages(
|
||
request.user_id,
|
||
request.visitor_biz_id,
|
||
sid,
|
||
limit=200,
|
||
offset=0,
|
||
)
|
||
if not data:
|
||
return "", False
|
||
items = data.get("items") or []
|
||
if not items:
|
||
return "", False
|
||
last_data = last_assistant_was_data_query(items)
|
||
# 方案A:仅在“续问/沿用口径”时注入少量上文;新话题直接清空,避免上下文污染 SQL。
|
||
max_pairs = 2 if is_likely_follow_up(user_text) else 0
|
||
block, _n = messages_to_text2sql_context(items, max_pairs=max_pairs)
|
||
return block, last_data
|
||
|
||
|
||
async def _append_session_if_needed(request: NLChatRequest, user_text: str, data: Dict[str, Any]) -> None:
|
||
if not (request.session_id and request.session_id.strip()):
|
||
return
|
||
try:
|
||
assistant_json = json.dumps(data, ensure_ascii=False)
|
||
await lite_nl_store.append_exchange(
|
||
request.user_id,
|
||
request.visitor_biz_id,
|
||
request.session_id.strip(),
|
||
user_text,
|
||
assistant_json,
|
||
)
|
||
except Exception as e:
|
||
logger.warning(f"[API] 会话落库跳过: {e}")
|
||
|
||
|
||
def _nl_dict_from_generation(result: GenerationResult) -> Dict[str, Any]:
|
||
sql = (result.sql or "").strip()
|
||
explain_parts: List[str] = []
|
||
if result.metadata.get("db_empty_feedback"):
|
||
explain_parts.append(str(result.metadata["db_empty_feedback"]))
|
||
if result.warnings:
|
||
explain_parts.extend(str(w) for w in result.warnings)
|
||
if result.errors:
|
||
explain_parts.extend(str(e) for e in result.errors)
|
||
explain = "; ".join(explain_parts) if explain_parts else None
|
||
branch_result: Dict[str, Any] = {
|
||
"sql": sql,
|
||
"columns": [],
|
||
"rows": [],
|
||
"row_count": 0,
|
||
"truncated": False,
|
||
}
|
||
if explain:
|
||
branch_result["sql_explain"] = explain
|
||
dbs = result.metadata.get("db_execution_status")
|
||
if dbs is not None:
|
||
branch_result["db_execution_status"] = dbs
|
||
dbe = result.metadata.get("db_empty_feedback")
|
||
if dbe:
|
||
branch_result["db_empty_feedback"] = dbe
|
||
branch_result["follow_up_required"] = bool(result.valid and dbs == 0)
|
||
sdm = result.metadata.get("sql_delivery_message")
|
||
if sdm:
|
||
branch_result["sql_delivery_message"] = str(sdm)
|
||
conf = 1.0 if result.valid else 0.0
|
||
if result.valid:
|
||
reason = f"使用了 {len(result.tables_used)} 张表"
|
||
else:
|
||
reason = result.errors[0] if result.errors else "SQL生成未通过验证"
|
||
payload: 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:
|
||
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
|
||
|
||
|
||
async def _run_generate(
|
||
question: str,
|
||
top_k: int = 20,
|
||
dialog_context: Optional[str] = None,
|
||
request: Optional[NLChatRequest] = None,
|
||
) -> GenerationResult:
|
||
orch = get_orchestrator()
|
||
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
|
||
dc = (dialog_context or "").strip() or None
|
||
q = question.strip()
|
||
logger.info(
|
||
"[GEN/API] dialect=%s top_k=%s dialog_context_chars=%s question_len=%s preview=%r",
|
||
dialect,
|
||
top_k,
|
||
len(dc) if dc else 0,
|
||
len(q),
|
||
q[:300] + ("…" if len(q) > 300 else ""),
|
||
)
|
||
|
||
async def _call_with_orch(o) -> GenerationResult:
|
||
def _call() -> GenerationResult:
|
||
return o.generate(
|
||
question=question.strip(),
|
||
dialect=dialect,
|
||
top_k_candidates=top_k,
|
||
dialog_context=dc,
|
||
)
|
||
|
||
return await asyncio.to_thread(_call)
|
||
|
||
if request is None:
|
||
return await _call_with_orch(orch)
|
||
async with _maybe_override_orch_llm(orch, request) as 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]:
|
||
if not request.message or not request.message.strip():
|
||
lang = _normalize_lang_code(request.lang_code)
|
||
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()
|
||
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
||
logger.info(
|
||
"[API/stream] 开始: user_id=%r visitor_biz_id=%r session_id=%r service_code=%r model=%r "
|
||
"lang_code=%r msg_chars=%s preview=%r dialog_context_chars=%s last_turn_was_data_query=%s",
|
||
request.user_id,
|
||
request.visitor_biz_id,
|
||
request.session_id,
|
||
request.service_code,
|
||
request.model,
|
||
request.lang_code,
|
||
len(text),
|
||
text[:400] + ("…" if len(text) > 400 else ""),
|
||
len(dialog_block) if dialog_block else 0,
|
||
last_data,
|
||
)
|
||
lang = _normalize_lang_code(request.lang_code)
|
||
async with _maybe_override_orch_llm(orch, request) as o:
|
||
classified = await asyncio.to_thread(
|
||
classify_dialog,
|
||
text,
|
||
last_turn_was_data_query=last_data,
|
||
dialog_context=dialog_block or None,
|
||
llm_client=o.deepseek,
|
||
)
|
||
if classified.intent == DialogIntent.CONVERSATION:
|
||
reply = (classified.reply_suggestion or "").strip()
|
||
if not reply:
|
||
reply = _localized_conversation_reply(lang)
|
||
logger.info(
|
||
"[API/stream] 对话意图 conversation(跳过 Text2SQL): reply_chars=%s reply_preview=%r",
|
||
len(reply),
|
||
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):
|
||
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"))
|
||
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:
|
||
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:
|
||
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))
|
||
# 仅 envelope 模式需要累积全文;dual/delta 只发 stage 增量,结束包再带完整 data
|
||
stream_acc = ""
|
||
while True:
|
||
piece = await chunk_queue.get()
|
||
if piece is None:
|
||
break
|
||
# 不做二次切分:直接按 LLM 原生 delta 逐条透传给前端
|
||
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
|
||
else:
|
||
# dual / delta:仅 stage 增量,避免每条再套一层完整 ApiEnvelope
|
||
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
|
||
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)
|
||
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()
|
||
logger.info(
|
||
"[API/stream] Text2SQL 完成: valid=%s attempts=%s tables_used=%s sql_chars=%s sql_head=%r",
|
||
result.valid,
|
||
result.attempts,
|
||
result.tables_used,
|
||
len(sql_out),
|
||
sql_out[:500] + ("…" if len(sql_out) > 500 else ""),
|
||
)
|
||
if result.errors:
|
||
logger.warning("[API/stream] 错误列表: %s", result.errors)
|
||
if result.warnings:
|
||
logger.info("[API/stream] 警告列表: %s", result.warnings)
|
||
|
||
# What this SQL does / SQL 说明:按 lang_code 本地化
|
||
try:
|
||
if isinstance(result.metadata, dict):
|
||
if result.metadata.get("sql_delivery_message"):
|
||
result.metadata["sql_delivery_message"] = await _translate_explain_text(
|
||
orch, str(result.metadata["sql_delivery_message"]), lang
|
||
)
|
||
if result.metadata.get("db_empty_feedback"):
|
||
result.metadata["db_empty_feedback"] = await _translate_explain_text(
|
||
orch, str(result.metadata["db_empty_feedback"]), lang
|
||
)
|
||
except Exception:
|
||
pass
|
||
|
||
data_dict = _nl_dict_from_generation(result)
|
||
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):
|
||
logger.info("=" * 60)
|
||
logger.info("Text2SQL API Server 启动中...")
|
||
logger.info("=" * 60)
|
||
|
||
try:
|
||
get_orchestrator()
|
||
logger.info("[OK] 服务已就绪")
|
||
except Exception as e:
|
||
logger.error(f"服务初始化失败: {e}")
|
||
raise
|
||
|
||
yield
|
||
|
||
logger.info("Text2SQL API Server 已关闭")
|
||
|
||
|
||
app = FastAPI(
|
||
title="Text2SQL NL Chat API",
|
||
description="自然语言转SQL的对话接口",
|
||
version="1.0.0",
|
||
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=["*"],
|
||
allow_credentials=True,
|
||
allow_methods=["*"],
|
||
allow_headers=["*"],
|
||
)
|
||
|
||
|
||
@app.get("/health")
|
||
async def health_check():
|
||
"""健康检查接口"""
|
||
return {"status": "ok", "message": "Text2SQL API is running"}
|
||
|
||
|
||
@app.post("/g3sb/api/nl/chat")
|
||
async def nl_chat(request: NLChatRequest):
|
||
"""
|
||
自然语言对话接口(非流式),响应结构与流式结束包一致。
|
||
寒暄/致谢等先经 classify_dialog,与 CLI 一致。
|
||
"""
|
||
try:
|
||
if not request.message or not request.message.strip():
|
||
lang = _normalize_lang_code(request.lang_code)
|
||
raise HTTPException(status_code=400, detail=_localized_empty_input_reply(lang))
|
||
text = request.message.strip()
|
||
orch = get_orchestrator()
|
||
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
||
logger.info(
|
||
"[API/chat] 请求: user_id=%r session_id=%r service_code=%r model=%r lang=%r "
|
||
"msg_chars=%s preview=%r dialog_context_chars=%s last_turn_was_data_query=%s",
|
||
request.user_id,
|
||
request.session_id,
|
||
request.service_code,
|
||
request.model,
|
||
request.lang_code,
|
||
len(text),
|
||
text[:400] + ("…" if len(text) > 400 else ""),
|
||
len(dialog_block) if dialog_block else 0,
|
||
last_data,
|
||
)
|
||
lang = _normalize_lang_code(request.lang_code)
|
||
async with _maybe_override_orch_llm(orch, request) as o:
|
||
classified = await asyncio.to_thread(
|
||
classify_dialog,
|
||
text,
|
||
last_turn_was_data_query=last_data,
|
||
dialog_context=dialog_block or None,
|
||
llm_client=o.deepseek,
|
||
)
|
||
if classified.intent == DialogIntent.CONVERSATION:
|
||
reply = (classified.reply_suggestion or "").strip()
|
||
if not reply:
|
||
reply = _localized_conversation_reply(lang)
|
||
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
|
||
data_dict = _conversation_nl_dict(reply)
|
||
data_dict["llm"] = _effective_llm_route_from_request(request)
|
||
await _append_session_if_needed(request, text, data_dict)
|
||
return {"code": 200, "msg": "success", "data": data_dict}
|
||
|
||
result = await _run_generate(text, dialog_context=dialog_block or None, request=request)
|
||
|
||
# What this SQL does / SQL 说明:按 lang_code 本地化
|
||
try:
|
||
if isinstance(result.metadata, dict):
|
||
if result.metadata.get("sql_delivery_message"):
|
||
result.metadata["sql_delivery_message"] = await _translate_explain_text(
|
||
orch, str(result.metadata["sql_delivery_message"]), lang
|
||
)
|
||
if result.metadata.get("db_empty_feedback"):
|
||
result.metadata["db_empty_feedback"] = await _translate_explain_text(
|
||
orch, str(result.metadata["db_empty_feedback"]), lang
|
||
)
|
||
except Exception:
|
||
pass
|
||
|
||
data_dict = _nl_dict_from_generation(result)
|
||
data_dict["llm"] = _effective_llm_route_from_request(request)
|
||
response_data = NLChatSuccessData.model_validate(data_dict)
|
||
msg = "success" if result.valid else "partial"
|
||
if result.valid:
|
||
logger.info(f"[API] 生成成功: {(result.sql or '')[:80]}...")
|
||
else:
|
||
logger.warning(f"[API] 生成失败: {result.errors}")
|
||
await _append_session_if_needed(request, text, data_dict)
|
||
return ApiEnvelope(code=200, msg=msg, data=response_data)
|
||
except HTTPException:
|
||
raise
|
||
except Exception as e:
|
||
logger.error(f"[API] 异常: {e}")
|
||
raise HTTPException(status_code=500, detail=str(e))
|
||
|
||
|
||
@app.post("/g3sb/api/nl/chat/stream")
|
||
async def nl_chat_stream(request: NLChatRequest):
|
||
"""
|
||
自然语言对话流式接口(SSE)。
|
||
|
||
默认(`SSE_STREAM_SHAPE=dual`)在 DATA_QUERY SQL 生成阶段只推送 stage delta(与 `delta` 一致);
|
||
寒暄等仍会按 dual 发送 delta + 增长的 `data.branch_result.answer`。结束包为完整
|
||
`{"code":200,"msg":"success|partial","data":{...}}`。
|
||
|
||
其它模式:
|
||
- `SSE_STREAM_SHAPE=envelope`:无 stage delta,`data.stream_narrative` 在 envelope 内逐步增长
|
||
- `SSE_STREAM_SHAPE=delta`:与当前 dual 在 SQL 流式段行为相同
|
||
"""
|
||
return StreamingResponse(
|
||
_chat_stream_events(request),
|
||
media_type="text/event-stream",
|
||
headers={
|
||
"Cache-Control": "no-cache",
|
||
"Connection": "keep-alive",
|
||
"X-Accel-Buffering": "no",
|
||
},
|
||
)
|
||
|
||
|
||
@app.post("/g3sb/api/nl/sessions")
|
||
async def nl_create_session(body: SessionCreateBody):
|
||
row = await lite_nl_store.create_session(body.user_id, body.visitor_biz_id, body.title)
|
||
return {"code": 200, "msg": "success", "data": row}
|
||
|
||
|
||
@app.get("/g3sb/api/nl/sessions")
|
||
async def nl_list_sessions(
|
||
user_id: Optional[str] = None,
|
||
visitor_biz_id: Optional[str] = None,
|
||
limit: int = 100,
|
||
offset: int = 0,
|
||
):
|
||
data = await lite_nl_store.list_sessions(user_id, visitor_biz_id, limit, offset)
|
||
return {"code": 200, "msg": "success", "data": data}
|
||
|
||
|
||
@app.get("/g3sb/api/nl/sessions/{session_id}/messages")
|
||
async def nl_get_messages(
|
||
session_id: str = FPath(..., description="会话ID"),
|
||
user_id: Optional[str] = None,
|
||
visitor_biz_id: Optional[str] = None,
|
||
limit: int = 500,
|
||
offset: int = 0,
|
||
):
|
||
data = await lite_nl_store.get_messages(user_id, visitor_biz_id, session_id, limit, offset)
|
||
if data is None:
|
||
return JSONResponse(status_code=200, content={"code": 404, "msg": "session not found", "data": None})
|
||
return {"code": 200, "msg": "success", "data": data}
|
||
|
||
|
||
@app.patch("/g3sb/api/nl/sessions/{session_id}/title")
|
||
async def nl_patch_session_title(session_id: str, body: SessionTitlePatchBody):
|
||
row = await lite_nl_store.update_session_title(
|
||
body.user_id, body.visitor_biz_id, session_id, body.title
|
||
)
|
||
if not row:
|
||
return JSONResponse(status_code=200, content={"code": 404, "msg": "session not found", "data": None})
|
||
return {"code": 200, "msg": "success", "data": row}
|
||
|
||
|
||
@app.patch("/g3sb/api/nl/sessions/{session_id}/messages/{message_id}")
|
||
async def nl_patch_session_message(
|
||
session_id: str,
|
||
message_id: int,
|
||
body: SessionMessagePatchBody,
|
||
):
|
||
row = await lite_nl_store.patch_message(
|
||
body.user_id, body.visitor_biz_id, session_id, message_id, body.content
|
||
)
|
||
if not row:
|
||
return JSONResponse(status_code=200, content={"code": 404, "msg": "message not found", "data": None})
|
||
return {"code": 200, "msg": "success", "data": row}
|
||
|
||
|
||
@app.delete("/g3sb/api/nl/sessions/{session_id}")
|
||
async def nl_delete_session(
|
||
session_id: str,
|
||
user_id: Optional[str] = None,
|
||
visitor_biz_id: Optional[str] = None,
|
||
):
|
||
row = await lite_nl_store.delete_session(user_id, visitor_biz_id, session_id)
|
||
if not row:
|
||
return JSONResponse(status_code=200, content={"code": 404, "msg": "session not found", "data": None})
|
||
return {"code": 200, "msg": "success", "data": row}
|
||
|
||
|
||
@app.post("/g3sb/api/nl/sql/execute")
|
||
async def nl_sql_execute():
|
||
return {
|
||
"code": 501,
|
||
"msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。",
|
||
"data": None,
|
||
}
|
||
|
||
|
||
@app.post("/g3sb/api/nl/operation-logs")
|
||
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(
|
||
limit: int = 50,
|
||
offset: int = 0,
|
||
):
|
||
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}
|
||
|
||
|
||
@app.get("/g3sb/api/nl/favorites")
|
||
async def nl_list_favorites(
|
||
user_id: Optional[str] = None,
|
||
visitor_biz_id: Optional[str] = None,
|
||
):
|
||
grouped = await lite_nl_store.get_favorites_grouped(user_id, visitor_biz_id)
|
||
return {"code": 200, "msg": "success", "data": grouped}
|
||
|
||
|
||
@app.post("/g3sb/api/nl/favorites")
|
||
async def nl_create_favorite(body: FavoriteCreateBody):
|
||
row = await lite_nl_store.add_favorite(body.user_id, body.visitor_biz_id, body.model_dump(exclude_none=True))
|
||
return {"code": 200, "msg": "success", "data": row}
|
||
|
||
|
||
@app.patch("/g3sb/api/nl/favorites/{fav_id}")
|
||
async def nl_patch_favorite(
|
||
fav_id: str,
|
||
body: Dict[str, Any] = Body(default_factory=dict),
|
||
user_id: Optional[str] = None,
|
||
visitor_biz_id: Optional[str] = None,
|
||
):
|
||
patch = {k: v for k, v in body.items() if k not in ("user_id", "visitor_biz_id")}
|
||
ok = await lite_nl_store.patch_favorite(user_id, visitor_biz_id, fav_id, patch)
|
||
if not ok:
|
||
return JSONResponse(status_code=200, content={"code": 404, "msg": "favorite not found", "data": None})
|
||
return {"code": 200, "msg": "success", "data": {}}
|
||
|
||
|
||
@app.delete("/g3sb/api/nl/favorites/{fav_id}")
|
||
async def nl_delete_favorite(
|
||
fav_id: str,
|
||
user_id: Optional[str] = None,
|
||
visitor_biz_id: Optional[str] = None,
|
||
):
|
||
ok = await lite_nl_store.delete_favorite(user_id, visitor_biz_id, fav_id)
|
||
if not ok:
|
||
return JSONResponse(status_code=200, content={"code": 404, "msg": "favorite not found", "data": None})
|
||
return {"code": 200, "msg": "success", "data": {"ok": True}}
|
||
|
||
|
||
@app.get("/g3sb/api/nl/knowledge-docs/uploads")
|
||
async def nl_list_knowledge_uploads(
|
||
limit: int = 50,
|
||
offset: int = 0,
|
||
):
|
||
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}
|
||
|
||
|
||
@app.post("/g3sb/api/nl/knowledge-docs/upload")
|
||
async def nl_upload_knowledge_doc():
|
||
return JSONResponse(
|
||
status_code=200,
|
||
content={"code": 501, "msg": "knowledge upload not supported on lite API", "data": None},
|
||
)
|
||
|
||
|
||
@app.get("/g3sb/api/nl/admin/visibility")
|
||
async def admin_visibility():
|
||
"""管理端可见性接口"""
|
||
return {
|
||
"code": 200,
|
||
"msg": "success",
|
||
"data": {
|
||
"admin_visible": True
|
||
}
|
||
}
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import uvicorn
|
||
import multiprocessing
|
||
|
||
port = int(os.getenv("API_PORT", "8041"))
|
||
host = os.getenv("API_HOST", "0.0.0.0")
|
||
|
||
logger.info(f"启动服务: http://{host}:{port}")
|
||
logger.info(f"API文档: http://{host}:{port}/docs")
|
||
|
||
# PyInstaller + multiprocessing(spawn)兼容
|
||
multiprocessing.freeze_support()
|
||
|
||
# PyInstaller 打包后必须使用 app 对象,不能使用字符串模块名
|
||
uvicorn.run(
|
||
app,
|
||
host=host,
|
||
port=port,
|
||
reload=False,
|
||
log_level="info",
|
||
# 必须为 None:False 仍会进入 uvicorn 的 fileConfig 分支并崩溃;None 才跳过覆盖 root 日志
|
||
log_config=None,
|
||
)
|