1use 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#[derive(Debug, Deserialize, JsonSchema)]
23struct MemorySearchInput {
24 query: String,
26 #[serde(default)]
28 limit: Option<usize>,
29}
30
31fn 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
47pub struct MemorySearchDispatcher {
53 store: Arc<dyn MemoryStore>,
54 tool_defs: Arc<[Arc<ToolDef>]>,
55}
56
57impl MemorySearchDispatcher {
58 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 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 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 #[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 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 #[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}