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 asyncio
|
|
|
|
|
|
import logging
|
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
|
from typing import Optional, List, Dict, Any, AsyncIterator
|
|
|
|
|
|
from contextlib import asynccontextmanager
|
|
|
|
|
|
|
2026-04-17 09:07:51 +08:00
|
|
|
|
from fastapi import FastAPI, HTTPException, Body, Path as FPath, Request
|
2026-04-14 10:28:22 +08:00
|
|
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
|
|
|
|
from fastapi.responses import StreamingResponse, JSONResponse
|
2026-04-17 09:07:51 +08:00
|
|
|
|
from fastapi.exceptions import RequestValidationError
|
2026-04-16 09:15:01 +08:00
|
|
|
|
from pydantic import AliasChoices, BaseModel, ConfigDict, Field
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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))
|
|
|
|
|
|
|
2026-04-16 10:53:10 +08:00
|
|
|
|
from utils.repo_logging import configure_text2sql_api_logging
|
|
|
|
|
|
|
|
|
|
|
|
configure_text2sql_api_logging(_REPO_DIR)
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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-15 17:33:07 +08:00
|
|
|
|
is_likely_follow_up,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
# per-request 覆盖 LLM client 时,使用锁避免并发串改 orchestrator.deepseek
|
|
|
|
|
|
_ORCH_LLM_LOCK = asyncio.Lock()
|
|
|
|
|
|
|
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,常见原因:"
|
2026-04-15 09:49:18 +08:00
|
|
|
|
"未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding;"
|
2026-04-16 09:15:01 +08:00
|
|
|
|
"或 Schema 文件路径不对、LLM Key 未设置)"
|
2026-04-14 18:02:12 +08:00
|
|
|
|
)
|
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:
|
2026-04-16 09:15:01 +08:00
|
|
|
|
# LLM 路由由 LLM_SERVICE_CODE + 对应 Key 决定;此处仅保留历史字段以兼容 create_orchestrator 签名
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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")
|
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):
|
2026-04-16 09:15:01 +08:00
|
|
|
|
model_config = ConfigDict(populate_by_name=True)
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
message: str = Field(..., description="用户输入的自然语言问题")
|
|
|
|
|
|
service_code: Optional[str] = Field(None, description="模型路由: openai | deepseek")
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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=自动",
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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="是否允许导出")
|
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
|
2026-04-17 09:07:51 +08:00
|
|
|
|
sql: Optional[str] = Field(None, description="便捷字段:等价于 branch_result.sql(可包含 -- 注释)")
|
|
|
|
|
|
explanation: Optional[str] = Field(
|
|
|
|
|
|
None,
|
|
|
|
|
|
description="SQL 自然语言解释/交付说明(便捷字段;通常来自 sql_delivery_message 或 sql_explain)",
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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")
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
# 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加
|
|
|
|
|
|
_DEFAULT_SSE_CHUNK_CHARS = int(os.getenv("SSE_STREAM_CHUNK_CHARS", "64"))
|
|
|
|
|
|
|
2026-04-17 09:07:51 +08:00
|
|
|
|
# 流式输出形态:
|
2026-04-17 09:58:08 +08:00
|
|
|
|
# - 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 流式段落行为一致)。
|
2026-04-17 09:07:51 +08:00
|
|
|
|
_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")
|
|
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
|
|
|
|
|
|
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]}
|
|
|
|
|
|
)
|
2026-04-16 14:29:24 +08:00
|
|
|
|
await asyncio.sleep(0)
|
2026-04-16 09:15:01 +08:00
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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-16 15:49:03 +08:00
|
|
|
|
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,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
async def _load_session_text2sql_context(
|
|
|
|
|
|
request: NLChatRequest,
|
2026-04-15 17:33:07 +08:00
|
|
|
|
user_text: str,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
) -> 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)
|
2026-04-15 17:33:07 +08:00
|
|
|
|
# 方案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)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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,
|
|
|
|
|
|
}
|
2026-04-17 09:07:51 +08:00
|
|
|
|
payload["sql"] = sql
|
|
|
|
|
|
payload["explanation"] = (
|
|
|
|
|
|
str(sdm).strip()
|
|
|
|
|
|
if sdm
|
|
|
|
|
|
else (str(explain).strip() if explain else None)
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
qo = result.metadata.get("question_original")
|
|
|
|
|
|
qz = result.metadata.get("question_zh_normalized")
|
|
|
|
|
|
if qo and qz:
|
|
|
|
|
|
payload["query_normalization"] = {"original": qo, "zh": qz}
|
2026-04-16 13:48:44 +08:00
|
|
|
|
if result.metadata.get("fewshot_golden_reuse"):
|
|
|
|
|
|
payload["fewshot_golden"] = {
|
|
|
|
|
|
"qid": result.metadata.get("fewshot_golden_qid"),
|
|
|
|
|
|
"score": result.metadata.get("fewshot_golden_score"),
|
|
|
|
|
|
}
|
2026-04-14 10:28:22 +08:00
|
|
|
|
return payload
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
async def _run_generate(
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
top_k: int = 20,
|
|
|
|
|
|
dialog_context: Optional[str] = None,
|
2026-04-16 09:15:01 +08:00
|
|
|
|
request: Optional[NLChatRequest] = None,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
) -> 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-16 10:53:10 +08:00
|
|
|
|
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 ""),
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
|
|
|
|
|
|
if not request.message or not request.message.strip():
|
2026-04-16 09:15:01 +08:00
|
|
|
|
lang = _normalize_lang_code(request.lang_code)
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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})
|
2026-04-14 10:28:22 +08:00
|
|
|
|
return
|
|
|
|
|
|
text = request.message.strip()
|
2026-04-14 18:02:12 +08:00
|
|
|
|
orch = get_orchestrator()
|
2026-04-15 17:33:07 +08:00
|
|
|
|
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
2026-04-16 10:53:10 +08:00
|
|
|
|
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,
|
|
|
|
|
|
)
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if classified.intent == DialogIntent.CONVERSATION:
|
2026-04-16 09:15:01 +08:00
|
|
|
|
reply = (classified.reply_suggestion or "").strip()
|
|
|
|
|
|
if not reply:
|
|
|
|
|
|
reply = _localized_conversation_reply(lang)
|
2026-04-16 10:53:10 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
"[API/stream] 对话意图 conversation(跳过 Text2SQL): reply_chars=%s reply_preview=%r",
|
|
|
|
|
|
len(reply),
|
|
|
|
|
|
reply[:500] + ("…" if len(reply) > 500 else ""),
|
|
|
|
|
|
)
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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
|
2026-04-14 10:28:22 +08:00
|
|
|
|
data_dict = _conversation_nl_dict(reply)
|
2026-04-16 15:49:03 +08:00
|
|
|
|
data_dict["llm"] = _effective_llm_route_from_request(request)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
await _append_session_if_needed(request, text, data_dict)
|
2026-04-17 09:07:51 +08:00
|
|
|
|
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
|
2026-04-14 10:28:22 +08:00
|
|
|
|
return
|
|
|
|
|
|
|
2026-04-16 13:48:44 +08:00
|
|
|
|
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 ""),
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
try:
|
2026-04-16 13:48:44 +08:00
|
|
|
|
async with _maybe_override_orch_llm(orch, request) as o:
|
|
|
|
|
|
chunk_queue: asyncio.Queue[Optional[str]] = asyncio.Queue()
|
|
|
|
|
|
holder: Dict[str, Any] = {}
|
2026-04-17 09:07:51 +08:00
|
|
|
|
# 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
|
2026-04-16 13:48:44 +08:00
|
|
|
|
|
|
|
|
|
|
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))
|
2026-04-17 09:58:08 +08:00
|
|
|
|
# 仅 envelope 模式需要累积全文;dual/delta 只发 stage 增量,结束包再带完整 data
|
2026-04-17 09:07:51 +08:00
|
|
|
|
stream_acc = ""
|
2026-04-16 13:48:44 +08:00
|
|
|
|
while True:
|
|
|
|
|
|
piece = await chunk_queue.get()
|
|
|
|
|
|
if piece is None:
|
|
|
|
|
|
break
|
2026-04-16 14:29:24 +08:00
|
|
|
|
# 不做二次切分:直接按 LLM 原生 delta 逐条透传给前端
|
|
|
|
|
|
if piece:
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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:
|
2026-04-17 09:58:08 +08:00
|
|
|
|
# dual / delta:仅 stage 增量,避免每条再套一层完整 ApiEnvelope
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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 = ""
|
2026-04-16 13:48:44 +08:00
|
|
|
|
await gen_task
|
|
|
|
|
|
if holder.get("error"):
|
|
|
|
|
|
raise holder["error"]
|
|
|
|
|
|
result = holder["result"]
|
2026-04-14 10:28:22 +08:00
|
|
|
|
except Exception as e:
|
2026-04-16 10:53:10 +08:00
|
|
|
|
logger.error(f"[API/stream] 生成异常: {e}", exc_info=True)
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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})
|
2026-04-14 10:28:22 +08:00
|
|
|
|
return
|
2026-04-16 09:15:01 +08:00
|
|
|
|
|
2026-04-16 10:53:10 +08:00
|
|
|
|
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)
|
|
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
# 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
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
data_dict = _nl_dict_from_generation(result)
|
2026-04-16 15:49:03 +08:00
|
|
|
|
data_dict["llm"] = _effective_llm_route_from_request(request)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
await _append_session_if_needed(request, text, data_dict)
|
|
|
|
|
|
|
2026-04-17 09:07:51 +08:00
|
|
|
|
# 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})
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
@asynccontextmanager
|
2026-04-17 09:07:51 +08:00
|
|
|
|
async def lifespan(_app: FastAPI):
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-04-17 09:07:51 +08:00
|
|
|
|
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))
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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():
|
2026-04-16 09:15:01 +08:00
|
|
|
|
lang = _normalize_lang_code(request.lang_code)
|
|
|
|
|
|
raise HTTPException(status_code=400, detail=_localized_empty_input_reply(lang))
|
2026-04-14 10:28:22 +08:00
|
|
|
|
text = request.message.strip()
|
2026-04-14 18:02:12 +08:00
|
|
|
|
orch = get_orchestrator()
|
2026-04-15 17:33:07 +08:00
|
|
|
|
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
2026-04-16 10:53:10 +08:00
|
|
|
|
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,
|
|
|
|
|
|
)
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if classified.intent == DialogIntent.CONVERSATION:
|
2026-04-16 09:15:01 +08:00
|
|
|
|
reply = (classified.reply_suggestion or "").strip()
|
|
|
|
|
|
if not reply:
|
|
|
|
|
|
reply = _localized_conversation_reply(lang)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
|
|
|
|
|
|
data_dict = _conversation_nl_dict(reply)
|
2026-04-16 15:49:03 +08:00
|
|
|
|
data_dict["llm"] = _effective_llm_route_from_request(request)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
await _append_session_if_needed(request, text, data_dict)
|
|
|
|
|
|
return {"code": 200, "msg": "success", "data": data_dict}
|
|
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
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
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
data_dict = _nl_dict_from_generation(result)
|
2026-04-16 15:49:03 +08:00
|
|
|
|
data_dict["llm"] = _effective_llm_route_from_request(request)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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)。
|
2026-04-17 09:07:51 +08:00
|
|
|
|
|
2026-04-17 09:58:08 +08:00
|
|
|
|
默认(`SSE_STREAM_SHAPE=dual`)在 DATA_QUERY SQL 生成阶段只推送 stage delta(与 `delta` 一致);
|
|
|
|
|
|
寒暄等仍会按 dual 发送 delta + 增长的 `data.branch_result.answer`。结束包为完整
|
|
|
|
|
|
`{"code":200,"msg":"success|partial","data":{...}}`。
|
2026-04-17 09:07:51 +08:00
|
|
|
|
|
|
|
|
|
|
其它模式:
|
2026-04-17 09:58:08 +08:00
|
|
|
|
- `SSE_STREAM_SHAPE=envelope`:无 stage delta,`data.stream_narrative` 在 envelope 内逐步增长
|
|
|
|
|
|
- `SSE_STREAM_SHAPE=delta`:与当前 dual 在 SQL 流式段行为相同
|
2026-04-14 10:28:22 +08:00
|
|
|
|
"""
|
|
|
|
|
|
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")
|
2026-04-17 09:07:51 +08:00
|
|
|
|
async def nl_sql_execute():
|
2026-04-14 10:28:22 +08:00
|
|
|
|
return {
|
|
|
|
|
|
"code": 501,
|
|
|
|
|
|
"msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。",
|
|
|
|
|
|
"data": None,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@app.post("/g3sb/api/nl/operation-logs")
|
2026-04-17 09:07:51 +08:00
|
|
|
|
async def nl_post_operation_log():
|
2026-04-14 10:28:22 +08:00
|
|
|
|
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
|
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,
|
2026-04-16 10:53:10 +08:00
|
|
|
|
log_level="info",
|
|
|
|
|
|
# 必须为 None:False 仍会进入 uvicorn 的 fileConfig 分支并崩溃;None 才跳过覆盖 root 日志
|
|
|
|
|
|
log_config=None,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
)
|