0.1.1 暂存
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# llm 包初始化
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,361 @@
|
||||
"""
|
||||
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 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 translate_nl_question_to_zh(self, question: str) -> str:
|
||||
"""
|
||||
将主要为英文的自然语言分析问题译为中文,便于与中文 Schema 注释 / 向量索引对齐。
|
||||
"""
|
||||
from config.prompts import TRANSLATE_NL_TO_ZH_SYSTEM, TRANSLATE_NL_TO_ZH_USER
|
||||
|
||||
q = (question or "").strip()
|
||||
if not q:
|
||||
return ""
|
||||
messages = [
|
||||
{"role": "system", "content": TRANSLATE_NL_TO_ZH_SYSTEM},
|
||||
{"role": "user", "content": TRANSLATE_NL_TO_ZH_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("「」\"'“”")
|
||||
|
||||
|
||||
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)
|
||||
Reference in New Issue
Block a user