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:
Binary file not shown.
Binary file not shown.
@@ -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
|
||||
|
||||
Binary file not shown.
@@ -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。"""
|
||||
|
||||
|
||||
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)
|
||||
|
||||
+42
-21
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user