408 lines
12 KiB
Python
408 lines
12 KiB
Python
"""
|
||
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}")
|
||
# 勿仅用 raw_content 判失败:空串时下游 `not raw.get("raw_content")` 会误判为成功
|
||
return {"_json_decode_failed": True, "raw_content": content}
|
||
|
||
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
|
||
)
|
||
}
|
||
]
|
||
|
||
kwargs.setdefault("temperature", 0.0)
|
||
kwargs.setdefault("top_p", 1.0)
|
||
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)
|
||
}
|
||
]
|
||
|
||
kwargs.setdefault("temperature", 0.0)
|
||
kwargs.setdefault("top_p", 1.0)
|
||
return self.chat_with_json(messages, **kwargs)
|
||
|
||
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
|
||
|
||
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()
|
||
|
||
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,
|
||
table_list=table_list,
|
||
),
|
||
},
|
||
]
|
||
|
||
# 选表为结构化决策:默认 temperature=0,避免同一问题多次选不同表/SQL 上下文
|
||
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:
|
||
"""
|
||
将中文/英文问句归一为**一句**标准中文(temperature=0),使同一语义的中英表述
|
||
走同一套向量检索与 SQL 生成路径。
|
||
"""
|
||
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_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:
|
||
"""
|
||
兼容旧名:与 :meth:`normalize_nl_question_for_text2sql` 相同(不再仅限英文)。
|
||
"""
|
||
return self.normalize_nl_question_for_text2sql(question)
|
||
|
||
|
||
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)
|