Files
ai-g3sb-backman2.0/api_server.py
T

704 lines
24 KiB
Python
Raw 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 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
from utils.dialog_context import (
last_assistant_was_data_query,
messages_to_text2sql_context,
)
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__)
# 加载 .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,常见原因:"
"USE_LOCAL_EMBEDDING=true 但未配置或找不到 EMBEDDING_MODEL_PATH;"
"或 USE_LOCAL_EMBEDDING=false 但未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_*;"
"或 Schema 文件路径不对、DEEPSEEK_API_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:
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")
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):
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="是否允许导出")
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
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},
}
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
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,
}
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 = ""
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>'
)
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)
async def _run_generate(
question: str,
top_k: int = 20,
dialog_context: Optional[str] = None,
) -> GenerationResult:
orch = get_orchestrator()
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
dc = (dialog_context or "").strip() or None
def _call() -> GenerationResult:
return orch.generate(
question=question.strip(),
dialect=dialect,
top_k_candidates=top_k,
dialog_context=dc,
)
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()
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,
)
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:
result = await _run_generate(text, dialog_context=dialog_block or None)
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]}...")
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,
)
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}
result = await _run_generate(text, dialog_context=dialog_block or None)
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
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"
)