2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
DeepSeek API 客户端封装
|
|
|
|
|
|
支持同步/异步调用,与CAMEL AI兼容
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
import os
|
|
|
|
|
|
import json
|
|
|
|
|
|
import logging
|
|
|
|
|
|
from typing import Dict, List, Optional, Any, Union
|
|
|
|
|
|
from dataclasses import dataclass, field
|
|
|
|
|
|
from openai import OpenAI, AsyncOpenAI
|
|
|
|
|
|
from openai.types.chat import ChatCompletion, ChatCompletionMessage
|
|
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
|
class DeepSeekConfig:
|
|
|
|
|
|
"""DeepSeek API配置"""
|
|
|
|
|
|
api_key: str
|
|
|
|
|
|
base_url: str = "https://api.deepseek.com"
|
|
|
|
|
|
model_name: str = "deepseek-chat"
|
|
|
|
|
|
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 DeepSeekClient:
|
|
|
|
|
|
"""
|
|
|
|
|
|
DeepSeek API 客户端(同步)
|
|
|
|
|
|
|
|
|
|
|
|
使用 OpenAI 兼容接口调用 DeepSeek 模型
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(self, config: DeepSeekConfig):
|
|
|
|
|
|
self.config = config
|
|
|
|
|
|
|
|
|
|
|
|
# 初始化OpenAI客户端(DeepSeek兼容OpenAI协议)
|
|
|
|
|
|
self.client = OpenAI(
|
|
|
|
|
|
api_key=config.api_key,
|
|
|
|
|
|
base_url=config.base_url,
|
|
|
|
|
|
timeout=config.timeout,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
f"[OK] DeepSeekClient初始化: model={config.model_name}, "
|
|
|
|
|
|
f"base_url={config.base_url}"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def chat(
|
|
|
|
|
|
self,
|
|
|
|
|
|
messages: List[Dict[str, str]],
|
|
|
|
|
|
**kwargs
|
|
|
|
|
|
) -> ChatCompletionMessage:
|
|
|
|
|
|
"""
|
|
|
|
|
|
发送聊天请求
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
messages: 消息列表,格式为 [{"role": "user", "content": "..."}, ...]
|
|
|
|
|
|
**kwargs: 覆盖默认参数的额外参数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
ChatCompletionMessage对象(包含.role和.content属性)
|
|
|
|
|
|
"""
|
|
|
|
|
|
# 合并参数
|
|
|
|
|
|
params = {
|
|
|
|
|
|
"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),
|
|
|
|
|
|
"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),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
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
|
|
|
|
|
|
logger.debug(
|
|
|
|
|
|
f"DeepSeek调用完成: "
|
|
|
|
|
|
f"prompt_tokens={usage.prompt_tokens}, "
|
|
|
|
|
|
f"completion_tokens={usage.completion_tokens}, "
|
|
|
|
|
|
f"total_tokens={usage.total_tokens}"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
return message
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"DeepSeek API调用失败: {e}")
|
|
|
|
|
|
raise
|
|
|
|
|
|
|
|
|
|
|
|
def chat_with_json(
|
|
|
|
|
|
self,
|
|
|
|
|
|
messages: List[Dict[str, str]],
|
|
|
|
|
|
**kwargs
|
|
|
|
|
|
) -> Dict[str, Any]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
发送聊天请求并期望JSON格式返回
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
messages: 消息列表
|
|
|
|
|
|
**kwargs: 额外参数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
解析后的JSON字典
|
|
|
|
|
|
"""
|
|
|
|
|
|
message = self.chat(messages, **kwargs)
|
|
|
|
|
|
content = message.content.strip()
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
|
# 尝试解析JSON(可能包含markdown代码块)
|
|
|
|
|
|
if "```json" in content:
|
|
|
|
|
|
# 提取```json```之间的内容
|
|
|
|
|
|
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(f"JSON解析失败,返回原始内容: {e}")
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# 勿仅用 raw_content 判失败:空串时下游 `not raw.get("raw_content")` 会误判为成功
|
|
|
|
|
|
return {"_json_decode_failed": True, "raw_content": content}
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
def generate_sql(
|
|
|
|
|
|
self,
|
|
|
|
|
|
prompt: str,
|
|
|
|
|
|
schema: str,
|
|
|
|
|
|
dialect: str = "tsql",
|
|
|
|
|
|
**kwargs
|
|
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
便捷方法:生成SQL
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
prompt: 用户问题
|
|
|
|
|
|
schema: Schema描述
|
|
|
|
|
|
dialect: SQL方言(默认 T-SQL / SQL Server)
|
|
|
|
|
|
**kwargs: 额外参数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
SQL字符串
|
|
|
|
|
|
"""
|
|
|
|
|
|
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
|
|
|
|
|
|
)
|
|
|
|
|
|
}
|
|
|
|
|
|
]
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
kwargs.setdefault("temperature", 0.0)
|
|
|
|
|
|
kwargs.setdefault("top_p", 1.0)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
response = self.chat(messages, **kwargs)
|
|
|
|
|
|
content = response.content.strip()
|
|
|
|
|
|
|
|
|
|
|
|
# 提取SQL(移除可能的markdown代码块标记)
|
|
|
|
|
|
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]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
便捷方法:验证SQL
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
sql: SQL语句
|
|
|
|
|
|
schema: Schema描述
|
|
|
|
|
|
**kwargs: 额外参数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
验证结果字典
|
|
|
|
|
|
"""
|
|
|
|
|
|
from config.prompts import VALIDATOR_SYSTEM, VALIDATOR_USER
|
|
|
|
|
|
|
|
|
|
|
|
messages = [
|
|
|
|
|
|
{"role": "system", "content": VALIDATOR_SYSTEM},
|
|
|
|
|
|
{
|
|
|
|
|
|
"role": "user",
|
|
|
|
|
|
"content": VALIDATOR_USER.format(sql=sql, schema=schema)
|
|
|
|
|
|
}
|
|
|
|
|
|
]
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
kwargs.setdefault("temperature", 0.0)
|
|
|
|
|
|
kwargs.setdefault("top_p", 1.0)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return self.chat_with_json(messages, **kwargs)
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
def empty_result_user_feedback(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
sql: str,
|
|
|
|
|
|
schema: str,
|
|
|
|
|
|
*,
|
|
|
|
|
|
max_schema_chars: int = 8000,
|
|
|
|
|
|
**kwargs: Any,
|
|
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
库探针为 0(执行成功但结果行数为 0)时,生成面向用户的中文补充说明,引导用户完善问题。
|
|
|
|
|
|
"""
|
|
|
|
|
|
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_tokens=512, **kwargs)
|
|
|
|
|
|
text = (msg.content or "").strip()
|
|
|
|
|
|
return text
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
def sql_probe_success_delivery_message(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
sql: str,
|
|
|
|
|
|
*,
|
|
|
|
|
|
max_sql_chars: int = 4000,
|
|
|
|
|
|
**kwargs: Any,
|
|
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
库探针为 1(执行成功且至少一行数据)时,生成面向用户的简短交付说明。
|
|
|
|
|
|
"""
|
|
|
|
|
|
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_tokens", 320)
|
|
|
|
|
|
msg = self.chat(messages, **kwargs)
|
|
|
|
|
|
return (msg.content or "").strip()
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def select_tables(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
table_list: str,
|
|
|
|
|
|
**kwargs
|
|
|
|
|
|
) -> Dict[str, Any]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
便捷方法:选择相关表
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
question: 用户问题
|
|
|
|
|
|
table_list: 可用表列表字符串
|
|
|
|
|
|
**kwargs: 额外参数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
包含relevant_tables和reasoning的字典
|
|
|
|
|
|
"""
|
|
|
|
|
|
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,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
table_list=table_list,
|
|
|
|
|
|
),
|
|
|
|
|
|
},
|
2026-04-10 16:52:07 +08:00
|
|
|
|
]
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# 选表为结构化决策:默认 temperature=0,避免同一问题多次选不同表/SQL 上下文
|
|
|
|
|
|
kwargs.setdefault("temperature", 0.0)
|
|
|
|
|
|
kwargs.setdefault("top_p", 1.0)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return self.chat_with_json(messages, **kwargs)
|
|
|
|
|
|
|
2026-04-15 11:25:19 +08:00
|
|
|
|
def normalize_nl_question_for_text2sql(self, question: str) -> str:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
"""
|
2026-04-15 11:25:19 +08:00
|
|
|
|
将中文/英文问句归一为**一句**标准中文(temperature=0),使同一语义的中英表述
|
|
|
|
|
|
走同一套向量检索与 SQL 生成路径。
|
2026-04-14 10:28:22 +08:00
|
|
|
|
"""
|
2026-04-15 11:25:19 +08:00
|
|
|
|
from config.prompts import (
|
|
|
|
|
|
CANONICALIZE_NL_FOR_TEXT2SQL_SYSTEM,
|
|
|
|
|
|
CANONICALIZE_NL_FOR_TEXT2SQL_USER,
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
q = (question or "").strip()
|
|
|
|
|
|
if not q:
|
|
|
|
|
|
return ""
|
|
|
|
|
|
messages = [
|
2026-04-15 11:25:19 +08:00
|
|
|
|
{"role": "system", "content": CANONICALIZE_NL_FOR_TEXT2SQL_SYSTEM},
|
|
|
|
|
|
{"role": "user", "content": CANONICALIZE_NL_FOR_TEXT2SQL_USER.format(question=q)},
|
2026-04-14 10:28:22 +08:00
|
|
|
|
]
|
|
|
|
|
|
msg = self.chat(messages, temperature=0.0, top_p=1.0, max_tokens=512)
|
|
|
|
|
|
text = (msg.content or "").strip()
|
|
|
|
|
|
line = text.splitlines()[0].strip() if text else ""
|
|
|
|
|
|
return line.strip("「」\"'“”")
|
|
|
|
|
|
|
2026-04-15 11:25:19 +08:00
|
|
|
|
def translate_nl_question_to_zh(self, question: str) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
兼容旧名:与 :meth:`normalize_nl_question_for_text2sql` 相同(不再仅限英文)。
|
|
|
|
|
|
"""
|
|
|
|
|
|
return self.normalize_nl_question_for_text2sql(question)
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
class AsyncDeepSeekClient:
|
|
|
|
|
|
"""
|
|
|
|
|
|
DeepSeek API 客户端(异步)
|
|
|
|
|
|
|
|
|
|
|
|
注意:CAMEL AI当前版本主要支持同步Agent,
|
|
|
|
|
|
异步客户端适用于自定义异步流程。
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(self, config: DeepSeekConfig):
|
|
|
|
|
|
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 = {
|
|
|
|
|
|
"model": self.config.model_name,
|
|
|
|
|
|
"messages": messages,
|
|
|
|
|
|
"temperature": kwargs.get("temperature", self.config.temperature),
|
|
|
|
|
|
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
|
|
|
|
|
|
}
|
|
|
|
|
|
response = await self.client.chat.completions.create(**params)
|
|
|
|
|
|
return response.choices[0].message
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# 便捷工厂函数
|
|
|
|
|
|
def create_deepseek_client(
|
|
|
|
|
|
api_key: Optional[str] = None,
|
|
|
|
|
|
**kwargs
|
|
|
|
|
|
) -> DeepSeekClient:
|
|
|
|
|
|
"""
|
|
|
|
|
|
创建DeepSeek客户端
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
api_key: API密钥(可从环境变量DEEPSEEK_API_KEY读取)
|
|
|
|
|
|
**kwargs: 覆盖默认配置的参数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
DeepSeekClient实例
|
|
|
|
|
|
"""
|
|
|
|
|
|
if api_key is None:
|
|
|
|
|
|
api_key = os.getenv("DEEPSEEK_API_KEY")
|
|
|
|
|
|
if not api_key:
|
|
|
|
|
|
raise ValueError(
|
|
|
|
|
|
"未提供api_key且环境变量DEEPSEEK_API_KEY未设置。"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
config = DeepSeekConfig(api_key=api_key, **kwargs)
|
|
|
|
|
|
return DeepSeekClient(config)
|