PicoBot/src/tools/todo_read.rs
oudecheng dc9211548a feat(model): 专家和子代理支持独立配置 provider/model
- 新增 ModelResolver 解析器,按 frontmatter 中的 provider/model 名覆盖基础 LLMProviderConfig

- Expert/SubagentDef 数据结构新增 provider/model 字段,frontmatter 解析与渲染支持往返

- AgentFactory 和 DefaultSubAgentRuntime 持有 ModelResolver,在创建 agent 时解析模型覆盖

- HTTP API 新增 /api/model-options 端点,ExpertResponse/Create/Update 和 SubagentUpdateRequest 支持 provider/model

- update_expert/update_subagent 支持 provider/model 字段写回 frontmatter

- 保持架构解耦:ModelResolver 位于 config 底层,不引入新的跨模块依赖
2026-07-30 22:44:10 +08:00

349 lines
11 KiB
Rust

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<TodoItem>,
count: usize,
scope_key: String,
source: &'static str,
}
// ── 工具实现 ──────────────────────────────────────────────
pub struct TodoReadTool {
/// 共享内存状态(与 TodoWriteTool 同一实例)
state: Arc<RwLock<HashMap<String, Vec<TodoItem>>>>,
/// SQLite 持久化层,用于进程重启后回填内存
repository: Arc<dyn TodoRepository>,
}
impl TodoReadTool {
pub(crate) fn new(
state: Arc<RwLock<HashMap<String, Vec<TodoItem>>>>,
repository: Arc<dyn TodoRepository>,
) -> 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<ToolResult> {
Ok(error_result("todo_read requires tool context (session_id)"))
}
async fn execute_with_context(
&self,
context: &ToolContext,
_args: serde_json::Value,
) -> anyhow::Result<ToolResult> {
// 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<TodoItem> = 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<RwLock<HashMap<String, Vec<TodoItem>>>> {
Arc::new(RwLock::new(HashMap::new()))
}
struct MockTodoRepository {
records: Vec<crate::storage::TodoRecord>,
}
impl TodoRepository for MockTodoRepository {
fn replace_todos(
&self,
_scope_key: &str,
_items: &[crate::storage::TodoRecord],
) -> Result<Vec<crate::storage::TodoRecord>, crate::storage::StorageError> {
Ok(vec![])
}
fn list_todos(
&self,
scope_key: &str,
) -> Result<Vec<crate::storage::TodoRecord>, 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");
}
}