2026-04-14 10:28:22 +08:00
|
|
|
|
#!/usr/bin/env python3
|
|
|
|
|
|
"""
|
|
|
|
|
|
FastAPI 服务入口 - Text2SQL NL Chat API
|
|
|
|
|
|
包装 Text2SQLOrchestrator 为 REST API 服务,供前端调用
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
|
|
import os
|
|
|
|
|
|
import sys
|
|
|
|
|
|
import json
|
|
|
|
|
|
import html as html_lib
|
|
|
|
|
|
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
|
|
|
|
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
|
|
|
|
from fastapi.responses import StreamingResponse, JSONResponse
|
|
|
|
|
|
from pydantic import BaseModel, 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 main 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
|
2026-04-14 18:02:12 +08:00
|
|
|
|
from utils.dialog_context import (
|
|
|
|
|
|
last_assistant_was_data_query,
|
|
|
|
|
|
messages_to_text2sql_context,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
logging.basicConfig(
|
|
|
|
|
|
level=logging.INFO,
|
|
|
|
|
|
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
|
|
|
|
|
|
datefmt='%Y-%m-%d %H:%M:%S'
|
|
|
|
|
|
)
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
# 加载 .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}")
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
orchestrator = None
|
|
|
|
|
|
schema_manager = None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def get_orchestrator():
|
|
|
|
|
|
"""获取或初始化 orchestrator"""
|
|
|
|
|
|
global orchestrator, schema_manager
|
|
|
|
|
|
|
|
|
|
|
|
if orchestrator is not None:
|
|
|
|
|
|
return orchestrator
|
|
|
|
|
|
|
|
|
|
|
|
if not setup_environment():
|
2026-04-14 18:02:12 +08:00
|
|
|
|
raise RuntimeError(
|
|
|
|
|
|
"环境检查失败(请查看上方 WARNING/INFO,常见原因:"
|
|
|
|
|
|
"USE_LOCAL_EMBEDDING=true 但未配置或找不到 EMBEDDING_MODEL_PATH;"
|
|
|
|
|
|
"或 USE_LOCAL_EMBEDDING=false 但未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_*;"
|
|
|
|
|
|
"或 Schema 文件路径不对、DEEPSEEK_API_KEY 未设置)"
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
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:
|
|
|
|
|
|
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"))
|
|
|
|
|
|
embedding_model = None
|
|
|
|
|
|
vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma")
|
2026-04-14 18:21:50 +08:00
|
|
|
|
no_vector_search = False # 启用向量搜索(ChromaDB 已修复)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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):
|
|
|
|
|
|
message: str = Field(..., description="用户输入的自然语言问题")
|
|
|
|
|
|
service_code: Optional[str] = Field(None, description="模型路由: openai | deepseek")
|
|
|
|
|
|
lang_code: Optional[str] = Field("auto", description="语言: zh | en | tc | 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")
|
|
|
|
|
|
streaming_throttle: Optional[int] = Field(None, description="流式节流参数")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
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="是否允许导出")
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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交付说明"
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class NLChatSuccessData(BaseModel):
|
|
|
|
|
|
intent: IntentPayload
|
|
|
|
|
|
branch_result: DataQueryResult
|
|
|
|
|
|
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 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 = ""
|
|
|
|
|
|
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")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
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},
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
async def _load_session_text2sql_context(
|
|
|
|
|
|
request: NLChatRequest,
|
|
|
|
|
|
) -> 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)
|
|
|
|
|
|
block, _n = messages_to_text2sql_context(items)
|
|
|
|
|
|
return block, last_data
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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,
|
|
|
|
|
|
}
|
|
|
|
|
|
qo = result.metadata.get("question_original")
|
|
|
|
|
|
qz = result.metadata.get("question_zh_normalized")
|
|
|
|
|
|
if qo and qz:
|
|
|
|
|
|
payload["query_normalization"] = {"original": qo, "zh": qz}
|
|
|
|
|
|
return payload
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _sql_gen_stream_html(result: GenerationResult) -> str:
|
|
|
|
|
|
sql = (result.sql or "").strip()
|
|
|
|
|
|
inner = json.dumps({"sql": sql}, ensure_ascii=False)
|
|
|
|
|
|
parts = [f"<data>{inner}</data>"]
|
|
|
|
|
|
explain = ""
|
2026-04-14 18:02:12 +08:00
|
|
|
|
if result.metadata.get("sql_delivery_message"):
|
|
|
|
|
|
parts.append(
|
|
|
|
|
|
f'<div class="sql-delivery-body">{html_lib.escape(str(result.metadata["sql_delivery_message"]))}</div>'
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if result.metadata.get("db_empty_feedback"):
|
|
|
|
|
|
explain = str(result.metadata["db_empty_feedback"])
|
|
|
|
|
|
elif result.warnings:
|
|
|
|
|
|
explain = str(result.warnings[0])
|
|
|
|
|
|
elif result.errors:
|
|
|
|
|
|
explain = "; ".join(str(e) for e in result.errors[:5])
|
|
|
|
|
|
if explain.strip():
|
|
|
|
|
|
parts.append(f'<div class="sql-explain-body">{html_lib.escape(explain)}</div>')
|
|
|
|
|
|
return "".join(parts)
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
async def _run_generate(
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
top_k: int = 20,
|
|
|
|
|
|
dialog_context: Optional[str] = None,
|
|
|
|
|
|
) -> GenerationResult:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
orch = get_orchestrator()
|
|
|
|
|
|
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dc = (dialog_context or "").strip() or None
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
def _call() -> GenerationResult:
|
|
|
|
|
|
return orch.generate(
|
|
|
|
|
|
question=question.strip(),
|
|
|
|
|
|
dialect=dialect,
|
|
|
|
|
|
top_k_candidates=top_k,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context=dc,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
return await asyncio.to_thread(_call)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
|
|
|
|
|
|
if not request.message or not request.message.strip():
|
|
|
|
|
|
yield _sse_data({"code": 400, "msg": "消息内容不能为空", "data": None})
|
|
|
|
|
|
return
|
|
|
|
|
|
text = request.message.strip()
|
2026-04-14 18:02:12 +08:00
|
|
|
|
orch = get_orchestrator()
|
|
|
|
|
|
dialog_block, last_data = await _load_session_text2sql_context(request)
|
|
|
|
|
|
classified = await asyncio.to_thread(
|
|
|
|
|
|
classify_dialog,
|
|
|
|
|
|
text,
|
|
|
|
|
|
last_turn_was_data_query=last_data,
|
|
|
|
|
|
dialog_context=dialog_block or None,
|
|
|
|
|
|
llm_client=orch.deepseek,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if classified.intent == DialogIntent.CONVERSATION:
|
|
|
|
|
|
reply = classified.reply_suggestion or ""
|
|
|
|
|
|
logger.info("[API/stream] 对话意图: conversation(跳过 Text2SQL,与 CLI single_query 一致)")
|
|
|
|
|
|
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "CHAT"})
|
|
|
|
|
|
yield _sse_data({"stage": "chat", "stream_kind": "content", "content": reply})
|
|
|
|
|
|
data_dict = _conversation_nl_dict(reply)
|
|
|
|
|
|
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
|
|
|
|
|
|
await _append_session_if_needed(request, text, data_dict)
|
|
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
|
|
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "DATA_QUERY"})
|
|
|
|
|
|
try:
|
2026-04-14 18:02:12 +08:00
|
|
|
|
result = await _run_generate(text, dialog_context=dialog_block or None)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"[API/stream] 生成异常: {e}")
|
|
|
|
|
|
yield _sse_data({"code": 500, "msg": str(e), "data": None})
|
|
|
|
|
|
return
|
|
|
|
|
|
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": _sql_gen_stream_html(result)})
|
|
|
|
|
|
data_dict = _nl_dict_from_generation(result)
|
|
|
|
|
|
msg = "success" if result.valid else "partial"
|
|
|
|
|
|
yield _sse_data({"code": 200, "msg": msg, "data": data_dict})
|
|
|
|
|
|
await _append_session_if_needed(request, text, 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
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
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():
|
|
|
|
|
|
raise HTTPException(status_code=400, detail="消息内容不能为空")
|
|
|
|
|
|
text = request.message.strip()
|
|
|
|
|
|
logger.info(f"[API] 问题: {text[:100]}...")
|
2026-04-14 18:02:12 +08:00
|
|
|
|
orch = get_orchestrator()
|
|
|
|
|
|
dialog_block, last_data = await _load_session_text2sql_context(request)
|
|
|
|
|
|
classified = await asyncio.to_thread(
|
|
|
|
|
|
classify_dialog,
|
|
|
|
|
|
text,
|
|
|
|
|
|
last_turn_was_data_query=last_data,
|
|
|
|
|
|
dialog_context=dialog_block or None,
|
|
|
|
|
|
llm_client=orch.deepseek,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if classified.intent == DialogIntent.CONVERSATION:
|
|
|
|
|
|
reply = classified.reply_suggestion or ""
|
|
|
|
|
|
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
|
|
|
|
|
|
data_dict = _conversation_nl_dict(reply)
|
|
|
|
|
|
await _append_session_if_needed(request, text, data_dict)
|
|
|
|
|
|
return {"code": 200, "msg": "success", "data": data_dict}
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
result = await _run_generate(text, dialog_context=dialog_block or None)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
data_dict = _nl_dict_from_generation(result)
|
|
|
|
|
|
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)。
|
|
|
|
|
|
事件体为 JSON:分片 delta `{stage, stream_kind, content}` 或结束包 `{code, msg, data}`。
|
|
|
|
|
|
"""
|
|
|
|
|
|
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(_body: SqlExecuteBody):
|
|
|
|
|
|
return {
|
|
|
|
|
|
"code": 501,
|
|
|
|
|
|
"msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。",
|
|
|
|
|
|
"data": None,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@app.post("/g3sb/api/nl/operation-logs")
|
|
|
|
|
|
async def nl_post_operation_log(_body: Dict[str, Any] = Body(...)):
|
|
|
|
|
|
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}}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@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(
|
|
|
|
|
|
doc_type: Optional[str] = None,
|
|
|
|
|
|
limit: int = 50,
|
|
|
|
|
|
offset: int = 0,
|
|
|
|
|
|
):
|
|
|
|
|
|
_ = doc_type
|
|
|
|
|
|
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
|
2026-04-14 18:21:50 +08:00
|
|
|
|
import multiprocessing
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
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")
|
|
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
# PyInstaller + multiprocessing(spawn)兼容
|
|
|
|
|
|
multiprocessing.freeze_support()
|
|
|
|
|
|
|
|
|
|
|
|
# PyInstaller 打包后必须使用 app 对象,不能使用字符串模块名
|
2026-04-14 10:28:22 +08:00
|
|
|
|
uvicorn.run(
|
2026-04-14 18:21:50 +08:00
|
|
|
|
app,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
host=host,
|
|
|
|
|
|
port=port,
|
|
|
|
|
|
reload=False,
|
|
|
|
|
|
log_level="info"
|
|
|
|
|
|
)
|