Files

445 lines
14 KiB
Python
Raw Permalink Normal View History

2026-04-10 16:52:07 +08:00
"""
DeepSeek API 客户端封装
支持同步/异步调用,与CAMEL AI兼容
"""
import os
import json
import logging
from typing import Any, Dict, Iterator, List, Optional, Union
2026-04-10 16:52:07 +08:00
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_stream(
self,
messages: List[Dict[str, str]],
**kwargs: Any,
) -> Iterator[str]:
"""
流式聊天:按 completion 增量产出文本片段(与 chat 同参,固定 stream=True)。
"""
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),
"frequency_penalty": kwargs.get(
"frequency_penalty", self.config.frequency_penalty
),
"presence_penalty": kwargs.get(
"presence_penalty", self.config.presence_penalty
),
"stream": True,
"timeout": kwargs.get("timeout", self.config.timeout),
}
if self.config.extra_headers:
params["extra_headers"] = self.config.extra_headers
try:
stream = self.client.chat.completions.create(**params)
for chunk in stream:
if not chunk.choices:
continue
delta = chunk.choices[0].delta
if delta and getattr(delta, "content", None):
yield delta.content
except Exception as e:
logger.error(f"DeepSeek API流式调用失败: {e}")
raise
2026-04-10 16:52:07 +08:00
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
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)
def normalize_nl_question_for_text2sql(self, question: str) -> str:
2026-04-14 10:28:22 +08:00
"""
将中文/英文问句归一为**一句**标准中文(temperature=0),使同一语义的中英表述
走同一套向量检索与 SQL 生成路径。
2026-04-14 10:28:22 +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 = [
{"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("「」\"'“”")
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)