Files

1315 lines
49 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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,
)