Files
ai-g3sb-backman2.0/llm/deepseek_client.py
T
2026-04-10 16:52:07 +08:00

305 lines
8.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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}")
return {"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
)
}
]
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)
}
]
return self.chat_with_json(messages, **kwargs)
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
)
}
]
return self.chat_with_json(messages, **kwargs)
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)