138 lines
4.3 KiB
Python
138 lines
4.3 KiB
Python
from openai import AsyncOpenAI
|
|
from typing import Any, Dict, List, Optional
|
|
from config.env import DeepSeekConfig
|
|
from utils.log_util import logger
|
|
|
|
|
|
class DeepSeekAPIError(Exception):
|
|
"""DeepSeek API调用异常"""
|
|
|
|
def __init__(self, status_code: int, message: str):
|
|
self.status_code = status_code
|
|
self.message = message
|
|
super().__init__(f"DeepSeek API error {status_code}: {message}")
|
|
|
|
|
|
class DeepSeekAPIClient:
|
|
"""DeepSeek API客户端 - 用于调用DeepSeek大语言模型
|
|
使用官方推荐的OpenAI SDK进行调用
|
|
"""
|
|
|
|
def __init__(self, base_url: str = None, api_key: str = None, timeout: float = 30.0):
|
|
self.base_url = base_url or DeepSeekConfig.DEEPSEEK_API_BASE
|
|
self.base_url = self.base_url.rstrip('/') if self.base_url else 'https://api.deepseek.com'
|
|
self.api_key = api_key or DeepSeekConfig.DEEPSEEK_API_KEY
|
|
self.timeout = timeout
|
|
self.model = DeepSeekConfig.DEEPSEEK_MODEL
|
|
|
|
# 初始化OpenAI客户端
|
|
if not self.api_key:
|
|
raise DeepSeekAPIError(401, 'DeepSeek API密钥未配置')
|
|
|
|
self.client = AsyncOpenAI(
|
|
api_key=self.api_key,
|
|
base_url=self.base_url,
|
|
timeout=self.timeout
|
|
)
|
|
|
|
async def chat_completion(
|
|
self,
|
|
messages: List[Dict[str, str]],
|
|
model: str = None,
|
|
temperature: float = 0.7,
|
|
max_tokens: int = 1024,
|
|
stream: bool = False,
|
|
) -> Dict[str, Any]:
|
|
"""调用DeepSeek聊天补全API
|
|
|
|
Args:
|
|
messages: 聊天消息列表,格式为 [{"role": "user", "content": "你的问题"}]
|
|
model: 模型名称,默认使用配置中的模型
|
|
temperature: 温度参数,控制生成内容的随机性
|
|
max_tokens: 最大生成token数
|
|
stream: 是否流式返回
|
|
|
|
Returns:
|
|
API返回的JSON响应
|
|
|
|
Raises:
|
|
DeepSeekAPIError: API调用失败时抛出
|
|
"""
|
|
model = model or self.model
|
|
|
|
try:
|
|
response = await self.client.chat.completions.create(
|
|
model=model,
|
|
messages=messages,
|
|
temperature=temperature,
|
|
max_tokens=max_tokens,
|
|
stream=stream
|
|
)
|
|
|
|
# 转换为字典格式返回,保持与原有接口兼容
|
|
if not stream:
|
|
return response.model_dump()
|
|
else:
|
|
return response # 流式响应直接返回
|
|
|
|
except Exception as exc:
|
|
logger.error(f"DeepSeek API调用失败: {str(exc)}")
|
|
raise DeepSeekAPIError(500, str(exc)) from exc
|
|
|
|
async def chat(
|
|
self,
|
|
question: str,
|
|
context: str = None,
|
|
model: str = None,
|
|
temperature: float = 0.7,
|
|
max_tokens: int = 1024,
|
|
) -> str:
|
|
"""简化的聊天接口
|
|
|
|
Args:
|
|
question: 用户问题
|
|
context: 上下文信息(可选)
|
|
model: 模型名称
|
|
temperature: 温度参数
|
|
max_tokens: 最大生成token数
|
|
|
|
Returns:
|
|
生成的回答
|
|
|
|
Raises:
|
|
DeepSeekAPIError: API调用失败时抛出
|
|
"""
|
|
# 构建消息列表
|
|
messages = []
|
|
|
|
# 如果有上下文,添加到系统消息中
|
|
if context:
|
|
messages.append({
|
|
"role": "system",
|
|
"content": f"基于以下上下文信息回答用户问题:\n{context}"
|
|
})
|
|
|
|
# 添加用户问题
|
|
messages.append({
|
|
"role": "user",
|
|
"content": question
|
|
})
|
|
|
|
# 调用API
|
|
response = await self.chat_completion(
|
|
messages=messages,
|
|
model=model,
|
|
temperature=temperature,
|
|
max_tokens=max_tokens,
|
|
stream=False
|
|
)
|
|
|
|
# 提取回答
|
|
answer = response.get('choices', [{}])[0].get('message', {}).get('content', '')
|
|
|
|
if not answer:
|
|
logger.error(f"DeepSeek API返回空回答: {response}")
|
|
raise DeepSeekAPIError(200, 'DeepSeek API返回空回答')
|
|
|
|
return answer.strip()
|