Skip to main content

meerkat_memory/
tool.rs

1//! Memory search tool — exposes `MemoryStore::search` as an `AgentToolDispatcher`.
2//!
3//! This tool allows agents to search their semantic memory for past conversation
4//! content that was indexed during compaction. It wraps an `Arc<dyn MemoryStore>`
5//! and delegates to its `search()` method.
6
7use std::sync::Arc;
8
9use async_trait::async_trait;
10use meerkat_core::AgentToolDispatcher;
11use meerkat_core::error::ToolError;
12use meerkat_core::memory::MemoryStore;
13use meerkat_core::types::{ToolCallView, ToolDef, ToolResult};
14use schemars::JsonSchema;
15use serde::Deserialize;
16use serde_json::{Map, Value, json};
17
18const TOOL_NAME: &str = "memory_search";
19const DEFAULT_LIMIT: usize = 5;
20
21/// Input schema for the memory_search tool.
22#[derive(Debug, Deserialize, JsonSchema)]
23struct MemorySearchInput {
24    /// Natural language search query describing what you want to recall.
25    query: String,
26    /// Maximum number of results to return (default: 5, max: 20).
27    #[serde(default)]
28    limit: Option<usize>,
29}
30
31/// Generate the JSON schema for the input type, ensuring `properties` and
32/// `required` keys are always present (tool schema contract).
33fn input_schema() -> Value {
34    let schema = schemars::schema_for!(MemorySearchInput);
35    let mut value = serde_json::to_value(&schema).unwrap_or(Value::Null);
36    if let Value::Object(ref mut obj) = value
37        && obj.get("type").and_then(Value::as_str) == Some("object")
38    {
39        obj.entry("properties".to_string())
40            .or_insert_with(|| Value::Object(Map::new()));
41        obj.entry("required".to_string())
42            .or_insert_with(|| Value::Array(Vec::new()));
43    }
44    value
45}
46
47/// Tool dispatcher that provides the `memory_search` tool.
48///
49/// Wraps an `Arc<dyn MemoryStore>` and exposes semantic search as an
50/// agent-callable tool. Created by the factory when the `memory-store-session`
51/// feature is enabled.
52pub struct MemorySearchDispatcher {
53    store: Arc<dyn MemoryStore>,
54    tool_defs: Arc<[Arc<ToolDef>]>,
55}
56
57impl MemorySearchDispatcher {
58    /// Create a new memory search dispatcher backed by the given store.
59    pub fn new(store: Arc<dyn MemoryStore>) -> Self {
60        let tool_def = Arc::new(ToolDef {
61            name: TOOL_NAME.to_string(),
62            description: "Search semantic memory for past conversation content. \
63                Memory contains text from earlier conversation turns that were \
64                compacted away to save context space. Use this to recall \
65                information from earlier in the conversation or from previous sessions."
66                .to_string(),
67            input_schema: input_schema(),
68        });
69
70        Self {
71            store,
72            tool_defs: Arc::from(vec![tool_def]),
73        }
74    }
75
76    /// Usage instructions for the system prompt.
77    pub fn usage_instructions() -> &'static str {
78        "# Semantic Memory\n\n\
79         You have access to a semantic memory store that contains text from earlier \
80         conversation turns that were compacted away. Use the `memory_search` tool \
81         to recall information that is no longer in your visible context."
82    }
83}
84
85#[async_trait]
86impl AgentToolDispatcher for MemorySearchDispatcher {
87    fn tools(&self) -> Arc<[Arc<ToolDef>]> {
88        Arc::clone(&self.tool_defs)
89    }
90
91    async fn dispatch(&self, call: ToolCallView<'_>) -> Result<ToolResult, ToolError> {
92        if call.name != TOOL_NAME {
93            return Err(ToolError::NotFound {
94                name: call.name.to_string(),
95            });
96        }
97
98        let input: MemorySearchInput =
99            serde_json::from_str(call.args.get()).map_err(|e| ToolError::InvalidArguments {
100                name: TOOL_NAME.to_string(),
101                reason: e.to_string(),
102            })?;
103
104        let limit = input.limit.unwrap_or(DEFAULT_LIMIT).min(20);
105
106        let results = self.store.search(&input.query, limit).await.map_err(|e| {
107            ToolError::ExecutionFailed {
108                message: format!("{TOOL_NAME}: {e}"),
109            }
110        })?;
111
112        let results_json: Vec<Value> = results
113            .into_iter()
114            .map(|r| {
115                json!({
116                    "content": r.content,
117                    "score": r.score,
118                    "session_id": r.metadata.session_id.to_string(),
119                    "turn": r.metadata.turn,
120                })
121            })
122            .collect();
123
124        Ok(ToolResult::new(
125            call.id.to_string(),
126            serde_json::to_string(&results_json).unwrap_or_else(|_| "[]".to_string()),
127            false,
128        ))
129    }
130}
131
132#[cfg(test)]
133#[allow(clippy::unwrap_used, clippy::expect_used)]
134mod tests {
135    use super::*;
136    use meerkat_core::memory::{MemoryMetadata, MemoryStore};
137    use meerkat_core::types::SessionId;
138    use serde_json::value::RawValue;
139    use std::time::SystemTime;
140
141    /// Helper to create a ToolCallView for testing.
142    fn make_call(args_json: &str) -> (String, Box<RawValue>, String) {
143        let id = "test-call-1".to_string();
144        let raw = RawValue::from_string(args_json.to_string()).unwrap();
145        let name = TOOL_NAME.to_string();
146        (id, raw, name)
147    }
148
149    fn call_view<'a>(id: &'a str, raw: &'a RawValue, name: &'a str) -> ToolCallView<'a> {
150        ToolCallView {
151            id,
152            name,
153            args: raw,
154        }
155    }
156
157    fn meta() -> MemoryMetadata {
158        MemoryMetadata {
159            session_id: SessionId::new(),
160            turn: Some(1),
161            indexed_at: SystemTime::now(),
162        }
163    }
164
165    // ==================== Tool Definition Tests ====================
166
167    #[test]
168    fn test_tool_name() {
169        let store = Arc::new(crate::SimpleMemoryStore::new());
170        let dispatcher = MemorySearchDispatcher::new(store);
171        let tools = dispatcher.tools();
172        assert_eq!(tools.len(), 1);
173        assert_eq!(tools[0].name, "memory_search");
174    }
175
176    #[test]
177    fn test_tool_schema_has_required_query() {
178        let store = Arc::new(crate::SimpleMemoryStore::new());
179        let dispatcher = MemorySearchDispatcher::new(store);
180        let tools = dispatcher.tools();
181        let schema = &tools[0].input_schema;
182
183        assert_eq!(schema["type"], "object");
184        assert!(schema["properties"]["query"].is_object());
185        assert_eq!(schema["properties"]["query"]["type"], "string");
186
187        let required = schema["required"].as_array().unwrap();
188        let required_strs: Vec<&str> = required.iter().filter_map(|v| v.as_str()).collect();
189        assert!(required_strs.contains(&"query"));
190    }
191
192    #[test]
193    fn test_tool_schema_has_optional_limit() {
194        let store = Arc::new(crate::SimpleMemoryStore::new());
195        let dispatcher = MemorySearchDispatcher::new(store);
196        let tools = dispatcher.tools();
197        let schema = &tools[0].input_schema;
198
199        assert!(schema["properties"]["limit"].is_object());
200
201        // limit is NOT in required
202        let required = schema["required"].as_array().unwrap();
203        let required_strs: Vec<&str> = required.iter().filter_map(|v| v.as_str()).collect();
204        assert!(!required_strs.contains(&"limit"));
205    }
206
207    // ==================== Dispatch Tests ====================
208
209    #[tokio::test]
210    async fn test_search_returns_results() {
211        let store = Arc::new(crate::SimpleMemoryStore::new());
212        store
213            .index("The project codename is AURORA-7", meta())
214            .await
215            .unwrap();
216        store
217            .index("The budget was set at $42,000", meta())
218            .await
219            .unwrap();
220        store
221            .index("Meeting scheduled for next Tuesday", meta())
222            .await
223            .unwrap();
224
225        let dispatcher = MemorySearchDispatcher::new(store);
226        let (id, raw, name) = make_call(r#"{"query": "project codename"}"#);
227        let view = call_view(&id, &raw, &name);
228
229        let result = dispatcher.dispatch(view).await.unwrap();
230        assert!(!result.is_error);
231
232        let parsed: Vec<Value> = serde_json::from_str(&result.text_content()).unwrap();
233        assert!(!parsed.is_empty());
234        assert!(parsed[0]["content"].as_str().unwrap().contains("AURORA"));
235        assert!(parsed[0]["score"].as_f64().unwrap() > 0.0);
236        assert!(parsed[0]["session_id"].is_string());
237    }
238
239    #[tokio::test]
240    async fn test_search_empty_store_returns_empty() {
241        let store = Arc::new(crate::SimpleMemoryStore::new());
242        let dispatcher = MemorySearchDispatcher::new(store);
243
244        let (id, raw, name) = make_call(r#"{"query": "anything"}"#);
245        let view = call_view(&id, &raw, &name);
246
247        let result = dispatcher.dispatch(view).await.unwrap();
248        assert!(!result.is_error);
249
250        let parsed: Vec<Value> = serde_json::from_str(&result.text_content()).unwrap();
251        assert!(parsed.is_empty());
252    }
253
254    #[tokio::test]
255    async fn test_search_with_limit() {
256        let store = Arc::new(crate::SimpleMemoryStore::new());
257        for i in 0..10 {
258            store
259                .index(&format!("Memory entry {i} about testing"), meta())
260                .await
261                .unwrap();
262        }
263
264        let dispatcher = MemorySearchDispatcher::new(store);
265        let (id, raw, name) = make_call(r#"{"query": "testing", "limit": 3}"#);
266        let view = call_view(&id, &raw, &name);
267
268        let result = dispatcher.dispatch(view).await.unwrap();
269        let parsed: Vec<Value> = serde_json::from_str(&result.text_content()).unwrap();
270        assert_eq!(parsed.len(), 3);
271    }
272
273    #[tokio::test]
274    async fn test_search_default_limit() {
275        let store = Arc::new(crate::SimpleMemoryStore::new());
276        for i in 0..10 {
277            store
278                .index(&format!("Entry {i} about Rust programming"), meta())
279                .await
280                .unwrap();
281        }
282
283        let dispatcher = MemorySearchDispatcher::new(store);
284        let (id, raw, name) = make_call(r#"{"query": "Rust"}"#);
285        let view = call_view(&id, &raw, &name);
286
287        let result = dispatcher.dispatch(view).await.unwrap();
288        let parsed: Vec<Value> = serde_json::from_str(&result.text_content()).unwrap();
289        assert_eq!(parsed.len(), DEFAULT_LIMIT);
290    }
291
292    #[tokio::test]
293    async fn test_search_no_match_returns_empty() {
294        let store = Arc::new(crate::SimpleMemoryStore::new());
295        store
296            .index("The weather is sunny today", meta())
297            .await
298            .unwrap();
299
300        let dispatcher = MemorySearchDispatcher::new(store);
301        let (id, raw, name) = make_call(r#"{"query": "quantum physics"}"#);
302        let view = call_view(&id, &raw, &name);
303
304        let result = dispatcher.dispatch(view).await.unwrap();
305        let parsed: Vec<Value> = serde_json::from_str(&result.text_content()).unwrap();
306        assert!(parsed.is_empty());
307    }
308
309    #[tokio::test]
310    async fn test_dispatch_wrong_tool_name() {
311        let store = Arc::new(crate::SimpleMemoryStore::new());
312        let dispatcher = MemorySearchDispatcher::new(store);
313
314        let id = "test-1".to_string();
315        let raw = RawValue::from_string(r#"{"query": "test"}"#.to_string()).unwrap();
316        let name = "wrong_tool";
317        let view = ToolCallView {
318            id: &id,
319            name,
320            args: &raw,
321        };
322
323        let result = dispatcher.dispatch(view).await;
324        assert!(matches!(result, Err(ToolError::NotFound { .. })));
325    }
326
327    #[tokio::test]
328    async fn test_dispatch_invalid_args() {
329        let store = Arc::new(crate::SimpleMemoryStore::new());
330        let dispatcher = MemorySearchDispatcher::new(store);
331
332        let (id, raw, name) = make_call(r#"{"not_query": "test"}"#);
333        let view = call_view(&id, &raw, &name);
334
335        let result = dispatcher.dispatch(view).await;
336        assert!(matches!(result, Err(ToolError::InvalidArguments { .. })));
337    }
338
339    #[tokio::test]
340    async fn test_limit_capped_at_20() {
341        let store = Arc::new(crate::SimpleMemoryStore::new());
342        for i in 0..30 {
343            store
344                .index(&format!("Data point {i} about science"), meta())
345                .await
346                .unwrap();
347        }
348
349        let dispatcher = MemorySearchDispatcher::new(store);
350        let (id, raw, name) = make_call(r#"{"query": "science", "limit": 100}"#);
351        let view = call_view(&id, &raw, &name);
352
353        let result = dispatcher.dispatch(view).await.unwrap();
354        let parsed: Vec<Value> = serde_json::from_str(&result.text_content()).unwrap();
355        assert!(parsed.len() <= 20);
356    }
357
358    #[test]
359    fn test_usage_instructions_not_empty() {
360        let instructions = MemorySearchDispatcher::usage_instructions();
361        assert!(!instructions.is_empty());
362        assert!(instructions.contains("memory_search"));
363    }
364}