Enhance LLM integration by adding OpenAI client support and enabling dynamic routing between DeepSeek and OpenAI services. Update environment configuration to include LLM_SERVICE_CODE for service selection, and modify API server to accommodate new request parameters for language and model. Implement streaming response improvements for chat interactions, allowing for segmented SSE output. Update documentation and impact analysis to reflect these changes.

This commit is contained in:
陈辅元
2026-04-16 09:15:01 +08:00
parent 85ad31348e
commit df33657099
17 changed files with 958 additions and 64 deletions
+336 -33
View File
@@ -19,7 +19,7 @@ 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 pydantic import AliasChoices, BaseModel, ConfigDict, Field
from dotenv import load_dotenv
_REPO_DIR = Path(__file__).resolve().parent
@@ -45,6 +45,9 @@ logging.basicConfig(
)
logger = logging.getLogger(__name__)
# per-request 覆盖 LLM client 时,使用锁避免并发串改 orchestrator.deepseek
_ORCH_LLM_LOCK = asyncio.Lock()
# 加载 .env 文件(支持 PyInstaller 打包后的目录结构)
def _find_env_file() -> Path:
"""查找 .env 文件,支持多种运行环境"""
@@ -100,7 +103,7 @@ def get_orchestrator():
raise RuntimeError(
"环境检查失败(请查看上方 WARNING/INFO,常见原因:"
"未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding;"
"或 Schema 文件路径不对、DEEPSEEK_API_KEY 未设置)"
"或 Schema 文件路径不对、LLM Key 未设置)"
)
schema_path = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
@@ -109,6 +112,7 @@ def get_orchestrator():
schema_manager = load_schema(schema_path, schema_meta_path)
class Args:
# LLM 路由由 LLM_SERVICE_CODE + 对应 Key 决定;此处仅保留历史字段以兼容 create_orchestrator 签名
api_key = os.getenv("DEEPSEEK_API_KEY")
model = os.getenv("MODEL_PRIMARY", "deepseek-chat")
temperature = float(os.getenv("TEMPERATURE", "0.3"))
@@ -133,9 +137,17 @@ def get_orchestrator():
class NLChatRequest(BaseModel):
model_config = ConfigDict(populate_by_name=True)
message: str = Field(..., description="用户输入的自然语言问题")
service_code: Optional[str] = Field(None, description="模型路由: openai | deepseek")
lang_code: Optional[str] = Field("auto", description="语言: zh | en | tc | auto")
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=自动",
)
taskId: Optional[str] = Field(None, description="任务ID")
session_id: Optional[str] = Field(None, description="会话ID")
visitor_biz_id: Optional[str] = Field(None, description="访客业务ID")
@@ -236,6 +248,27 @@ def _sse_data(obj: Dict[str, Any]) -> bytes:
return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n".encode("utf-8")
# 与前端 chatStore onDelta 一致:多次 { stage, stream_kind: content, content } 累加
_DEFAULT_SSE_CHUNK_CHARS = int(os.getenv("SSE_STREAM_CHUNK_CHARS", "64"))
async def _sse_stream_text_chunks(
stage: str,
content: str,
*,
chunk_size: Optional[int] = None,
) -> AsyncIterator[bytes]:
"""将长文本拆成多段 SSE,便于浏览器逐段渲染(流式)。"""
if not content:
return
size = max(8, chunk_size or _DEFAULT_SSE_CHUNK_CHARS)
for i in range(0, len(content), size):
yield _sse_data(
{"stage": stage, "stream_kind": "content", "content": content[i : i + size]}
)
await asyncio.sleep(0)
def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
"""寒暄/元问题等:与前端 BUSINESS_MANUAL + branch_result.answer 一致。"""
r = (reply or "").strip() or "您好。"
@@ -245,6 +278,226 @@ def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
}
def _normalize_lang_code(code: Optional[str]) -> str:
raw = (code or "auto").strip()
if not raw:
return "auto"
# 兼容中文取值
if raw in ("简体中文", "简体", "中文(简体)"):
return "zh"
if raw in ("繁体中文", "繁體中文", "繁体", "繁體"):
return "tc"
if raw in ("英语", "英文", "英語"):
return "en"
c = raw.lower().replace("_", "-")
if c in ("zh", "tc", "en", "auto"):
return c
# 兼容前端可能传的语言标签
if c in ("english", "en-us", "en-gb"):
return "en"
if c in ("zh-cn", "zh-hans", "zh-sg"):
return "zh"
if c in ("zh-tw", "zh-hant", "zh-hk", "traditional", "traditional-chinese"):
return "tc"
return "auto"
def _normalize_model_label(model: Optional[str]) -> Optional[str]:
"""
前端可能传展示文案(如 'GPT-4o mini' / 'DeepSeek V3')。
这里统一映射到真实模型名(openai/deepseek 各自可用的 model id)。
"""
raw = (model or "").strip()
if not raw:
return None
m = raw.strip().lower()
m = m.replace("_", "-").replace(" ", "-")
while "--" in m:
m = m.replace("--", "-")
# OpenAI
if m in ("gpt-4o-mini", "gpt4o-mini", "gpt-4o-mini"):
return "gpt-4o-mini"
if m in ("gpt-4o", "gpt4o"):
return "gpt-4o"
# DeepSeek
if m in ("deepseek-v3", "deepseekv3", "deepseek-v3.0", "deepseekv3.0", "deepseek-v3"):
return "deepseek-chat"
if m in ("deepseek-chat", "deepseek-reasoner"):
return m
return raw # 未知时原样透传(交由网关/后端判定)
def _infer_service_code(service_code: Optional[str], model_name: Optional[str]) -> Optional[str]:
sc = (service_code or "").strip().lower()
if sc:
if sc in ("openai", "deepseek"):
return sc
if sc in ("gpt", "chatgpt"):
return "openai"
if "deepseek" in sc:
return "deepseek"
m = (model_name or "").strip().lower()
if m:
if "deepseek" in m:
return "deepseek"
if m.startswith("gpt"):
return "openai"
return None
def _has_openai_key() -> bool:
return bool((os.getenv("OPENAI_API_KEY") or "").strip())
def _has_deepseek_key() -> bool:
return bool((os.getenv("DEEPSEEK_API_KEY") or "").strip())
def _localized_conversation_reply(lang_code: str) -> str:
if lang_code == "en":
return (
"Hello, I'm the Text2SQL assistant.\n"
"Describe what you want to query or aggregate in natural language "
"(e.g., available balance of an account; summarize unsettled trades by broker).\n"
"Type quit or exit to leave."
)
if lang_code == "tc":
return (
"您好,我是業務庫 Text2SQL 助手。\n"
"請用自然語言描述要查詢或統計的內容(例如:查詢某帳戶可用餘額、按經紀商匯總未結算交易筆數)。\n"
"輸入 quit 或 exit 可退出。"
)
# zh / auto
return (
"您好,我是业务库 Text2SQL 助手。\n"
"请用自然语言描述要查询或统计的内容(例如:查询某账户可用余额、按经纪商汇总未结算交易笔数)。\n"
"输入 quit 或 exit 可退出。"
)
def _localized_empty_input_reply(lang_code: str) -> str:
if lang_code == "en":
return "Please enter a concrete business query, or type quit to exit."
if lang_code == "tc":
return "請輸入具體的業務查詢問題,或輸入 quit 退出。"
return "请输入具体的业务查询问题,或输入 quit 退出。"
def _lang_label(lang_code: str) -> str:
if lang_code == "en":
return "English"
if lang_code == "tc":
return "Traditional Chinese"
return "Simplified Chinese"
async def _translate_explain_text(
orch,
text: Optional[str],
lang_code: str,
) -> Optional[str]:
"""
将“说明类”文本翻译到 lang_code(仅用于 sql_delivery_message / db_empty_feedback / sql_explain)。
zh 直接返回;en/tc 使用 LLM 翻译(短输出,避免引入格式)。
"""
t = (text or "").strip()
if not t:
return text
if lang_code in ("auto", "zh"):
return text
target = _lang_label(lang_code)
messages = [
{
"role": "system",
"content": (
"You are a translation assistant.\n"
f"Translate the following text to {target}.\n"
"Rules:\n"
"- Keep the meaning identical; do not add new information.\n"
"- Keep SQL keywords/code unchanged if present.\n"
"- Output plain text only (no Markdown, no code fences, no JSON).\n"
),
},
{"role": "user", "content": t},
]
def _call() -> str:
msg = orch.deepseek.chat(messages, temperature=0.0, top_p=1.0, max_completion_tokens=320)
return (msg.content or "").strip()
try:
out = await asyncio.to_thread(_call)
return out or text
except Exception as e:
logger.warning("[API] explain translation skipped: %s", e)
return text
@asynccontextmanager
async def _maybe_override_orch_llm(orch, request: NLChatRequest):
"""
若请求传 service_code / model,则临时覆盖 orchestrator.deepseek;
使用锁保证同一时刻仅一个请求修改该引用。
"""
model = _normalize_model_label(request.model)
sc = _infer_service_code(request.service_code, model)
lang = _normalize_lang_code(request.lang_code)
needs_override = (sc is not None) or (model is not None) or (lang == "en")
if not needs_override:
yield orch
return
from llm.router import create_llm_client
tmp = None
if sc is not None or model is not None:
# 兜底:前端切什么就用什么;但若对应 key 未配置,则自动回退到另一家,保证可用
sc2 = sc
if sc2 == "deepseek" and not _has_deepseek_key():
if _has_openai_key():
logger.warning("[API] deepseek requested but key missing, fallback to openai")
sc2 = "openai"
if sc2 == "openai" and not _has_openai_key():
if _has_deepseek_key():
logger.warning("[API] openai requested but key missing, fallback to deepseek")
sc2 = "deepseek"
llm_kwargs: Dict[str, Any] = {}
if model is not None:
# 保护:openai 网关下 deepseek-chat 会 404;deepseek 官方下 gpt-* 也会失败
if (sc2 or "").strip().lower() == "openai" and "deepseek" in model.strip().lower():
pass
elif (sc2 or "").strip().lower() == "deepseek" and model.strip().lower().startswith("gpt"):
pass
else:
llm_kwargs["model_name"] = model
tmp = create_llm_client(sc2, **llm_kwargs)
async with _ORCH_LLM_LOCK:
old = getattr(orch, "deepseek", None)
old_translate = getattr(orch, "translate_english_to_zh", None)
if tmp is not None:
orch.deepseek = tmp
# 英文界面:不要把英文问句归一成中文,否则下游解释会倾向中文
if lang == "en" and hasattr(orch, "translate_english_to_zh"):
orch.translate_english_to_zh = False
try:
yield orch
finally:
if tmp is not None:
orch.deepseek = old
if old_translate is not None and hasattr(orch, "translate_english_to_zh"):
orch.translate_english_to_zh = old_translate
async def _load_session_text2sql_context(
request: NLChatRequest,
user_text: str,
@@ -359,41 +612,54 @@ async def _run_generate(
question: str,
top_k: int = 20,
dialog_context: Optional[str] = None,
request: Optional[NLChatRequest] = None,
) -> GenerationResult:
orch = get_orchestrator()
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
dc = (dialog_context or "").strip() or None
def _call() -> GenerationResult:
return orch.generate(
question=question.strip(),
dialect=dialect,
top_k_candidates=top_k,
dialog_context=dc,
)
async def _call_with_orch(o) -> GenerationResult:
def _call() -> GenerationResult:
return o.generate(
question=question.strip(),
dialect=dialect,
top_k_candidates=top_k,
dialog_context=dc,
)
return await asyncio.to_thread(_call)
return await asyncio.to_thread(_call)
if request is None:
return await _call_with_orch(orch)
async with _maybe_override_orch_llm(orch, request) as o:
return await _call_with_orch(o)
async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
if not request.message or not request.message.strip():
yield _sse_data({"code": 400, "msg": "消息内容不能为空", "data": None})
lang = _normalize_lang_code(request.lang_code)
yield _sse_data({"code": 400, "msg": _localized_empty_input_reply(lang), "data": None})
return
text = request.message.strip()
orch = get_orchestrator()
dialog_block, last_data = await _load_session_text2sql_context(request, text)
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,
)
lang = _normalize_lang_code(request.lang_code)
async with _maybe_override_orch_llm(orch, request) as o:
classified = await asyncio.to_thread(
classify_dialog,
text,
last_turn_was_data_query=last_data,
dialog_context=dialog_block or None,
llm_client=o.deepseek,
)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
reply = (classified.reply_suggestion or "").strip()
if not reply:
reply = _localized_conversation_reply(lang)
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})
async for pkt in _sse_stream_text_chunks("chat", reply):
yield pkt
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)
@@ -401,12 +667,29 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "DATA_QUERY"})
try:
result = await _run_generate(text, dialog_context=dialog_block or None)
result = await _run_generate(text, dialog_context=dialog_block or None, request=request)
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)})
# 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
html = _sql_gen_stream_html(result)
async for pkt in _sse_stream_text_chunks("sql_gen", html):
yield pkt
data_dict = _nl_dict_from_generation(result)
msg = "success" if result.valid else "partial"
yield _sse_data({"code": 200, "msg": msg, "data": data_dict})
@@ -461,26 +744,46 @@ async def nl_chat(request: NLChatRequest):
"""
try:
if not request.message or not request.message.strip():
raise HTTPException(status_code=400, detail="消息内容不能为空")
lang = _normalize_lang_code(request.lang_code)
raise HTTPException(status_code=400, detail=_localized_empty_input_reply(lang))
text = request.message.strip()
logger.info(f"[API] 问题: {text[:100]}...")
orch = get_orchestrator()
dialog_block, last_data = await _load_session_text2sql_context(request, text)
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,
)
lang = _normalize_lang_code(request.lang_code)
async with _maybe_override_orch_llm(orch, request) as o:
classified = await asyncio.to_thread(
classify_dialog,
text,
last_turn_was_data_query=last_data,
dialog_context=dialog_block or None,
llm_client=o.deepseek,
)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
reply = (classified.reply_suggestion or "").strip()
if not reply:
reply = _localized_conversation_reply(lang)
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
data_dict = _conversation_nl_dict(reply)
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)
result = await _run_generate(text, dialog_context=dialog_block or None, request=request)
# What this SQL does / SQL 说明:按 lang_code 本地化
try:
if isinstance(result.metadata, dict):
if result.metadata.get("sql_delivery_message"):
result.metadata["sql_delivery_message"] = await _translate_explain_text(
orch, str(result.metadata["sql_delivery_message"]), lang
)
if result.metadata.get("db_empty_feedback"):
result.metadata["db_empty_feedback"] = await _translate_explain_text(
orch, str(result.metadata["db_empty_feedback"]), lang
)
except Exception:
pass
data_dict = _nl_dict_from_generation(result)
response_data = NLChatSuccessData.model_validate(data_dict)
msg = "success" if result.valid else "partial"