Files
ai-g3sb-backman2.0/backend/llm/deepseek_client.py
T

408 lines
12 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}")
# 勿仅用 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)