use crate::types::*;
use crate::tool_registry::ToolRegistry;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use anyhow::{Result, Context};
use serde_json;
use tracing::error;
use reqwest::Client as HttpClient;
use serde::Deserialize;
pub struct DeepseekClient {
http_client: HttpClient,
config: DeepseekClientConfig,
history: Arc<RwLock<Vec<serde_json::Value>>>,
initialized: Arc<RwLock<bool>>,
tool_registry: Arc<RwLock<Option<Arc<ToolRegistry>>>>,
max_retries: u32,
timeout_ms: u64,
debug_mode: bool,
verbose: bool,
}
impl DeepseekClient {
pub fn new(config: DeepseekClientConfig) -> Result<Self> {
if config.api_key.is_empty() {
error!("DeepSeek API key is required");
eprintln!("❌ DeepSeek API key is required");
eprintln!("请设置环境变量 DEEPSEEK_API_KEY 或创建 .env 文件");
eprintln!("示例: DEEPSEEK_API_KEY=your_api_key_here");
return Err(anyhow::anyhow!("DeepSeek API key is required"));
}
let http_client = HttpClient::builder()
.timeout(std::time::Duration::from_millis(config.timeout.unwrap_or(30000)))
.build()
.context("Failed to create HTTP client")?;
Ok(Self {
http_client,
config,
history: Arc::new(RwLock::new(Vec::new())),
initialized: Arc::new(RwLock::new(false)),
tool_registry: Arc::new(RwLock::new(None)),
max_retries: 3,
timeout_ms: 30000, debug_mode: false,
verbose: false,
})
}
pub async fn initialize(&mut self) -> Result<()> {
self.log("开始初始化...", "debug");
if self.config.api_key.is_empty() {
return Err(anyhow::anyhow!("DeepSeek API key is required"));
}
self.log("发现MCP工具...", "debug");
let tool_registry = Arc::new(ToolRegistry::new());
tool_registry.discover_all_mcp_tools(self.debug_mode).await?;
let tool_count = tool_registry.get_tool_count().await;
println!("✅ 发现 {} 个工具", tool_count);
let tool_names = tool_registry.get_tool_names().await;
if !tool_names.is_empty() {
println!("📋 可用工具列表:");
for (i, name) in tool_names.iter().enumerate() {
println!(" {}. {}", i + 1, name);
}
}
{
let mut registry = self.tool_registry.write().await;
*registry = Some(tool_registry);
}
{
let mut initialized = self.initialized.write().await;
*initialized = true;
}
self.log("初始化完成", "debug");
Ok(())
}
pub async fn is_initialized(&self) -> bool {
let initialized = self.initialized.read().await;
*initialized
}
pub async fn get_history(&self) -> Vec<serde_json::Value> {
let history = self.history.read().await;
history.clone()
}
pub async fn get_available_tools(&self) -> Vec<String> {
let registry = self.tool_registry.read().await;
if let Some(registry) = registry.as_ref() {
registry.get_tool_names().await
} else {
Vec::new()
}
}
pub async fn set_history(&self, history: Vec<serde_json::Value>, _options: Option<serde_json::Value>) {
let mut hist = self.history.write().await;
*hist = history;
}
pub async fn set_tool_registry(&self, tool_registry: Arc<ToolRegistry>) {
let mut registry = self.tool_registry.write().await;
*registry = Some(tool_registry);
}
pub async fn generate_content(
&self,
prompt: &str,
tools: Option<Vec<serde_json::Value>>,
model: Option<&str>,
conversation_history: Option<Vec<serde_json::Value>>,
disable_thinking_tools: bool,
) -> Result<DeepseekResponse> {
self.log("生成内容...", "debug");
if !self.is_initialized().await {
return Err(anyhow::anyhow!("DeepSeek client not initialized"));
}
let messages = self.build_messages(prompt, conversation_history).await;
let model_name = model.unwrap_or("deepseek-chat");
let filtered_tools = if disable_thinking_tools {
tools.map(|t| {
t.into_iter()
.filter(|tool| {
!tool.get("function")
.and_then(|f| f.get("name"))
.and_then(|n| n.as_str())
.map_or(false, |name| name.contains("sequentialthinking"))
})
.collect()
})
} else {
tools
};
if let Some(ref tools) = filtered_tools {
self.log(&format!("已禁用思考工具,剩余工具数量: {}", tools.len()), "debug");
}
let response = self.call_deepseek_api(&messages, &filtered_tools, model_name).await?;
let content = response.choices[0].message.content.clone().unwrap_or_default();
let tool_calls = response.choices[0].message.tool_calls.as_ref().map(|calls| {
calls.iter().map(|tc| {
if let Some(function) = &tc.function {
ToolCall {
id: tc.id.clone(),
call_type: ToolCallType::Function,
function: Some(FunctionCallInfo {
name: function.name.clone(),
arguments: function.arguments.clone(),
}),
custom: None,
}
} else {
ToolCall {
id: tc.id.clone(),
call_type: ToolCallType::Custom,
function: None,
custom: None,
}
}
}).collect()
});
Ok(DeepseekResponse {
content,
tool_calls,
})
}
pub async fn execute_tool_calls(&self, tool_calls: Vec<ToolCall>) -> Result<Vec<ToolExecutionResult>> {
let registry = self.tool_registry.read().await;
let registry = registry.as_ref().ok_or_else(|| anyhow::anyhow!("Tool registry not set"))?;
let mut results = Vec::new();
for tool_call in tool_calls {
if tool_call.call_type != ToolCallType::Function || tool_call.function.is_none() {
self.log(&format!("跳过非函数工具调用: {:?}", tool_call.call_type), "warn");
results.push(ToolExecutionResult {
tool_call_id: tool_call.id,
name: tool_call.function.as_ref().map_or("unknown".to_string(), |f| f.name.clone()),
content: format!("跳过非函数工具调用: {:?}", tool_call.call_type),
success: false,
error: Some(format!("不支持的工具调用类型: {:?}", tool_call.call_type)),
});
continue;
}
let function = tool_call.function.as_ref().unwrap();
let mut last_error: Option<String> = None;
let mut success = false;
let mut result_content = String::new();
for attempt in 1..=self.max_retries {
self.log(&format!("执行工具: {} (尝试 {}/{})", function.name, attempt, self.max_retries), "debug");
match registry.get_tool(&function.name).await {
Some(tool) => {
match serde_json::from_str::<HashMap<String, serde_json::Value>>(&function.arguments) {
Ok(args) => {
let timeout_duration = std::time::Duration::from_millis(self.timeout_ms);
match tokio::time::timeout(timeout_duration, tool.execute(args)).await {
Ok(result) => {
match result {
Ok(tool_result) => {
if let Some(llm_content) = &tool_result.llm_content {
if let Ok(parts) = serde_json::from_value::<Vec<Part>>(llm_content.clone()) {
for part in parts {
if let Some(text) = part.text {
result_content.push_str(&text);
result_content.push('\n');
}
}
} else if let Ok(text) = serde_json::from_value::<String>(llm_content.clone()) {
result_content = text;
}
}
if result_content.is_empty() {
result_content = tool_result.return_display.unwrap_or_default();
}
success = true;
println!("✅ {}", function.name);
break;
}
Err(e) => {
last_error = Some(e.to_string());
self.log(&format!("工具执行失败 (尝试 {}/{}): {} - {}", attempt, self.max_retries, function.name, e), "warn");
}
}
}
Err(_) => {
last_error = Some("工具执行超时".to_string());
self.log(&format!("工具执行超时 (尝试 {}/{})", attempt, self.max_retries), "warn");
}
}
}
Err(e) => {
last_error = Some(format!("参数解析失败: {}", e));
break;
}
}
}
None => {
last_error = Some(format!("Tool {} not found", function.name));
break;
}
}
if attempt < self.max_retries {
tokio::time::sleep(std::time::Duration::from_millis(1000 * attempt as u64)).await;
}
}
results.push(ToolExecutionResult {
tool_call_id: tool_call.id,
name: function.name.clone(),
content: if result_content.is_empty() {
if success { "工具执行成功".to_string() } else { "工具执行失败".to_string() }
} else {
result_content
},
success,
error: if success { None } else { last_error },
});
}
Ok(results)
}
pub async fn chat_with_tools(
&self,
prompt: &str,
model: Option<&str>,
max_iterations: usize,
) -> Result<String> {
self.log("开始对话...", "debug");
if !self.is_initialized().await {
return Err(anyhow::anyhow!("DeepseekClient not initialized. Call initialize() first."));
}
let mut current_prompt = prompt.to_string();
let mut iteration = 0;
let mut conversation_history: Vec<serde_json::Value> = Vec::new();
let mut thinking_count = 0; let max_thinking_count = 3; let mut use_simple_model = false;
conversation_history.push(serde_json::json!({
"role": "user",
"content": current_prompt
}));
while iteration < max_iterations {
iteration += 1;
self.log(&format!("第 {} 轮对话", iteration), "debug");
let tool_schemas = {
let registry = self.tool_registry.read().await;
if let Some(registry) = registry.as_ref() {
Some(registry.get_function_declarations().await)
} else {
None
}
};
match self.generate_content(¤t_prompt, tool_schemas, model, Some(conversation_history.clone()), use_simple_model).await {
Ok(response) => {
conversation_history.push(serde_json::json!({
"role": "assistant",
"content": response.content
}));
if let Some(tool_calls) = response.tool_calls {
if !tool_calls.is_empty() {
self.log(&format!("发现 {} 个工具调用", tool_calls.len()), "debug");
let thinking_tools = tool_calls.iter().filter(|call| {
call.function.as_ref()
.map_or(false, |f| f.name.contains("sequentialthinking"))
}).count();
if thinking_tools > 0 {
thinking_count += thinking_tools;
self.log(&format!("思考工具调用次数: {}/{}", thinking_count, max_thinking_count), "debug");
if thinking_count > max_thinking_count && !use_simple_model {
self.log("思考次数超限,切换到简单模式", "warn");
use_simple_model = true;
return Ok(format!("我已经思考了 {} 次,现在切换到简单模式。请直接告诉我您需要什么帮助,我会直接执行而不进行深度思考。", max_thinking_count));
}
}
let tool_results = self.execute_tool_calls(tool_calls).await?;
let failed_tools = tool_results.iter().filter(|result| !result.success).count();
let successful_tools = tool_results.iter().filter(|result| result.success).count();
let mut tool_results_content = String::from("\n\n工具执行结果:\n");
for result in &tool_results {
let status = if result.success { "✅" } else { "❌" };
let mut display_content = result.content.clone();
if display_content.len() > 200 {
let truncated = display_content.chars().take(200).collect::<String>();
display_content = format!("{}... (总长度: {} 字符)",
truncated, display_content.len());
}
tool_results_content.push_str(&format!("{} {}: {}\n", status, result.name, display_content));
}
if failed_tools > 0 {
self.log(&format!("{} 个工具执行失败,重新思考", failed_tools), "warn");
let is_memory_task = tool_results.iter().any(|tool| {
tool.name.contains("memory") || tool.name.contains("create_entities")
});
let retry_prompt = if is_memory_task {
"记忆任务已完成。即使工具执行遇到问题,任务目标已经达成。".to_string()
} else {
let failed_tools_info = tool_results.iter()
.filter(|t| !t.success)
.map(|t| format!("- {}: {}", t.name, t.error.as_deref().unwrap_or("未知错误")))
.collect::<Vec<_>>()
.join("\n");
let successful_tools_info = tool_results.iter()
.filter(|t| t.success)
.map(|t| format!("- {}: {}", t.name, t.content))
.collect::<Vec<_>>()
.join("\n");
format!("刚才的工具调用中有一些失败了。请分析失败的原因并尝试其他方法来解决用户的问题。
失败的工具:
{}
成功的工具:
{}
请重新思考并尝试其他方法。如果所有任务都已完成,请明确说明。", failed_tools_info, successful_tools_info)
};
let mut full_tool_results_content = String::from("\n\n工具执行结果:\n");
for result in &tool_results {
let status = if result.success { "✅" } else { "❌" };
full_tool_results_content.push_str(&format!("{} {}: {}\n", status, result.name, result.content));
}
conversation_history.push(serde_json::json!({
"role": "user",
"content": format!("{}\n{}", full_tool_results_content, retry_prompt)
}));
current_prompt = retry_prompt;
continue; } else {
self.log("所有工具执行成功", "debug");
let completion_indicators = [
"任务完成", "已完成", "完成", "结束", "完成所有",
"所有任务", "任务结束", "工作完成", "执行完毕"
];
let response_text = response.content.to_lowercase();
let is_task_complete = completion_indicators.iter().any(|indicator| {
response_text.contains(indicator)
});
if is_task_complete {
return Ok(response.content + &tool_results_content);
} else {
self.log("任务可能未完成,继续执行", "debug");
let mut full_tool_results_content = String::from("\n\n工具执行结果:\n");
for result in &tool_results {
let status = if result.success { "✅" } else { "❌" };
full_tool_results_content.push_str(&format!("{} {}: {}\n", status, result.name, result.content));
}
let continue_prompt = format!("请继续执行剩余的任务。如果所有任务都已完成,请明确说明\"任务完成\"。
当前已完成的工作:
{}
请继续执行或确认任务完成。", full_tool_results_content);
conversation_history.push(serde_json::json!({
"role": "user",
"content": continue_prompt
}));
current_prompt = continue_prompt;
continue;
}
}
}
} else {
self.log("没有工具调用", "debug");
return Ok(response.content);
}
}
Err(e) => {
self.log(&format!("第 {} 轮对话出错: {}", iteration, e), "error");
if iteration >= max_iterations - 1 {
return Ok(format!("抱歉,在处理您的请求时遇到了问题:{}。我已经尝试了 {} 次,建议您重新描述需求或检查网络连接。", e, max_iterations));
}
let retry_prompt = match e.to_string() {
ref error_str if error_str.contains("maximum context length") || error_str.contains("tokens") => {
self.log("检测到token限制错误,清理历史记录", "debug");
self.clear_old_history().await;
"刚才出现了上下文长度限制问题,我已经清理了历史记录。请重新描述您的需求。".to_string()
}
ref error_str if error_str.contains("Connection failed") || error_str.contains("timeout") => {
format!("刚才出现了连接错误:{}。请稍等片刻后重新尝试,或者尝试使用其他工具。", e)
}
ref error_str if error_str.contains("Invalid") || error_str.contains("400") => {
format!("刚才出现了参数错误:{}。请检查工具参数是否正确,或者尝试不同的方法。", e)
}
ref error_str if error_str.contains("empty array") || error_str.contains("tools") => {
format!("刚才出现了工具配置问题:{}。请尝试重新描述您的需求,或者使用其他可用的工具。", e)
}
_ => {
format!("刚才出现了错误:{}。请重新尝试解决用户的问题,可以考虑使用不同的方法或工具。", e)
}
};
current_prompt = retry_prompt;
conversation_history.push(serde_json::json!({
"role": "user",
"content": current_prompt.clone()
}));
tokio::time::sleep(std::time::Duration::from_millis(1000)).await;
}
}
}
Ok(format!("经过 {} 轮尝试,仍然无法完全解决您的问题。建议您:\n1. 重新描述您的需求,提供更详细的信息\n2. 检查网络连接是否正常\n3. 尝试将复杂任务分解为更小的步骤\n4. 或者稍后再试", max_iterations))
}
async fn call_deepseek_api(
&self,
messages: &[serde_json::Value],
tools: &Option<Vec<serde_json::Value>>,
model: &str,
) -> Result<DeepseekApiResponse> {
let request_body = serde_json::json!({
"model": model,
"messages": messages,
"tools": tools,
"tool_choice": if tools.is_some() { serde_json::Value::String("auto".to_string()) } else { serde_json::Value::Null },
"temperature": 0.1,
"max_tokens": 4096,
});
let response = self.http_client
.post(&format!("{}/chat/completions", self.config.base_url.as_deref().unwrap_or("https://api.deepseek.com/v1")))
.header("Authorization", format!("Bearer {}", self.config.api_key))
.header("Content-Type", "application/json")
.json(&request_body)
.send()
.await
.context("Failed to send request to DeepSeek API")?;
if !response.status().is_success() {
let error_text = response.text().await.unwrap_or_default();
return Err(anyhow::anyhow!("DeepSeek API error: {}", error_text));
}
let api_response: DeepseekApiResponse = response
.json()
.await
.context("Failed to parse DeepSeek API response")?;
Ok(api_response)
}
async fn build_messages(&self, prompt: &str, conversation_history: Option<Vec<serde_json::Value>>) -> Vec<serde_json::Value> {
let mut messages = Vec::new();
if let Some(history) = conversation_history {
messages.extend(history);
} else {
let limited_history = self.get_limited_history().await;
for content in limited_history {
if let Some(role) = content.get("role").and_then(|v| v.as_str()) {
if let Some(text) = content.get("content").and_then(|v| v.as_str()) {
messages.push(serde_json::json!({
"role": role,
"content": text
}));
}
}
}
}
messages.push(serde_json::json!({
"role": "user",
"content": prompt
}));
self.limit_message_tokens(&mut messages);
messages
}
async fn get_limited_history(&self) -> Vec<serde_json::Value> {
let history = self.history.read().await;
let max_history_length = 20; if history.len() <= max_history_length {
history.clone()
} else {
history[history.len() - max_history_length..].to_vec()
}
}
fn limit_message_tokens(&self, messages: &mut Vec<serde_json::Value>) {
const MAX_TOKENS: usize = 100000; let mut total_tokens = 0;
let mut limited_messages = Vec::new();
for message in messages.iter().rev() {
let message_tokens = self.estimate_tokens(
message.get("content").and_then(|v| v.as_str()).unwrap_or("")
);
if total_tokens + message_tokens > MAX_TOKENS {
break;
}
total_tokens += message_tokens;
limited_messages.insert(0, message.clone()); }
self.log(&format!("消息token数量: {}, 消息数量: {}", total_tokens, limited_messages.len()), "debug");
*messages = limited_messages;
}
fn estimate_tokens(&self, text: &str) -> usize {
(text.len() + 2) / 3
}
async fn clear_old_history(&self) {
const KEEP_MESSAGES: usize = 10;
let mut history = self.history.write().await;
if history.len() > KEEP_MESSAGES {
*history = history[history.len() - KEEP_MESSAGES..].to_vec();
self.log(&format!("已清理历史记录,保留最近 {} 条消息", KEEP_MESSAGES), "debug");
}
}
fn log(&self, message: &str, level: &str) {
if level == "debug" && !self.verbose { return; }
if level == "info" && !self.verbose { return; }
let prefix = match level {
"error" => "❌",
"warn" => "⚠️",
"debug" => "🔍",
_ => "ℹ️",
};
println!("{} [DeepseekClient] {}", prefix, message);
}
}
#[derive(Debug, Deserialize)]
struct DeepseekApiResponse {
choices: Vec<Choice>,
}
#[derive(Debug, Deserialize)]
struct Choice {
message: Message,
}
#[derive(Debug, Deserialize)]
struct Message {
content: Option<String>,
tool_calls: Option<Vec<ToolCallInfo>>,
}
#[derive(Debug, Deserialize)]
struct ToolCallInfo {
id: String,
#[serde(rename = "type")]
call_type: String,
function: Option<FunctionInfo>,
}
#[derive(Debug, Deserialize)]
struct FunctionInfo {
name: String,
arguments: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_deepseek_client_creation() {
let config = DeepseekClientConfig {
api_key: "test_key".to_string(),
base_url: Some("https://api.deepseek.com/v1".to_string()),
timeout: Some(30000),
max_retries: Some(3),
debug_mode: Some(false),
target_dir: Some(".".to_string()),
};
let client = DeepseekClient::new(config);
assert!(client.is_ok());
}
#[test]
fn test_estimate_tokens() {
let config = DeepseekClientConfig {
api_key: "test".to_string(),
base_url: None,
timeout: None,
max_retries: None,
debug_mode: None,
target_dir: None,
};
let client = DeepseekClient::new(config).unwrap();
assert_eq!(client.estimate_tokens("hello"), 2);
assert_eq!(client.estimate_tokens("你好世界"), 2);
}
}