diff --git a/.env b/.env index 49104de..51a1bd1 100644 --- a/.env +++ b/.env @@ -2,12 +2,23 @@ # 复制为 .env 并填入实际值 # ========== 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 +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 # 确定性模式:每次生成相同结果 MAX_TOKENS=4096 # 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等 diff --git a/IMPACT_ANALYSIS.md b/IMPACT_ANALYSIS.md index 1540f67..d1866cb 100644 --- a/IMPACT_ANALYSIS.md +++ b/IMPACT_ANALYSIS.md @@ -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(防上下文污染) ## 1. 改动概览 @@ -409,3 +524,30 @@ ## 4. 配置 | `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`。 | diff --git a/TASK_SUMMARY.md b/TASK_SUMMARY.md index e8b7600..e00b69d 100644 --- a/TASK_SUMMARY.md +++ b/TASK_SUMMARY.md @@ -12,6 +12,10 @@ - **新话题**:不注入历史摘要(`max_pairs=0`) - **目的**:减少多话题情况下历史 SQL/条件回流导致的串话与错误 SQL。 +- **新增(本次补充)**:`backend/llm/openai_client.py` + - 提供 `OpenAIClient.chat` / `chat_with_json`(以及与 DeepSeekClient 对齐的若干便捷方法),用于接入 OpenAI 或 OpenAI 兼容网关。 + - `.env` 增补 `OPENAI_MODEL/OPENAI_CHAT_MODEL` 的注释示例(不影响现有配置)。 + ## 影响与风险 - **破坏性变更**:否(对外 API 不变) @@ -26,6 +30,22 @@ - 同一 session:先做话题A数据查询,再提一个完全不同话题B(无“再/按上面”等续问词),应不再引用 A 的表/过滤条件。 - 同一 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)。 diff --git a/__pycache__/api_server.cpython-312.pyc b/__pycache__/api_server.cpython-312.pyc index 00a519a..6600344 100644 Binary files a/__pycache__/api_server.cpython-312.pyc and b/__pycache__/api_server.cpython-312.pyc differ diff --git a/api_server.py b/api_server.py index 7ac816a..5b6c08e 100644 --- a/api_server.py +++ b/api_server.py @@ -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" diff --git a/backend/__pycache__/main.cpython-312.pyc b/backend/__pycache__/main.cpython-312.pyc index 7efc9fb..d6f118d 100644 Binary files a/backend/__pycache__/main.cpython-312.pyc and b/backend/__pycache__/main.cpython-312.pyc differ diff --git a/backend/agents/__pycache__/orchestrator.cpython-312.pyc b/backend/agents/__pycache__/orchestrator.cpython-312.pyc index 8b01279..276de29 100644 Binary files a/backend/agents/__pycache__/orchestrator.cpython-312.pyc and b/backend/agents/__pycache__/orchestrator.cpython-312.pyc differ diff --git a/backend/agents/orchestrator.py b/backend/agents/orchestrator.py index 79b548b..692ab7a 100644 --- a/backend/agents/orchestrator.py +++ b/backend/agents/orchestrator.py @@ -44,6 +44,7 @@ class Text2SQLOrchestrator: def __init__( self, schema_manager: SchemaManager, + llm_client: Optional[object] = None, deepseek_api_key: Optional[str] = None, deepseek_config: Optional[DeepSeekConfig] = None, vector_db_path: str = "./data/embeddings/chroma", @@ -61,6 +62,7 @@ class Text2SQLOrchestrator: Args: schema_manager: Schema管理器实例 + llm_client: 可选:外部传入的 LLM Client(需具备 chat/chat_with_json 等方法)。 deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY) deepseek_config: DeepSeek配置对象(优先于api_key) vector_db_path: 向量数据库路径 @@ -75,13 +77,14 @@ class Text2SQLOrchestrator: self.use_vector_search = use_vector_search self.translate_english_to_zh = translate_english_to_zh - # 初始化DeepSeek客户端 - if deepseek_config: - self.deepseek = DeepSeekClient(deepseek_config) + # 初始化 LLM 客户端(历史属性名保留为 deepseek,避免大范围改动) + if llm_client is not None: + self.deepseek = llm_client else: - self.deepseek = DeepSeekClient( - DeepSeekConfig(api_key=deepseek_api_key) - ) + if deepseek_config: + self.deepseek = DeepSeekClient(deepseek_config) + else: + self.deepseek = DeepSeekClient(DeepSeekConfig(api_key=deepseek_api_key)) # 初始化向量索引(延迟加载) self._vector_index: Optional[SchemaIndexer] = None diff --git a/backend/config/__pycache__/prompts.cpython-312.pyc b/backend/config/__pycache__/prompts.cpython-312.pyc index df7b127..8a4418c 100644 Binary files a/backend/config/__pycache__/prompts.cpython-312.pyc and b/backend/config/__pycache__/prompts.cpython-312.pyc differ diff --git a/backend/config/prompts.py b/backend/config/prompts.py index cc685d4..77e6982 100644 --- a/backend/config/prompts.py +++ b/backend/config/prompts.py @@ -270,7 +270,8 @@ VALIDATOR_USER = """需要验证的SQL: 请输出验证结果JSON:""" EMPTY_RESULT_FEEDBACK_SYSTEM = """你是数据分析助手。用户的自然语言问题已转成 SQL,且在目标库执行成功,但**当前结果在列与行上均为空**(无可用结果集)。 -请用 2~5 句简洁中文说明可能原因(如条件过严、时间范围无数据、对象不存在等),并**友好引导用户补充**时间、筛选条件、业务对象等,便于下次提问更精确。 +请用 2~5 句简洁说明可能原因(如条件过严、时间范围无数据、对象不存在等),并**友好引导用户补充**时间、筛选条件、业务对象等,便于下次提问更精确。 +若用户问题主要为英文,请用英文输出;否则用中文输出。 不要编造 Schema 中不存在的表或字段;不要重复输出整段 SQL;不要输出 JSON 或 Markdown 代码块。""" EMPTY_RESULT_FEEDBACK_USER = """用户原始问题: @@ -288,9 +289,10 @@ EMPTY_RESULT_FEEDBACK_USER = """用户原始问题: SQL_PROBE_SUCCESS_DELIVERY_SYSTEM = """你是证券/期货类业务库的 Text2SQL 助手。 用户的自然语言问题已转为 SQL,且在目标库**试执行成功且至少返回一行数据**。 -请用 1~3 句简洁中文向用户说明: +请用 1~3 句简洁说明: - 该 SQL 大致在查询或统计什么(业务语义); - 可提示用户可在下方查看完整 SQL 并自行执行或导出。 +若用户问题主要为英文,请用英文输出;否则用中文输出。 不要编造具体数据值;不要逐列复述;不要输出 Markdown 代码块或 JSON。""" diff --git a/backend/llm/__pycache__/deepseek_client.cpython-312.pyc b/backend/llm/__pycache__/deepseek_client.cpython-312.pyc index e418af6..8abd54e 100644 Binary files a/backend/llm/__pycache__/deepseek_client.cpython-312.pyc and b/backend/llm/__pycache__/deepseek_client.cpython-312.pyc differ diff --git a/backend/llm/__pycache__/openai_client.cpython-312.pyc b/backend/llm/__pycache__/openai_client.cpython-312.pyc new file mode 100644 index 0000000..91de5d6 Binary files /dev/null and b/backend/llm/__pycache__/openai_client.cpython-312.pyc differ diff --git a/backend/llm/__pycache__/router.cpython-312.pyc b/backend/llm/__pycache__/router.cpython-312.pyc new file mode 100644 index 0000000..4922461 Binary files /dev/null and b/backend/llm/__pycache__/router.cpython-312.pyc differ diff --git a/backend/llm/openai_client.py b/backend/llm/openai_client.py index e69de29..a2a94dd 100644 --- a/backend/llm/openai_client.py +++ b/backend/llm/openai_client.py @@ -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) + diff --git a/backend/llm/router.py b/backend/llm/router.py new file mode 100644 index 0000000..3e01d54 --- /dev/null +++ b/backend/llm/router.py @@ -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) + diff --git a/backend/main.py b/backend/main.py index 386f674..2cbecde 100644 --- a/backend/main.py +++ b/backend/main.py @@ -108,16 +108,27 @@ def setup_environment(): logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录") return False - # 检查API Key - api_key = os.getenv("DEEPSEEK_API_KEY", "").strip() - if not api_key: - logger.warning("环境变量 DEEPSEEK_API_KEY 未设置") - logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key") + # 检查 LLM Key(DeepSeek / OpenAI 可切换) + llm_sc = (os.getenv("LLM_SERVICE_CODE") or "").strip().lower() + if llm_sc and llm_sc not in ("deepseek", "openai"): + 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") + return False logger.info(f"[OK] 环境检查通过") 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 @@ -160,6 +171,7 @@ def load_schema(schema_path: str, schema_meta_path: Optional[str] = None): def create_orchestrator(schema_mgr, args): """创建编排器""" from agents.orchestrator import Text2SQLOrchestrator + from llm.router import create_llm_client, resolve_llm_service_code from llm.deepseek_client import DeepSeekConfig translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() not in ( @@ -171,25 +183,34 @@ def create_orchestrator(schema_mgr, args): if getattr(args, "no_translate_en", False): translate_en = False - api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip() - if not api_key: - raise ValueError( - "未配置 DeepSeek API Key:请在 .env 中设置 DEEPSEEK_API_KEY," - "或使用命令行参数 --api-key" + sc = resolve_llm_service_code() + if sc == "deepseek": + api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip() + if not api_key: + raise ValueError( + "未配置 DeepSeek API Key:请在 .env 中设置 DEEPSEEK_API_KEY," + "或使用命令行参数 --api-key" + ) + base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip() + cfg = DeepSeekConfig( + api_key=api_key, + base_url=base_url, + model_name=args.model or "deepseek-chat", + temperature=args.temperature, + 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, ) - base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip() - - config = DeepSeekConfig( - api_key=api_key, - base_url=base_url, - model_name=args.model or "deepseek-chat", - temperature=args.temperature, - max_tokens=args.max_tokens, - ) orchestrator = Text2SQLOrchestrator( schema_manager=schema_mgr, - deepseek_config=config, + llm_client=llm_client, vector_db_path=args.vector_db, max_retry=args.max_retry, use_vector_search=not args.no_vector_search, diff --git a/data/embeddings/chroma_fewshot/9756e6ea-6d51-4e24-a490-4a58f0a585b9/length.bin b/data/embeddings/chroma_fewshot/9756e6ea-6d51-4e24-a490-4a58f0a585b9/length.bin index e601c07..2d95673 100644 Binary files a/data/embeddings/chroma_fewshot/9756e6ea-6d51-4e24-a490-4a58f0a585b9/length.bin and b/data/embeddings/chroma_fewshot/9756e6ea-6d51-4e24-a490-4a58f0a585b9/length.bin differ