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:
@@ -2,12 +2,23 @@
|
|||||||
# 复制为 .env 并填入实际值
|
# 复制为 .env 并填入实际值
|
||||||
|
|
||||||
# ========== DeepSeek API 配置 ==========
|
# ========== DeepSeek API 配置 ==========
|
||||||
DEEPSEEK_API_KEY=sk-2e3cc9f79a2c40e9a6b7c960d8c090d9
|
# LLM 路由(deepseek | openai),不填则按 Key 自动选择(优先 deepseek)
|
||||||
|
LLM_SERVICE_CODE=openai
|
||||||
|
|
||||||
|
|
||||||
|
DEEPSEEK_API_KEY=sk-041b3995415c4445a1713750fcc9e023
|
||||||
DEEPSEEK_BASE_URL=https://api.deepseek.com
|
DEEPSEEK_BASE_URL=https://api.deepseek.com
|
||||||
|
MODEL_PRIMARY=deepseek-chat
|
||||||
|
|
||||||
|
|
||||||
|
OPENAI_API_KEY=sk-proj-FEb2ChHZK5Llm0tBkmNIT3BlbkFJbu0b0zQpTW8762yD6HDv
|
||||||
|
OPENAI_BASE_URL=http://113.192.49.54:9080/v1
|
||||||
|
OPENAI_MODEL=gpt-5.4
|
||||||
|
# OPENAI_CHAT_MODEL=gpt-4o-mini
|
||||||
|
|
||||||
# 模型配置
|
# 模型配置
|
||||||
# 主生成模型
|
# 主生成模型
|
||||||
MODEL_PRIMARY=deepseek-chat
|
|
||||||
TEMPERATURE=0 # 确定性模式:每次生成相同结果
|
TEMPERATURE=0 # 确定性模式:每次生成相同结果
|
||||||
MAX_TOKENS=4096
|
MAX_TOKENS=4096
|
||||||
# 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等
|
# 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等
|
||||||
|
|||||||
@@ -42,6 +42,121 @@
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
# Impact Analysis Report — 新增 OpenAI LLM Client(对话/JSON 输出)
|
||||||
|
|
||||||
|
## 1. 改动概览
|
||||||
|
|
||||||
|
- **背景与目标**:补齐 `backend/llm/openai_client.py`,提供 OpenAI(或 OpenAI 兼容网关)调用封装,支持同步/异步 `chat` 与 `chat_with_json`,便于与既有 prompt/JSON 输出链路复用。
|
||||||
|
- **涉及模块**:`backend/llm/openai_client.py`、`.env`(仅补充注释示例,不影响现有运行配置)。
|
||||||
|
- **改动类型**:功能新增。
|
||||||
|
|
||||||
|
## 2. 方法级改动分析
|
||||||
|
|
||||||
|
| 位置 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `OpenAIClient.chat` | 新增:基于 `openai` SDK 的 Chat Completions 调用封装,支持覆盖 `temperature/max_tokens/top_p` 等参数。 |
|
||||||
|
| `OpenAIClient.chat_with_json` | 新增:抽取/解析 JSON(兼容 markdown code fence);解析失败时返回 `{"_json_decode_failed": true, "raw_content": ...}`。 |
|
||||||
|
| `create_openai_client` | 新增:从环境变量读取 `OPENAI_API_KEY`、可选 `OPENAI_BASE_URL`、`OPENAI_MODEL/OPENAI_CHAT_MODEL`。 |
|
||||||
|
|
||||||
|
## 3. 调用方与影响范围分析
|
||||||
|
|
||||||
|
- **调用方**:当前仓库运行链路仍默认使用 `DeepSeekClient`;本次新增仅提供可选能力,未修改既有编排器/接口路由。
|
||||||
|
- **破坏性变更**:否。
|
||||||
|
|
||||||
|
## 4. 风险与回滚
|
||||||
|
|
||||||
|
- **风险级别**:低(新增模块,不改变既有默认路径)。
|
||||||
|
- **回滚**:删除新增文件与对应文档/注释即可。
|
||||||
|
|
||||||
|
**回滚方式是否简单**:是。
|
||||||
|
|
||||||
|
## 5. 验证与测试
|
||||||
|
|
||||||
|
- 已执行:`python -m py_compile backend/llm/openai_client.py`。
|
||||||
|
|
||||||
|
## 6. 配置变更
|
||||||
|
|
||||||
|
- `.env`:补充 `OPENAI_MODEL/OPENAI_CHAT_MODEL` 注释示例(不修改现有真实配置)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
# Impact Analysis Report — DeepSeek/OpenAI 双模型切换(LLM_SERVICE_CODE)
|
||||||
|
|
||||||
|
## 1. 改动概览
|
||||||
|
|
||||||
|
- **背景与目标**:支持在 DeepSeek 与 OpenAI(或兼容网关)之间切换 LLM 调用来源,便于在不同环境/配额下切换推理服务。
|
||||||
|
- **涉及模块**:`backend/llm/router.py`(新)、`backend/agents/orchestrator.py`、`backend/main.py`、`api_server.py`、`.env`(新增开关)。
|
||||||
|
- **改动类型**:功能增强(可配置路由)。
|
||||||
|
|
||||||
|
## 2. 方法级改动分析
|
||||||
|
|
||||||
|
| 位置 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `backend/llm/router.py` | 新增:`resolve_llm_service_code` + `create_llm_client`,按 `LLM_SERVICE_CODE` 或 Key 存在性创建 LLM Client。 |
|
||||||
|
| `Text2SQLOrchestrator.__init__` | 新增可选 `llm_client` 参数;若传入则作为 `self.deepseek` 使用(保留属性名避免大范围改动)。 |
|
||||||
|
| `backend/main.py` | `setup_environment` 增加 `LLM_SERVICE_CODE` 校验;`create_orchestrator` 通过 router 创建 LLM Client。 |
|
||||||
|
| `api_server.py` | 错误提示文案更新;编排器初始化保持通过 `create_orchestrator` 完成。 |
|
||||||
|
|
||||||
|
## 3. 调用方与影响范围分析
|
||||||
|
|
||||||
|
- **调用方**:CLI(`backend/main.py`)与 API(`api_server.py`)初始化编排器路径。
|
||||||
|
- **行为变化**:
|
||||||
|
- `LLM_SERVICE_CODE=openai`:使用 `OPENAI_*` 做 LLM 调用(意图分类 / 归一 / 选表 / 生成 / 说明等)。
|
||||||
|
- `LLM_SERVICE_CODE=deepseek` 或未设置但存在 `DEEPSEEK_API_KEY`:仍默认 DeepSeek(与历史一致)。
|
||||||
|
- **破坏性变更**:否(对外 API 入参/出参不变;仅初始化与内部客户端来源可切换)。
|
||||||
|
|
||||||
|
## 4. 风险与回滚
|
||||||
|
|
||||||
|
- **风险级别**:中(不同模型对 JSON 严格性/输出格式偏好不同,可能影响 `chat_with_json` 的解析成功率与稳定性;失败时仍会返回 `_json_decode_failed` 供上游处理)。
|
||||||
|
- **回滚**:将 `LLM_SERVICE_CODE` 切回 `deepseek` 或回退相关文件改动。
|
||||||
|
|
||||||
|
**回滚方式是否简单**:是。
|
||||||
|
|
||||||
|
## 5. 验证与测试
|
||||||
|
|
||||||
|
- 已执行:`python -m py_compile` 覆盖 `backend/llm/router.py`、`backend/agents/orchestrator.py`、`backend/main.py`、`api_server.py`。
|
||||||
|
|
||||||
|
## 6. 配置变更
|
||||||
|
|
||||||
|
| 配置项 | 含义 | 取值 |
|
||||||
|
|--------|------|------|
|
||||||
|
| `LLM_SERVICE_CODE` | 选择 LLM 路由 | `openai` / `deepseek`(留空自动按 Key 选择,优先 deepseek) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
# Impact Analysis Report — 适配前端参数(lang_code / model)
|
||||||
|
|
||||||
|
## 1. 改动概览
|
||||||
|
|
||||||
|
- **背景与目标**:前端请求会携带 `lang_code`(`zh`/`tc`/`en`)与 `model`(模型名覆盖);后端需接收并在对话分支/LLM 调用侧生效。
|
||||||
|
- **涉及模块**:`api_server.py`。
|
||||||
|
- **改动类型**:兼容性增强。
|
||||||
|
|
||||||
|
## 2. 方法级改动分析
|
||||||
|
|
||||||
|
| 位置 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `NLChatRequest` | 新增 `model` 字段;保留 `lang_code`(规范化为 `zh/tc/en/auto`)。 |
|
||||||
|
| `nl_chat` / `_chat_stream_events` | 对话分支:空输入/默认引导文案按 `lang_code` 返回(中/繁/英)。 |
|
||||||
|
| `_maybe_override_orch_llm`(新) | 若请求带 `service_code/model`,临时覆盖 `orchestrator.deepseek` 为按请求创建的 LLM client;使用锁避免并发串改。 |
|
||||||
|
| `_run_generate` | 增加 `request` 参数,用于按请求覆盖 LLM client 后再生成。 |
|
||||||
|
|
||||||
|
## 3. 调用方与影响范围分析
|
||||||
|
|
||||||
|
- **调用方**:`/g3sb/api/nl/chat`、`/g3sb/api/nl/chat/stream`。
|
||||||
|
- **破坏性变更**:否(仅新增可选请求字段与内部适配;未提供时行为保持不变)。
|
||||||
|
|
||||||
|
## 4. 风险与回滚
|
||||||
|
|
||||||
|
- **风险级别**:中(带 `model/service_code` 的请求会串行化执行 LLM 调用,避免并发污染;在高并发时可能降低吞吐)。
|
||||||
|
- **回滚**:回退 `api_server.py` 的按请求覆盖逻辑即可。
|
||||||
|
|
||||||
|
**回滚方式是否简单**:是。
|
||||||
|
|
||||||
|
## 5. 验证与测试
|
||||||
|
|
||||||
|
- 已执行:`python -m py_compile api_server.py`。
|
||||||
|
|
||||||
# Impact Analysis Report — 方案A:新话题禁用 dialog_context(防上下文污染)
|
# Impact Analysis Report — 方案A:新话题禁用 dialog_context(防上下文污染)
|
||||||
|
|
||||||
## 1. 改动概览
|
## 1. 改动概览
|
||||||
@@ -409,3 +524,30 @@
|
|||||||
## 4. 配置
|
## 4. 配置
|
||||||
|
|
||||||
| `FEWSHOT_CHROMA_EPHEMERAL` | 可选;`true` 时与旧版「仅内存」行为一致。 |
|
| `FEWSHOT_CHROMA_EPHEMERAL` | 可选;`true` 时与旧版「仅内存」行为一致。 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
# Impact Analysis Report — `/g3sb/api/nl/chat/stream` 分片 SSE(追加)
|
||||||
|
|
||||||
|
## 1. 改动概览
|
||||||
|
|
||||||
|
- **背景与目标**:前端 `chatStore` 按多次 `{ stage, stream_kind, content }` 累加 `chatText` / `sqlGenHtml`;原先后端在 `chat` / `sql_gen` 阶段各只发一条大包,网络侧无渐进感。
|
||||||
|
- **涉及模块**:`api_server.py`。
|
||||||
|
- **改动类型**:行为优化(SSE 事件形态不变,仅增加条数)。
|
||||||
|
|
||||||
|
## 2. 方法级改动
|
||||||
|
|
||||||
|
| 位置 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `_sse_stream_text_chunks` | 新增:按字符窗口(默认 64,可用 `SSE_STREAM_CHUNK_CHARS` 覆盖)拆成多段 SSE;段间 `asyncio.sleep(0)` 让出事件循环。 |
|
||||||
|
| `_chat_stream_events` | `chat` / `sql_gen` 正文改为 `async for` 分片 `yield`。 |
|
||||||
|
|
||||||
|
## 3. 破坏性变更
|
||||||
|
|
||||||
|
- **否**(结束包 `code/msg/data` 仍与原先一致)。
|
||||||
|
|
||||||
|
## 4. 配置变更
|
||||||
|
|
||||||
|
| 项 | 说明 |
|
||||||
|
|----|------|
|
||||||
|
| `SSE_STREAM_CHUNK_CHARS` | 可选;每段 SSE 的 `content` 最大字符数,默认 `64`。 |
|
||||||
|
|||||||
@@ -12,6 +12,10 @@
|
|||||||
- **新话题**:不注入历史摘要(`max_pairs=0`)
|
- **新话题**:不注入历史摘要(`max_pairs=0`)
|
||||||
- **目的**:减少多话题情况下历史 SQL/条件回流导致的串话与错误 SQL。
|
- **目的**:减少多话题情况下历史 SQL/条件回流导致的串话与错误 SQL。
|
||||||
|
|
||||||
|
- **新增(本次补充)**:`backend/llm/openai_client.py`
|
||||||
|
- 提供 `OpenAIClient.chat` / `chat_with_json`(以及与 DeepSeekClient 对齐的若干便捷方法),用于接入 OpenAI 或 OpenAI 兼容网关。
|
||||||
|
- `.env` 增补 `OPENAI_MODEL/OPENAI_CHAT_MODEL` 的注释示例(不影响现有配置)。
|
||||||
|
|
||||||
## 影响与风险
|
## 影响与风险
|
||||||
|
|
||||||
- **破坏性变更**:否(对外 API 不变)
|
- **破坏性变更**:否(对外 API 不变)
|
||||||
@@ -26,6 +30,22 @@
|
|||||||
- 同一 session:先做话题A数据查询,再提一个完全不同话题B(无“再/按上面”等续问词),应不再引用 A 的表/过滤条件。
|
- 同一 session:先做话题A数据查询,再提一个完全不同话题B(无“再/按上面”等续问词),应不再引用 A 的表/过滤条件。
|
||||||
- 同一 session:首轮查询后,第二轮使用“再/按上面/沿用口径”等续问词,应仍能续用上轮口径生成 SQL。
|
- 同一 session:首轮查询后,第二轮使用“再/按上面/沿用口径”等续问词,应仍能续用上轮口径生成 SQL。
|
||||||
|
|
||||||
|
- 已执行:
|
||||||
|
- `python -m py_compile backend/llm/openai_client.py`
|
||||||
|
|
||||||
|
## 配置说明(补充)
|
||||||
|
|
||||||
|
- **LLM 切换**:通过 `.env` 的 `LLM_SERVICE_CODE` 在 `deepseek` / `openai` 之间切换(留空则按 Key 自动选择,优先 deepseek)。
|
||||||
|
|
||||||
|
## 前端参数适配(补充)
|
||||||
|
|
||||||
|
- **lang_code**:支持 `zh`(简体)、`tc`(繁体)、`en`(英语);在对话分支的默认引导文案/空输入提示中生效。
|
||||||
|
- **model**:请求体可传 `model` 覆盖默认模型名(与 `service_code` 搭配使用),用于按请求临时切换模型。
|
||||||
|
|
||||||
|
## 流式输出(补充)
|
||||||
|
|
||||||
|
- **`/g3sb/api/nl/chat/stream`**:`chat` 与 `sql_gen` 阶段的正文改为多段 SSE 分片推送(默认每段约 64 字符,可用环境变量 `SSE_STREAM_CHUNK_CHARS` 调整),与前端 `onDelta` 累加逻辑一致。
|
||||||
|
|
||||||
## 后续事项
|
## 后续事项
|
||||||
|
|
||||||
- 若误判率偏高,可迭代 `is_likely_follow_up` 规则(补充关键词/短语),或升级为轻量 LLM 判别器(方案B)。
|
- 若误判率偏高,可迭代 `is_likely_follow_up` 规则(补充关键词/短语),或升级为轻量 LLM 判别器(方案B)。
|
||||||
|
|||||||
Binary file not shown.
+317
-14
@@ -19,7 +19,7 @@ from contextlib import asynccontextmanager
|
|||||||
from fastapi import FastAPI, HTTPException, Body, Path as FPath
|
from fastapi import FastAPI, HTTPException, Body, Path as FPath
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from fastapi.responses import StreamingResponse, JSONResponse
|
from fastapi.responses import StreamingResponse, JSONResponse
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import AliasChoices, BaseModel, ConfigDict, Field
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
_REPO_DIR = Path(__file__).resolve().parent
|
_REPO_DIR = Path(__file__).resolve().parent
|
||||||
@@ -45,6 +45,9 @@ logging.basicConfig(
|
|||||||
)
|
)
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# per-request 覆盖 LLM client 时,使用锁避免并发串改 orchestrator.deepseek
|
||||||
|
_ORCH_LLM_LOCK = asyncio.Lock()
|
||||||
|
|
||||||
# 加载 .env 文件(支持 PyInstaller 打包后的目录结构)
|
# 加载 .env 文件(支持 PyInstaller 打包后的目录结构)
|
||||||
def _find_env_file() -> Path:
|
def _find_env_file() -> Path:
|
||||||
"""查找 .env 文件,支持多种运行环境"""
|
"""查找 .env 文件,支持多种运行环境"""
|
||||||
@@ -100,7 +103,7 @@ def get_orchestrator():
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"环境检查失败(请查看上方 WARNING/INFO,常见原因:"
|
"环境检查失败(请查看上方 WARNING/INFO,常见原因:"
|
||||||
"未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding;"
|
"未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding;"
|
||||||
"或 Schema 文件路径不对、DEEPSEEK_API_KEY 未设置)"
|
"或 Schema 文件路径不对、LLM Key 未设置)"
|
||||||
)
|
)
|
||||||
|
|
||||||
schema_path = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
|
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)
|
schema_manager = load_schema(schema_path, schema_meta_path)
|
||||||
|
|
||||||
class Args:
|
class Args:
|
||||||
|
# LLM 路由由 LLM_SERVICE_CODE + 对应 Key 决定;此处仅保留历史字段以兼容 create_orchestrator 签名
|
||||||
api_key = os.getenv("DEEPSEEK_API_KEY")
|
api_key = os.getenv("DEEPSEEK_API_KEY")
|
||||||
model = os.getenv("MODEL_PRIMARY", "deepseek-chat")
|
model = os.getenv("MODEL_PRIMARY", "deepseek-chat")
|
||||||
temperature = float(os.getenv("TEMPERATURE", "0.3"))
|
temperature = float(os.getenv("TEMPERATURE", "0.3"))
|
||||||
@@ -133,9 +137,17 @@ def get_orchestrator():
|
|||||||
|
|
||||||
|
|
||||||
class NLChatRequest(BaseModel):
|
class NLChatRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(populate_by_name=True)
|
||||||
|
|
||||||
message: str = Field(..., description="用户输入的自然语言问题")
|
message: str = Field(..., description="用户输入的自然语言问题")
|
||||||
service_code: Optional[str] = Field(None, description="模型路由: openai | deepseek")
|
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")
|
taskId: Optional[str] = Field(None, description="任务ID")
|
||||||
session_id: Optional[str] = Field(None, description="会话ID")
|
session_id: Optional[str] = Field(None, description="会话ID")
|
||||||
visitor_biz_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")
|
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]:
|
def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
|
||||||
"""寒暄/元问题等:与前端 BUSINESS_MANUAL + branch_result.answer 一致。"""
|
"""寒暄/元问题等:与前端 BUSINESS_MANUAL + branch_result.answer 一致。"""
|
||||||
r = (reply or "").strip() or "您好。"
|
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(
|
async def _load_session_text2sql_context(
|
||||||
request: NLChatRequest,
|
request: NLChatRequest,
|
||||||
user_text: str,
|
user_text: str,
|
||||||
@@ -359,13 +612,15 @@ async def _run_generate(
|
|||||||
question: str,
|
question: str,
|
||||||
top_k: int = 20,
|
top_k: int = 20,
|
||||||
dialog_context: Optional[str] = None,
|
dialog_context: Optional[str] = None,
|
||||||
|
request: Optional[NLChatRequest] = None,
|
||||||
) -> GenerationResult:
|
) -> GenerationResult:
|
||||||
orch = get_orchestrator()
|
orch = get_orchestrator()
|
||||||
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
|
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
|
||||||
dc = (dialog_context or "").strip() or None
|
dc = (dialog_context or "").strip() or None
|
||||||
|
|
||||||
|
async def _call_with_orch(o) -> GenerationResult:
|
||||||
def _call() -> GenerationResult:
|
def _call() -> GenerationResult:
|
||||||
return orch.generate(
|
return o.generate(
|
||||||
question=question.strip(),
|
question=question.strip(),
|
||||||
dialect=dialect,
|
dialect=dialect,
|
||||||
top_k_candidates=top_k,
|
top_k_candidates=top_k,
|
||||||
@@ -374,26 +629,37 @@ async def _run_generate(
|
|||||||
|
|
||||||
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]:
|
async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
|
||||||
if not request.message or not request.message.strip():
|
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
|
return
|
||||||
text = request.message.strip()
|
text = request.message.strip()
|
||||||
orch = get_orchestrator()
|
orch = get_orchestrator()
|
||||||
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
||||||
|
lang = _normalize_lang_code(request.lang_code)
|
||||||
|
async with _maybe_override_orch_llm(orch, request) as o:
|
||||||
classified = await asyncio.to_thread(
|
classified = await asyncio.to_thread(
|
||||||
classify_dialog,
|
classify_dialog,
|
||||||
text,
|
text,
|
||||||
last_turn_was_data_query=last_data,
|
last_turn_was_data_query=last_data,
|
||||||
dialog_context=dialog_block or None,
|
dialog_context=dialog_block or None,
|
||||||
llm_client=orch.deepseek,
|
llm_client=o.deepseek,
|
||||||
)
|
)
|
||||||
if classified.intent == DialogIntent.CONVERSATION:
|
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 一致)")
|
logger.info("[API/stream] 对话意图: conversation(跳过 Text2SQL,与 CLI single_query 一致)")
|
||||||
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "CHAT"})
|
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)
|
data_dict = _conversation_nl_dict(reply)
|
||||||
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
|
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
|
||||||
await _append_session_if_needed(request, text, 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"})
|
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "DATA_QUERY"})
|
||||||
try:
|
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:
|
except Exception as e:
|
||||||
logger.error(f"[API/stream] 生成异常: {e}")
|
logger.error(f"[API/stream] 生成异常: {e}")
|
||||||
yield _sse_data({"code": 500, "msg": str(e), "data": None})
|
yield _sse_data({"code": 500, "msg": str(e), "data": None})
|
||||||
return
|
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)
|
data_dict = _nl_dict_from_generation(result)
|
||||||
msg = "success" if result.valid else "partial"
|
msg = "success" if result.valid else "partial"
|
||||||
yield _sse_data({"code": 200, "msg": msg, "data": data_dict})
|
yield _sse_data({"code": 200, "msg": msg, "data": data_dict})
|
||||||
@@ -461,26 +744,46 @@ async def nl_chat(request: NLChatRequest):
|
|||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
if not request.message or not request.message.strip():
|
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()
|
text = request.message.strip()
|
||||||
logger.info(f"[API] 问题: {text[:100]}...")
|
logger.info(f"[API] 问题: {text[:100]}...")
|
||||||
orch = get_orchestrator()
|
orch = get_orchestrator()
|
||||||
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
dialog_block, last_data = await _load_session_text2sql_context(request, text)
|
||||||
|
lang = _normalize_lang_code(request.lang_code)
|
||||||
|
async with _maybe_override_orch_llm(orch, request) as o:
|
||||||
classified = await asyncio.to_thread(
|
classified = await asyncio.to_thread(
|
||||||
classify_dialog,
|
classify_dialog,
|
||||||
text,
|
text,
|
||||||
last_turn_was_data_query=last_data,
|
last_turn_was_data_query=last_data,
|
||||||
dialog_context=dialog_block or None,
|
dialog_context=dialog_block or None,
|
||||||
llm_client=orch.deepseek,
|
llm_client=o.deepseek,
|
||||||
)
|
)
|
||||||
if classified.intent == DialogIntent.CONVERSATION:
|
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)")
|
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
|
||||||
data_dict = _conversation_nl_dict(reply)
|
data_dict = _conversation_nl_dict(reply)
|
||||||
await _append_session_if_needed(request, text, data_dict)
|
await _append_session_if_needed(request, text, data_dict)
|
||||||
return {"code": 200, "msg": "success", "data": 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)
|
data_dict = _nl_dict_from_generation(result)
|
||||||
response_data = NLChatSuccessData.model_validate(data_dict)
|
response_data = NLChatSuccessData.model_validate(data_dict)
|
||||||
msg = "success" if result.valid else "partial"
|
msg = "success" if result.valid else "partial"
|
||||||
|
|||||||
Binary file not shown.
Binary file not shown.
@@ -44,6 +44,7 @@ class Text2SQLOrchestrator:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
schema_manager: SchemaManager,
|
schema_manager: SchemaManager,
|
||||||
|
llm_client: Optional[object] = None,
|
||||||
deepseek_api_key: Optional[str] = None,
|
deepseek_api_key: Optional[str] = None,
|
||||||
deepseek_config: Optional[DeepSeekConfig] = None,
|
deepseek_config: Optional[DeepSeekConfig] = None,
|
||||||
vector_db_path: str = "./data/embeddings/chroma",
|
vector_db_path: str = "./data/embeddings/chroma",
|
||||||
@@ -61,6 +62,7 @@ class Text2SQLOrchestrator:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
schema_manager: Schema管理器实例
|
schema_manager: Schema管理器实例
|
||||||
|
llm_client: 可选:外部传入的 LLM Client(需具备 chat/chat_with_json 等方法)。
|
||||||
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
|
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
|
||||||
deepseek_config: DeepSeek配置对象(优先于api_key)
|
deepseek_config: DeepSeek配置对象(优先于api_key)
|
||||||
vector_db_path: 向量数据库路径
|
vector_db_path: 向量数据库路径
|
||||||
@@ -75,13 +77,14 @@ class Text2SQLOrchestrator:
|
|||||||
self.use_vector_search = use_vector_search
|
self.use_vector_search = use_vector_search
|
||||||
self.translate_english_to_zh = translate_english_to_zh
|
self.translate_english_to_zh = translate_english_to_zh
|
||||||
|
|
||||||
# 初始化DeepSeek客户端
|
# 初始化 LLM 客户端(历史属性名保留为 deepseek,避免大范围改动)
|
||||||
|
if llm_client is not None:
|
||||||
|
self.deepseek = llm_client
|
||||||
|
else:
|
||||||
if deepseek_config:
|
if deepseek_config:
|
||||||
self.deepseek = DeepSeekClient(deepseek_config)
|
self.deepseek = DeepSeekClient(deepseek_config)
|
||||||
else:
|
else:
|
||||||
self.deepseek = DeepSeekClient(
|
self.deepseek = DeepSeekClient(DeepSeekConfig(api_key=deepseek_api_key))
|
||||||
DeepSeekConfig(api_key=deepseek_api_key)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 初始化向量索引(延迟加载)
|
# 初始化向量索引(延迟加载)
|
||||||
self._vector_index: Optional[SchemaIndexer] = None
|
self._vector_index: Optional[SchemaIndexer] = None
|
||||||
|
|||||||
Binary file not shown.
@@ -270,7 +270,8 @@ VALIDATOR_USER = """需要验证的SQL:
|
|||||||
请输出验证结果JSON:"""
|
请输出验证结果JSON:"""
|
||||||
|
|
||||||
EMPTY_RESULT_FEEDBACK_SYSTEM = """你是数据分析助手。用户的自然语言问题已转成 SQL,且在目标库执行成功,但**当前结果在列与行上均为空**(无可用结果集)。
|
EMPTY_RESULT_FEEDBACK_SYSTEM = """你是数据分析助手。用户的自然语言问题已转成 SQL,且在目标库执行成功,但**当前结果在列与行上均为空**(无可用结果集)。
|
||||||
请用 2~5 句简洁中文说明可能原因(如条件过严、时间范围无数据、对象不存在等),并**友好引导用户补充**时间、筛选条件、业务对象等,便于下次提问更精确。
|
请用 2~5 句简洁说明可能原因(如条件过严、时间范围无数据、对象不存在等),并**友好引导用户补充**时间、筛选条件、业务对象等,便于下次提问更精确。
|
||||||
|
若用户问题主要为英文,请用英文输出;否则用中文输出。
|
||||||
不要编造 Schema 中不存在的表或字段;不要重复输出整段 SQL;不要输出 JSON 或 Markdown 代码块。"""
|
不要编造 Schema 中不存在的表或字段;不要重复输出整段 SQL;不要输出 JSON 或 Markdown 代码块。"""
|
||||||
|
|
||||||
EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
|
EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
|
||||||
@@ -288,9 +289,10 @@ EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
|
|||||||
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM = """你是证券/期货类业务库的 Text2SQL 助手。
|
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM = """你是证券/期货类业务库的 Text2SQL 助手。
|
||||||
用户的自然语言问题已转为 SQL,且在目标库**试执行成功且至少返回一行数据**。
|
用户的自然语言问题已转为 SQL,且在目标库**试执行成功且至少返回一行数据**。
|
||||||
|
|
||||||
请用 1~3 句简洁中文向用户说明:
|
请用 1~3 句简洁说明:
|
||||||
- 该 SQL 大致在查询或统计什么(业务语义);
|
- 该 SQL 大致在查询或统计什么(业务语义);
|
||||||
- 可提示用户可在下方查看完整 SQL 并自行执行或导出。
|
- 可提示用户可在下方查看完整 SQL 并自行执行或导出。
|
||||||
|
若用户问题主要为英文,请用英文输出;否则用中文输出。
|
||||||
|
|
||||||
不要编造具体数据值;不要逐列复述;不要输出 Markdown 代码块或 JSON。"""
|
不要编造具体数据值;不要逐列复述;不要输出 Markdown 代码块或 JSON。"""
|
||||||
|
|
||||||
|
|||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,340 @@
|
|||||||
|
"""
|
||||||
|
OpenAI API 客户端封装(同步/异步)
|
||||||
|
|
||||||
|
提供与 :mod:`llm.deepseek_client` 同形态的接口,便于在需要时切换 LLM 路由。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from openai import AsyncOpenAI, OpenAI # type: ignore[import-not-found]
|
||||||
|
from openai.types.chat import ( # type: ignore[import-not-found]
|
||||||
|
ChatCompletion,
|
||||||
|
ChatCompletionMessage,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OpenAIConfig:
|
||||||
|
"""OpenAI API 配置(支持 OpenAI 兼容网关)。"""
|
||||||
|
|
||||||
|
api_key: str
|
||||||
|
base_url: str = "https://api.openai.com/v1"
|
||||||
|
model_name: str = "gpt-4o-mini"
|
||||||
|
temperature: float = 0.3
|
||||||
|
max_tokens: int = 4096
|
||||||
|
top_p: float = 0.9
|
||||||
|
frequency_penalty: float = 0.0
|
||||||
|
presence_penalty: float = 0.0
|
||||||
|
timeout: float = 60.0
|
||||||
|
extra_headers: Optional[Dict[str, str]] = None
|
||||||
|
|
||||||
|
|
||||||
|
class OpenAIClient:
|
||||||
|
"""
|
||||||
|
OpenAI API 客户端(同步)。
|
||||||
|
|
||||||
|
说明:本项目其他模块仅要求具备 ``chat`` / ``chat_with_json`` 方法;
|
||||||
|
这里额外提供若干便捷方法,与 DeepSeekClient 对齐,便于复用同一套 prompt。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: OpenAIConfig):
|
||||||
|
self.config = config
|
||||||
|
self.client = OpenAI(
|
||||||
|
api_key=config.api_key,
|
||||||
|
base_url=config.base_url,
|
||||||
|
timeout=config.timeout,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"[OK] OpenAIClient初始化: model=%s, base_url=%s",
|
||||||
|
config.model_name,
|
||||||
|
config.base_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
def chat(self, messages: List[Dict[str, str]], **kwargs) -> ChatCompletionMessage:
|
||||||
|
max_tokens = kwargs.get("max_tokens", self.config.max_tokens)
|
||||||
|
max_completion_tokens = kwargs.get("max_completion_tokens", None)
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"model": self.config.model_name,
|
||||||
|
"messages": messages,
|
||||||
|
"temperature": kwargs.get("temperature", self.config.temperature),
|
||||||
|
"top_p": kwargs.get("top_p", self.config.top_p),
|
||||||
|
"frequency_penalty": kwargs.get(
|
||||||
|
"frequency_penalty", self.config.frequency_penalty
|
||||||
|
),
|
||||||
|
"presence_penalty": kwargs.get(
|
||||||
|
"presence_penalty", self.config.presence_penalty
|
||||||
|
),
|
||||||
|
"timeout": kwargs.get("timeout", self.config.timeout),
|
||||||
|
}
|
||||||
|
|
||||||
|
# 兼容不同模型/网关:
|
||||||
|
# - 传统 Chat Completions 使用 max_tokens
|
||||||
|
# - 部分新模型(如 gpt-5.*)要求 max_completion_tokens
|
||||||
|
if max_completion_tokens is not None:
|
||||||
|
params["max_completion_tokens"] = max_completion_tokens
|
||||||
|
else:
|
||||||
|
params["max_tokens"] = max_tokens
|
||||||
|
|
||||||
|
if self.config.extra_headers:
|
||||||
|
params["extra_headers"] = self.config.extra_headers
|
||||||
|
|
||||||
|
try:
|
||||||
|
response: ChatCompletion = self.client.chat.completions.create(**params)
|
||||||
|
message = response.choices[0].message
|
||||||
|
usage = response.usage
|
||||||
|
if usage is not None:
|
||||||
|
logger.debug(
|
||||||
|
"OpenAI调用完成: prompt_tokens=%s, completion_tokens=%s, total_tokens=%s",
|
||||||
|
usage.prompt_tokens,
|
||||||
|
usage.completion_tokens,
|
||||||
|
usage.total_tokens,
|
||||||
|
)
|
||||||
|
return message
|
||||||
|
except Exception as e:
|
||||||
|
# 自动兼容:若网关提示 max_tokens 不支持,则改用 max_completion_tokens 重试一次
|
||||||
|
msg = str(e)
|
||||||
|
if (
|
||||||
|
"Unsupported parameter" in msg
|
||||||
|
and "max_tokens" in msg
|
||||||
|
and "max_completion_tokens" in msg
|
||||||
|
and "max_completion_tokens" not in params
|
||||||
|
):
|
||||||
|
params.pop("max_tokens", None)
|
||||||
|
params["max_completion_tokens"] = max_tokens
|
||||||
|
try:
|
||||||
|
response = self.client.chat.completions.create(**params)
|
||||||
|
return response.choices[0].message
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
logger.error("OpenAI API调用失败: %s", e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
def chat_with_json(
|
||||||
|
self, messages: List[Dict[str, str]], **kwargs
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
message = self.chat(messages, **kwargs)
|
||||||
|
content = (message.content or "").strip()
|
||||||
|
|
||||||
|
try:
|
||||||
|
if "```json" in content:
|
||||||
|
start = content.find("```json") + 7
|
||||||
|
end = content.find("```", start)
|
||||||
|
content = content[start:end].strip()
|
||||||
|
elif "```" in content:
|
||||||
|
start = content.find("```") + 3
|
||||||
|
end = content.find("```", start)
|
||||||
|
content = content[start:end].strip()
|
||||||
|
return json.loads(content)
|
||||||
|
except json.JSONDecodeError as e:
|
||||||
|
logger.warning("JSON解析失败,返回原始内容: %s", e)
|
||||||
|
return {"_json_decode_failed": True, "raw_content": content}
|
||||||
|
|
||||||
|
# ===== 便捷方法(与 DeepSeekClient 对齐)=====
|
||||||
|
def generate_sql(
|
||||||
|
self, prompt: str, schema: str, dialect: str = "tsql", **kwargs
|
||||||
|
) -> str:
|
||||||
|
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": SQL_GENERATOR_USER.format(
|
||||||
|
schema=schema, question=prompt, dialect=dialect
|
||||||
|
),
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
kwargs.setdefault("temperature", 0.0)
|
||||||
|
kwargs.setdefault("top_p", 1.0)
|
||||||
|
response = self.chat(messages, **kwargs)
|
||||||
|
content = (response.content or "").strip()
|
||||||
|
|
||||||
|
if "```sql" in content:
|
||||||
|
start = content.find("```sql") + 6
|
||||||
|
end = content.find("```", start)
|
||||||
|
content = content[start:end].strip()
|
||||||
|
elif "```" in content:
|
||||||
|
start = content.find("```") + 3
|
||||||
|
end = content.find("```", start)
|
||||||
|
content = content[start:end].strip()
|
||||||
|
|
||||||
|
return content
|
||||||
|
|
||||||
|
def validate_sql(self, sql: str, schema: str, **kwargs) -> Dict[str, Any]:
|
||||||
|
from config.prompts import VALIDATOR_SYSTEM, VALIDATOR_USER
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": VALIDATOR_SYSTEM},
|
||||||
|
{"role": "user", "content": VALIDATOR_USER.format(sql=sql, schema=schema)},
|
||||||
|
]
|
||||||
|
|
||||||
|
kwargs.setdefault("temperature", 0.0)
|
||||||
|
kwargs.setdefault("top_p", 1.0)
|
||||||
|
return self.chat_with_json(messages, **kwargs)
|
||||||
|
|
||||||
|
def select_tables(
|
||||||
|
self, question: str, table_list: str, **kwargs
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
from config.prompts import SCHEMA_LINKER_SYSTEM, SCHEMA_LINKER_USER
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": SCHEMA_LINKER_SYSTEM},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": SCHEMA_LINKER_USER.format(
|
||||||
|
question=question, table_list=table_list
|
||||||
|
),
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
kwargs.setdefault("temperature", 0.0)
|
||||||
|
kwargs.setdefault("top_p", 1.0)
|
||||||
|
return self.chat_with_json(messages, **kwargs)
|
||||||
|
|
||||||
|
def normalize_nl_question_for_text2sql(self, question: str) -> str:
|
||||||
|
from config.prompts import (
|
||||||
|
CANONICALIZE_NL_FOR_TEXT2SQL_SYSTEM,
|
||||||
|
CANONICALIZE_NL_FOR_TEXT2SQL_USER,
|
||||||
|
)
|
||||||
|
|
||||||
|
q = (question or "").strip()
|
||||||
|
if not q:
|
||||||
|
return ""
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": CANONICALIZE_NL_FOR_TEXT2SQL_SYSTEM},
|
||||||
|
{"role": "user", "content": CANONICALIZE_NL_FOR_TEXT2SQL_USER.format(question=q)},
|
||||||
|
]
|
||||||
|
msg = self.chat(messages, temperature=0.0, top_p=1.0, max_completion_tokens=512)
|
||||||
|
text = (msg.content or "").strip()
|
||||||
|
line = text.splitlines()[0].strip() if text else ""
|
||||||
|
return line.strip("「」\"'“”")
|
||||||
|
|
||||||
|
def translate_nl_question_to_zh(self, question: str) -> str:
|
||||||
|
return self.normalize_nl_question_for_text2sql(question)
|
||||||
|
|
||||||
|
def empty_result_user_feedback(
|
||||||
|
self,
|
||||||
|
question: str,
|
||||||
|
sql: str,
|
||||||
|
schema: str,
|
||||||
|
*,
|
||||||
|
max_schema_chars: int = 8000,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> str:
|
||||||
|
from config.prompts import EMPTY_RESULT_FEEDBACK_SYSTEM, EMPTY_RESULT_FEEDBACK_USER
|
||||||
|
|
||||||
|
schema_snip = (schema or "")[:max_schema_chars]
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": EMPTY_RESULT_FEEDBACK_SYSTEM},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": EMPTY_RESULT_FEEDBACK_USER.format(
|
||||||
|
question=question or "(无)",
|
||||||
|
sql=sql,
|
||||||
|
schema=schema_snip,
|
||||||
|
),
|
||||||
|
},
|
||||||
|
]
|
||||||
|
msg = self.chat(messages, temperature=0.4, max_completion_tokens=512, **kwargs)
|
||||||
|
return ((msg.content or "").strip())
|
||||||
|
|
||||||
|
def sql_probe_success_delivery_message(
|
||||||
|
self,
|
||||||
|
question: str,
|
||||||
|
sql: str,
|
||||||
|
*,
|
||||||
|
max_sql_chars: int = 4000,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> str:
|
||||||
|
from config.prompts import (
|
||||||
|
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM,
|
||||||
|
SQL_PROBE_SUCCESS_DELIVERY_USER,
|
||||||
|
)
|
||||||
|
|
||||||
|
q = (question or "").strip() or "(无)"
|
||||||
|
s = (sql or "").strip()
|
||||||
|
if len(s) > max_sql_chars:
|
||||||
|
s = s[: max_sql_chars - 20].rstrip() + "\n-- …(已截断)"
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": SQL_PROBE_SUCCESS_DELIVERY_SYSTEM},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": SQL_PROBE_SUCCESS_DELIVERY_USER.format(question=q, sql=s),
|
||||||
|
},
|
||||||
|
]
|
||||||
|
kwargs.setdefault("temperature", 0.2)
|
||||||
|
kwargs.setdefault("top_p", 1.0)
|
||||||
|
kwargs.setdefault("max_completion_tokens", 320)
|
||||||
|
msg = self.chat(messages, **kwargs)
|
||||||
|
return (msg.content or "").strip()
|
||||||
|
|
||||||
|
|
||||||
|
class AsyncOpenAIClient:
|
||||||
|
"""OpenAI API 客户端(异步)。"""
|
||||||
|
|
||||||
|
def __init__(self, config: OpenAIConfig):
|
||||||
|
self.config = config
|
||||||
|
self.client = AsyncOpenAI(
|
||||||
|
api_key=config.api_key,
|
||||||
|
base_url=config.base_url,
|
||||||
|
timeout=config.timeout,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def chat(self, messages: List[Dict[str, str]], **kwargs) -> ChatCompletionMessage:
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"model": self.config.model_name,
|
||||||
|
"messages": messages,
|
||||||
|
"temperature": kwargs.get("temperature", self.config.temperature),
|
||||||
|
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
|
||||||
|
"top_p": kwargs.get("top_p", self.config.top_p),
|
||||||
|
}
|
||||||
|
response = await self.client.chat.completions.create(**params)
|
||||||
|
return response.choices[0].message
|
||||||
|
|
||||||
|
|
||||||
|
def create_openai_client(
|
||||||
|
api_key: Optional[str] = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> OpenAIClient:
|
||||||
|
"""
|
||||||
|
便捷工厂:创建 OpenAIClient。
|
||||||
|
|
||||||
|
读取环境变量:
|
||||||
|
- OPENAI_API_KEY(必填)
|
||||||
|
- OPENAI_BASE_URL(可选)
|
||||||
|
- OPENAI_MODEL / OPENAI_CHAT_MODEL(可选)
|
||||||
|
"""
|
||||||
|
if api_key is None:
|
||||||
|
api_key = (os.getenv("OPENAI_API_KEY") or "").strip()
|
||||||
|
if not api_key:
|
||||||
|
raise ValueError("未提供 api_key 且环境变量 OPENAI_API_KEY 未设置。")
|
||||||
|
|
||||||
|
base_url = (kwargs.pop("base_url", None) or os.getenv("OPENAI_BASE_URL") or "").strip()
|
||||||
|
model_name = (
|
||||||
|
kwargs.pop("model_name", None)
|
||||||
|
or os.getenv("OPENAI_MODEL")
|
||||||
|
or os.getenv("OPENAI_CHAT_MODEL")
|
||||||
|
or ""
|
||||||
|
).strip()
|
||||||
|
|
||||||
|
cfg_kwargs: Dict[str, Any] = dict(kwargs)
|
||||||
|
if base_url:
|
||||||
|
cfg_kwargs["base_url"] = base_url
|
||||||
|
if model_name:
|
||||||
|
cfg_kwargs["model_name"] = model_name
|
||||||
|
|
||||||
|
config = OpenAIConfig(api_key=api_key, **cfg_kwargs)
|
||||||
|
return OpenAIClient(config)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
"""
|
||||||
|
LLM Client 路由器:在 DeepSeek / OpenAI 之间切换。
|
||||||
|
|
||||||
|
约定:
|
||||||
|
- 调用方只依赖 duck-typing:需要 ``chat`` / ``chat_with_json`` 等方法。
|
||||||
|
- 通过环境变量 ``LLM_SERVICE_CODE``(openai|deepseek)决定默认路由;
|
||||||
|
若未设置则按 Key 存在性自动选择(优先 deepseek)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
|
from llm.deepseek_client import DeepSeekClient, DeepSeekConfig
|
||||||
|
from llm.openai_client import create_openai_client
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_llm_service_code(service_code: Optional[str] = None) -> str:
|
||||||
|
sc = (service_code or os.getenv("LLM_SERVICE_CODE") or "").strip().lower()
|
||||||
|
if sc in ("openai", "deepseek"):
|
||||||
|
return sc
|
||||||
|
|
||||||
|
# 自动选择:优先 DeepSeek(与历史默认一致)
|
||||||
|
if (os.getenv("DEEPSEEK_API_KEY") or "").strip():
|
||||||
|
return "deepseek"
|
||||||
|
if (os.getenv("OPENAI_API_KEY") or "").strip():
|
||||||
|
return "openai"
|
||||||
|
return "deepseek"
|
||||||
|
|
||||||
|
|
||||||
|
def create_llm_client(service_code: Optional[str] = None, **kwargs: Any) -> Any:
|
||||||
|
"""
|
||||||
|
创建 LLM Client。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
DeepSeekClient 或 OpenAIClient(同形态接口)。
|
||||||
|
"""
|
||||||
|
sc = resolve_llm_service_code(service_code)
|
||||||
|
if sc == "openai":
|
||||||
|
# OpenAI 侧:默认从 OPENAI_* 读取
|
||||||
|
return create_openai_client(**kwargs)
|
||||||
|
|
||||||
|
# DeepSeek 侧:从 DEEPSEEK_* 读取
|
||||||
|
api_key = (kwargs.pop("api_key", None) or os.getenv("DEEPSEEK_API_KEY") or "").strip()
|
||||||
|
if not api_key:
|
||||||
|
raise ValueError("未配置 DEEPSEEK_API_KEY(LLM_SERVICE_CODE=deepseek)")
|
||||||
|
base_url = (kwargs.pop("base_url", None) or os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
|
||||||
|
model = (kwargs.pop("model_name", None) or os.getenv("MODEL_PRIMARY") or "deepseek-chat").strip()
|
||||||
|
cfg = DeepSeekConfig(api_key=api_key, base_url=base_url, model_name=model, **kwargs)
|
||||||
|
return DeepSeekClient(cfg)
|
||||||
|
|
||||||
+29
-8
@@ -108,16 +108,27 @@ def setup_environment():
|
|||||||
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
|
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# 检查API Key
|
# 检查 LLM Key(DeepSeek / OpenAI 可切换)
|
||||||
api_key = os.getenv("DEEPSEEK_API_KEY", "").strip()
|
llm_sc = (os.getenv("LLM_SERVICE_CODE") or "").strip().lower()
|
||||||
if not api_key:
|
if llm_sc and llm_sc not in ("deepseek", "openai"):
|
||||||
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置")
|
logger.warning("未知 LLM_SERVICE_CODE=%r(仅支持 deepseek/openai)", llm_sc)
|
||||||
|
return False
|
||||||
|
if llm_sc == "openai":
|
||||||
|
if not (os.getenv("OPENAI_API_KEY") or "").strip():
|
||||||
|
logger.warning("LLM_SERVICE_CODE=openai 但 OPENAI_API_KEY 未设置")
|
||||||
|
return False
|
||||||
|
else:
|
||||||
|
if not (os.getenv("DEEPSEEK_API_KEY") or "").strip():
|
||||||
|
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置(默认 LLM_SERVICE_CODE=deepseek)")
|
||||||
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
|
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
logger.info(f"[OK] 环境检查通过")
|
logger.info(f"[OK] 环境检查通过")
|
||||||
logger.info(f" - Schema: {schema_path}")
|
logger.info(f" - Schema: {schema_path}")
|
||||||
logger.info(f" - API Key: {'已配置' if api_key else '未配置'}")
|
logger.info(
|
||||||
|
" - LLM: %s",
|
||||||
|
(llm_sc or "deepseek(auto)"),
|
||||||
|
)
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -160,6 +171,7 @@ def load_schema(schema_path: str, schema_meta_path: Optional[str] = None):
|
|||||||
def create_orchestrator(schema_mgr, args):
|
def create_orchestrator(schema_mgr, args):
|
||||||
"""创建编排器"""
|
"""创建编排器"""
|
||||||
from agents.orchestrator import Text2SQLOrchestrator
|
from agents.orchestrator import Text2SQLOrchestrator
|
||||||
|
from llm.router import create_llm_client, resolve_llm_service_code
|
||||||
from llm.deepseek_client import DeepSeekConfig
|
from llm.deepseek_client import DeepSeekConfig
|
||||||
|
|
||||||
translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() not in (
|
translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() not in (
|
||||||
@@ -171,6 +183,8 @@ def create_orchestrator(schema_mgr, args):
|
|||||||
if getattr(args, "no_translate_en", False):
|
if getattr(args, "no_translate_en", False):
|
||||||
translate_en = False
|
translate_en = False
|
||||||
|
|
||||||
|
sc = resolve_llm_service_code()
|
||||||
|
if sc == "deepseek":
|
||||||
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
|
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
|
||||||
if not api_key:
|
if not api_key:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -178,18 +192,25 @@ def create_orchestrator(schema_mgr, args):
|
|||||||
"或使用命令行参数 --api-key"
|
"或使用命令行参数 --api-key"
|
||||||
)
|
)
|
||||||
base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
|
base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
|
||||||
|
cfg = DeepSeekConfig(
|
||||||
config = DeepSeekConfig(
|
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
model_name=args.model or "deepseek-chat",
|
model_name=args.model or "deepseek-chat",
|
||||||
temperature=args.temperature,
|
temperature=args.temperature,
|
||||||
max_tokens=args.max_tokens,
|
max_tokens=args.max_tokens,
|
||||||
)
|
)
|
||||||
|
llm_client = create_llm_client("deepseek", **cfg.__dict__)
|
||||||
|
else:
|
||||||
|
# openai:完全由 OPENAI_* 决定;同时沿用 temperature/max_tokens 作为默认值覆盖
|
||||||
|
llm_client = create_llm_client(
|
||||||
|
"openai",
|
||||||
|
temperature=args.temperature,
|
||||||
|
max_tokens=args.max_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
orchestrator = Text2SQLOrchestrator(
|
orchestrator = Text2SQLOrchestrator(
|
||||||
schema_manager=schema_mgr,
|
schema_manager=schema_mgr,
|
||||||
deepseek_config=config,
|
llm_client=llm_client,
|
||||||
vector_db_path=args.vector_db,
|
vector_db_path=args.vector_db,
|
||||||
max_retry=args.max_retry,
|
max_retry=args.max_retry,
|
||||||
use_vector_search=not args.no_vector_search,
|
use_vector_search=not args.no_vector_search,
|
||||||
|
|||||||
Binary file not shown.
Reference in New Issue
Block a user