use crate::types::*;
use crate::deepseek_client::DeepseekClient;
use crate::tool_registry::ToolRegistry;
use crate::prompt_registry::PromptRegistry;
use crate::config::Config;
use crate::prompts::get_mcp_system_prompt;
use std::collections::HashMap;
use std::sync::Arc;
use anyhow::Result;
use tracing::{info, debug};
use serde_json;
use chrono::Utc;
#[derive(Clone)]
pub struct FileOperationAgent {
deepseek_client: Arc<DeepseekClient>,
debug_mode: bool,
config: Config,
tool_registry: Arc<ToolRegistry>,
prompt_registry: Arc<PromptRegistry>,
}
impl FileOperationAgent {
pub fn new(
config: Config,
tool_registry: Arc<ToolRegistry>,
prompt_registry: Arc<PromptRegistry>,
) -> Result<Self> {
let debug_mode = false;
let client_config = DeepseekClientConfig {
api_key: config.deepseek_api.api_key.clone(),
base_url: Some(config.deepseek_api.api_endpoint.clone()),
timeout: Some(config.models.request_timeout * 1000), max_retries: Some(config.models.max_retries as u32),
debug_mode: Some(debug_mode),
target_dir: Some(config.workspace_root.to_string_lossy().to_string()),
};
let deepseek_client = Arc::new(DeepseekClient::new(client_config)?);
Ok(Self {
deepseek_client,
debug_mode,
config,
tool_registry,
prompt_registry,
})
}
pub async fn new_initialized(
config: Config,
tool_registry: Arc<ToolRegistry>,
prompt_registry: Arc<PromptRegistry>,
) -> Result<Self> {
let mut agent = Self::new(config, tool_registry, prompt_registry)?;
agent.initialize().await?;
Ok(agent)
}
pub async fn initialize(&mut self) -> Result<()> {
if self.debug_mode {
info!("开始初始化智能体...");
}
if let Some(deepseek_client) = Arc::get_mut(&mut self.deepseek_client) {
deepseek_client.initialize().await?;
} else {
return Err(anyhow::anyhow!("无法获取DeepSeek客户端的可变引用"));
}
if self.debug_mode {
info!("智能体初始化完成");
}
Ok(())
}
pub async fn process_request(&self, user_input: &str) -> Result<AgentResponse> {
if self.debug_mode {
info!("处理请求: {}...", &user_input[..user_input.len().min(50)]);
}
let system_prompt = get_mcp_system_prompt(&self.config.workspace_root.to_string_lossy());
let memory_context = self.get_memory_context(user_input).await;
let full_prompt = format!("{}\n\n{}\n\n用户请求: \"{}\"", system_prompt, memory_context, user_input);
let response = self.deepseek_client.chat_with_tools(&full_prompt, None, 10).await?;
self.save_to_memory(user_input, &response).await;
Ok(AgentResponse {
message: response.clone(),
actions: vec![],
content: Some(response),
memory_updates: None, })
}
pub async fn get_conversation_history(&self) -> Vec<serde_json::Value> {
self.deepseek_client.get_history().await
}
pub async fn get_available_tools(&self) -> Vec<String> {
self.deepseek_client.get_available_tools().await
}
async fn get_memory_context(&self, user_input: &str) -> String {
if !self.debug_mode {
return String::new();
}
match self.search_memory(user_input).await {
Ok(search_response) => {
if search_response.contains("相关记忆") {
return format!("\n## 相关记忆上下文\n{}\n", search_response);
}
}
Err(error) => {
if self.debug_mode {
debug!("获取记忆上下文失败: {}", error);
}
}
}
String::new()
}
async fn search_memory(&self, query: &str) -> Result<String> {
let search_prompt = format!("请使用 mcp_memory_search_nodes 工具搜索与以下内容相关的记忆: \"{}\"", query);
self.deepseek_client.chat_with_tools(&search_prompt, None, 3).await
}
async fn save_to_memory(&self, user_input: &str, response: &str) {
if !self.should_save_to_memory(user_input, response) {
return;
}
match self.save_conversation_to_memory(user_input, response).await {
Ok(_) => {
if self.debug_mode {
info!("已保存对话到记忆中");
}
}
Err(error) => {
if self.debug_mode {
debug!("保存到记忆失败: {}", error);
}
}
}
}
async fn save_conversation_to_memory(&self, user_input: &str, response: &str) -> Result<()> {
let timestamp = Utc::now().to_rfc3339();
let save_prompt = format!(r#"请使用 mcp_memory_create_entities 工具保存以下对话到记忆中:
实体名称: 对话_{}
实体类型: conversation
观察记录:
- 用户输入: {}
- AI响应: {}...
- 时间: {}"#,
timestamp,
user_input,
&response.chars().take(200).collect::<String>(),
timestamp
);
self.deepseek_client.chat_with_tools(&save_prompt, None, 3).await?;
Ok(())
}
fn should_save_to_memory(&self, user_input: &str, response: &str) -> bool {
let important_keywords = [
"项目", "配置", "设置", "偏好", "习惯", "重要", "记住", "保存",
"项目结构", "文件路径", "工作目录", "开发环境"
];
let input_lower = user_input.to_lowercase();
let response_lower = response.to_lowercase();
important_keywords.iter().any(|keyword|
input_lower.contains(keyword) || response_lower.contains(keyword)
) || response.len() > 500 }
pub async fn get_status(&self) -> AgentStatus {
let is_initialized = self.deepseek_client.is_initialized().await;
let tool_count = self.get_available_tools().await.len();
let memory_count = 0;
AgentStatus {
is_initialized,
tool_count,
memory_count,
debug_mode: self.debug_mode,
}
}
pub async fn reset(&self) -> Result<()> {
self.deepseek_client.set_history(Vec::new(), None).await;
info!("智能体状态已重置");
Ok(())
}
pub fn set_debug_mode(&mut self, debug_mode: bool) {
self.debug_mode = debug_mode;
}
pub fn is_debug_mode(&self) -> bool {
self.debug_mode
}
}
#[derive(Debug, Clone)]
pub struct AgentStatus {
pub is_initialized: bool,
pub tool_count: usize,
pub memory_count: usize,
pub debug_mode: bool,
}
pub struct AgentBuilder {
config: Option<Config>,
tool_registry: Option<Arc<ToolRegistry>>,
prompt_registry: Option<Arc<PromptRegistry>>,
}
impl AgentBuilder {
pub fn new() -> Self {
Self {
config: None,
tool_registry: None,
prompt_registry: None,
}
}
pub fn config(mut self, config: Config) -> Self {
self.config = Some(config);
self
}
pub fn tool_registry(mut self, tool_registry: Arc<ToolRegistry>) -> Self {
self.tool_registry = Some(tool_registry);
self
}
pub fn prompt_registry(mut self, prompt_registry: Arc<PromptRegistry>) -> Self {
self.prompt_registry = Some(prompt_registry);
self
}
pub fn build(self) -> Result<FileOperationAgent> {
let config = self.config.ok_or_else(|| anyhow::anyhow!("Config not provided"))?;
let tool_registry = self.tool_registry.ok_or_else(|| anyhow::anyhow!("ToolRegistry not provided"))?;
let prompt_registry = self.prompt_registry.ok_or_else(|| anyhow::anyhow!("PromptRegistry not provided"))?;
FileOperationAgent::new(config, tool_registry, prompt_registry)
}
}
impl Default for AgentBuilder {
fn default() -> Self {
Self::new()
}
}
pub struct AgentFactory;
impl AgentFactory {
pub fn create_default() -> Result<FileOperationAgent> {
let config = Config::new()?;
let tool_registry = Arc::new(ToolRegistry::new());
let prompt_registry = Arc::new(PromptRegistry::new());
AgentBuilder::new()
.config(config)
.tool_registry(tool_registry)
.prompt_registry(prompt_registry)
.build()
}
pub fn create_debug() -> Result<FileOperationAgent> {
let config = Config::new()?;
let tool_registry = Arc::new(ToolRegistry::new());
let prompt_registry = Arc::new(PromptRegistry::new());
AgentBuilder::new()
.config(config)
.tool_registry(tool_registry)
.prompt_registry(prompt_registry)
.build()
}
pub fn from_env() -> Result<FileOperationAgent> {
Self::create_default()
}
pub fn create_high_performance() -> Result<FileOperationAgent> {
let mut config = Config::new()?;
config.models.request_timeout = 10; config.models.max_retries = 5; let tool_registry = Arc::new(ToolRegistry::new());
let prompt_registry = Arc::new(PromptRegistry::new());
AgentBuilder::new()
.config(config)
.tool_registry(tool_registry)
.prompt_registry(prompt_registry)
.build()
}
}
pub struct AgentManager {
agents: Arc<tokio::sync::RwLock<HashMap<String, Arc<FileOperationAgent>>>>,
}
impl AgentManager {
pub fn new() -> Self {
Self {
agents: Arc::new(tokio::sync::RwLock::new(HashMap::new())),
}
}
pub async fn add_agent(&self, name: String, agent: Arc<FileOperationAgent>) {
let mut agents = self.agents.write().await;
agents.insert(name, agent);
}
pub async fn get_agent(&self, name: &str) -> Option<Arc<FileOperationAgent>> {
let agents = self.agents.read().await;
agents.get(name).cloned()
}
pub async fn remove_agent(&self, name: &str) -> Option<Arc<FileOperationAgent>> {
let mut agents = self.agents.write().await;
agents.remove(name)
}
pub async fn list_agents(&self) -> Vec<String> {
let agents = self.agents.read().await;
agents.keys().cloned().collect()
}
pub async fn agent_count(&self) -> usize {
let agents = self.agents.read().await;
agents.len()
}
}
impl Default for AgentManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_agent_builder() {
}
#[test]
fn test_should_save_to_memory() {
let config = Config::default();
let tool_registry = Arc::new(ToolRegistry::new());
let prompt_registry = Arc::new(PromptRegistry::new());
let agent = FileOperationAgent::new(config, tool_registry, prompt_registry).unwrap();
assert!(agent.should_save_to_memory("这是一个重要的项目配置", ""));
assert!(agent.should_save_to_memory("", "这是项目结构信息"));
let long_response = "a".repeat(600);
assert!(agent.should_save_to_memory("普通问题", &long_response));
assert!(!agent.should_save_to_memory("你好", "你好!"));
}
#[tokio::test]
async fn test_agent_manager() {
let manager = AgentManager::new();
assert_eq!(manager.agent_count().await, 0);
let config = Config::default();
let tool_registry = Arc::new(ToolRegistry::new());
let prompt_registry = Arc::new(PromptRegistry::new());
let agent = Arc::new(FileOperationAgent::new(config, tool_registry, prompt_registry).unwrap());
manager.add_agent("test_agent".to_string(), agent).await;
assert_eq!(manager.agent_count().await, 1);
let retrieved = manager.get_agent("test_agent").await;
assert!(retrieved.is_some());
let removed = manager.remove_agent("test_agent").await;
assert!(removed.is_some());
assert_eq!(manager.agent_count().await, 0);
}
}