use std::collections::HashMap; use std::sync::Arc; use async_trait::async_trait; use serde::Serialize; use serde_json::json; use tokio::sync::RwLock; use crate::storage::TodoRepository; use crate::tools::traits::{Tool, ToolContext, ToolResult}; use crate::tools::todo_write::{TodoItem, scope_key_from_context}; // ── 输出结构 ────────────────────────────────────────────── #[derive(Debug, Clone, Serialize)] struct TodoReadOutput { todos: Vec, count: usize, scope_key: String, source: &'static str, } // ── 工具实现 ────────────────────────────────────────────── pub struct TodoReadTool { /// 共享内存状态(与 TodoWriteTool 同一实例) state: Arc>>>, /// SQLite 持久化层,用于进程重启后回填内存 repository: Arc, } impl TodoReadTool { pub(crate) fn new( state: Arc>>>, repository: Arc, ) -> Self { Self { state, repository } } } #[async_trait] impl Tool for TodoReadTool { fn name(&self) -> &str { "todo_read" } fn description(&self) -> &str { "Read the current todo list for this conversation without modifying it. \ Returns all tracked tasks with their id, content, and status. \ No parameters required." } fn parameters_schema(&self) -> serde_json::Value { json!({ "type": "object", "properties": {}, "required": [] }) } fn read_only(&self) -> bool { true } fn concurrency_safe(&self) -> bool { true } async fn execute(&self, _args: serde_json::Value) -> anyhow::Result { Ok(error_result("todo_read requires tool context (session_id)")) } async fn execute_with_context( &self, context: &ToolContext, _args: serde_json::Value, ) -> anyhow::Result { // 1. 计算 scope_key let scope_key = match scope_key_from_context(context) { Some(key) => key, None => { return Ok(error_result( "todo_read requires session_id or topic_id in tool context", )) } }; // 2. 读锁查内存 { let guard = self.state.read().await; if let Some(items) = guard.get(&scope_key) { if !items.is_empty() { return Ok(success_result(items, &scope_key, "memory")); } } } // 3. 内存为空 → 查 SQLite 并回填 let records = match self.repository.list_todos(&scope_key) { Ok(records) => records, Err(e) => { tracing::warn!( error = %e, %scope_key, "TodoReadTool: failed to load todos from SQLite" ); return Ok(success_result(&[], &scope_key, "memory")); } }; if records.is_empty() { return Ok(success_result(&[], &scope_key, "sqlite")); } let items: Vec = records .into_iter() .map(|r| TodoItem { id: r.id, content: r.content, status: r.status, created_by_message_id: r.created_by_message_id, }) .collect(); // 回填内存 { let mut guard = self.state.write().await; guard.insert(scope_key.clone(), items.clone()); } tracing::info!( scope_key = %scope_key, todo_count = items.len(), "TodoReadTool: backfilled memory from SQLite" ); Ok(success_result(&items, &scope_key, "sqlite")) } } // ── 辅助函数 ────────────────────────────────────────────── fn success_result( items: &[TodoItem], scope_key: &str, source: &'static str, ) -> ToolResult { let output = TodoReadOutput { todos: items.to_vec(), count: items.len(), scope_key: scope_key.to_string(), source, }; ToolResult { success: true, output: serde_json::to_string_pretty(&output).unwrap_or_default(), error: None, } } fn error_result(message: &str) -> ToolResult { ToolResult { success: false, output: String::new(), error: Some(message.to_string()), } } // ── 测试 ────────────────────────────────────────────────── #[cfg(test)] mod tests { use super::*; use crate::tools::traits::ToolContext; fn test_context() -> ToolContext { ToolContext { channel_name: Some("cli".to_string()), sender_id: Some("user-1".to_string()), chat_id: Some("chat-1".to_string()), session_id: Some("cli:chat-1".to_string()), topic_id: None, message_id: Some("msg-1".to_string()), message_seq: Some(1), subagent_description: None, nesting_depth: 0, task_id: None, parent_task_id: None, tool_call_id: None, parent_capability: None, } } fn test_state() -> Arc>>> { Arc::new(RwLock::new(HashMap::new())) } struct MockTodoRepository { records: Vec, } impl TodoRepository for MockTodoRepository { fn replace_todos( &self, _scope_key: &str, _items: &[crate::storage::TodoRecord], ) -> Result, crate::storage::StorageError> { Ok(vec![]) } fn list_todos( &self, scope_key: &str, ) -> Result, crate::storage::StorageError> { Ok(self.records.iter().filter(|r| r.scope_key == scope_key).cloned().collect()) } } fn mock_record(scope_key: &str, id: &str, content: &str, status: &str) -> crate::storage::TodoRecord { crate::storage::TodoRecord { id: id.to_string(), scope_key: scope_key.to_string(), session_id: "cli:chat-1".to_string(), topic_id: None, content: content.to_string(), status: status.to_string(), priority: "medium".to_string(), created_at: 1000, updated_at: 1000, created_by_message_id: None, } } #[tokio::test] async fn test_read_from_memory() { let state = test_state(); { let mut guard = state.write().await; guard.insert( "cli:chat-1".to_string(), vec![TodoItem { id: "a1".to_string(), content: "任务A".to_string(), status: "pending".to_string(), created_by_message_id: None, }], ); } let repo = Arc::new(MockTodoRepository { records: vec![] }); let tool = TodoReadTool::new(state, repo); let result = tool.execute_with_context(&test_context(), json!({})).await.unwrap(); assert!(result.success); let output: serde_json::Value = serde_json::from_str(&result.output).unwrap(); assert_eq!(output["count"], 1); assert_eq!(output["source"], "memory"); assert_eq!(output["todos"][0]["id"], "a1"); } #[tokio::test] async fn test_read_from_sqlite_backfill() { let state = test_state(); let repo = Arc::new(MockTodoRepository { records: vec![ mock_record("cli:chat-1", "b1", "任务B", "in_progress"), mock_record("cli:chat-1", "b2", "任务C", "pending"), ], }); let tool = TodoReadTool::new(state.clone(), repo); let result = tool.execute_with_context(&test_context(), json!({})).await.unwrap(); assert!(result.success); let output: serde_json::Value = serde_json::from_str(&result.output).unwrap(); assert_eq!(output["count"], 2); assert_eq!(output["source"], "sqlite"); // 验证内存已被回填 let guard = state.read().await; let items = guard.get("cli:chat-1").unwrap(); assert_eq!(items.len(), 2); } #[tokio::test] async fn test_read_empty_list() { let state = test_state(); let repo = Arc::new(MockTodoRepository { records: vec![] }); let tool = TodoReadTool::new(state, repo); let result = tool.execute_with_context(&test_context(), json!({})).await.unwrap(); assert!(result.success); let output: serde_json::Value = serde_json::from_str(&result.output).unwrap(); assert_eq!(output["count"], 0); } #[tokio::test] async fn test_read_no_context() { let state = test_state(); let repo = Arc::new(MockTodoRepository { records: vec![] }); let tool = TodoReadTool::new(state, repo); let result = tool.execute(json!({})).await.unwrap(); assert!(!result.success); assert!(result.error.unwrap().contains("requires tool context")); } #[tokio::test] async fn test_read_topic_isolation() { let state = test_state(); { let mut guard = state.write().await; guard.insert( "cli:chat-1".to_string(), vec![TodoItem { id: "m1".to_string(), content: "主会话任务".to_string(), status: "pending".to_string(), created_by_message_id: None, }], ); } let repo = Arc::new(MockTodoRepository { records: vec![ mock_record("topic-xyz", "t1", "话题任务", "completed"), ], }); let tool = TodoReadTool::new(state, repo); let topic_ctx = ToolContext { session_id: Some("cli:chat-1".to_string()), topic_id: Some("topic-xyz".to_string()), ..ToolContext::default() }; let result = tool.execute_with_context(&topic_ctx, json!({})).await.unwrap(); assert!(result.success); let output: serde_json::Value = serde_json::from_str(&result.output).unwrap(); assert_eq!(output["scope_key"], "topic-xyz"); assert_eq!(output["count"], 1); assert_eq!(output["todos"][0]["content"], "话题任务"); assert_eq!(output["source"], "sqlite"); } }