use std::sync::Arc;
use async_trait::async_trait;
use meerkat_core::AgentToolDispatcher;
use meerkat_core::error::ToolError;
use meerkat_core::memory::{MemorySearchScope, MemoryStore};
use meerkat_core::types::{ToolCallView, ToolDef, ToolProvenance, ToolResult, ToolSourceKind};
use schemars::JsonSchema;
use serde::Deserialize;
use serde_json::{Map, Value, json};
const TOOL_NAME: &str = "memory_search";
const DEFAULT_LIMIT: usize = 5;
#[derive(Debug, Deserialize, JsonSchema)]
struct MemorySearchInput {
query: String,
#[serde(default)]
limit: Option<usize>,
}
fn input_schema() -> Value {
let schema = schemars::schema_for!(MemorySearchInput);
let mut value = serde_json::to_value(&schema).unwrap_or(Value::Null);
if let Value::Object(ref mut obj) = value
&& obj.get("type").and_then(Value::as_str) == Some("object")
{
obj.entry("properties".to_string())
.or_insert_with(|| Value::Object(Map::new()));
obj.entry("required".to_string())
.or_insert_with(|| Value::Array(Vec::new()));
}
value
}
pub struct MemorySearchDispatcher {
store: Arc<dyn MemoryStore>,
scope: MemorySearchScope,
tool_defs: Arc<[Arc<ToolDef>]>,
}
impl MemorySearchDispatcher {
pub fn new(store: Arc<dyn MemoryStore>, scope: MemorySearchScope) -> Self {
let tool_def = Arc::new(ToolDef {
name: TOOL_NAME.into(),
description: "Search semantic memory for past conversation content. \
Memory contains text from earlier conversation turns that were \
compacted away to save context space. Use this to recall \
information from earlier in the conversation or from previous sessions."
.to_string(),
input_schema: input_schema(),
provenance: Some(ToolProvenance {
kind: ToolSourceKind::Memory,
source_id: "memory".into(),
}),
});
Self {
store,
scope,
tool_defs: Arc::from(vec![tool_def]),
}
}
pub fn for_session(
store: Arc<dyn MemoryStore>,
session_id: meerkat_core::types::SessionId,
) -> Self {
Self::new(store, MemorySearchScope::for_session(session_id))
}
pub fn usage_instructions() -> &'static str {
"# Semantic Memory\n\n\
You have access to a semantic memory store that contains text from earlier \
conversation turns that were compacted away. Use the `memory_search` tool \
to recall information that is no longer in your visible context."
}
}
#[async_trait]
impl AgentToolDispatcher for MemorySearchDispatcher {
fn tools(&self) -> Arc<[Arc<ToolDef>]> {
Arc::clone(&self.tool_defs)
}
async fn dispatch(
&self,
call: ToolCallView<'_>,
) -> Result<meerkat_core::ops::ToolDispatchOutcome, ToolError> {
if call.name != TOOL_NAME {
return Err(ToolError::NotFound {
name: call.name.into(),
});
}
let input: MemorySearchInput =
serde_json::from_str(call.args.get()).map_err(|e| ToolError::InvalidArguments {
name: TOOL_NAME.into(),
reason: e.to_string(),
})?;
let limit = input.limit.unwrap_or(DEFAULT_LIMIT).min(20);
let results = self
.store
.search(&self.scope, &input.query, limit)
.await
.map_err(|e| ToolError::ExecutionFailed {
message: e.to_string(),
})?;
let items: Vec<Value> = results
.into_iter()
.map(|r| {
json!({
"content": r.content,
"score": r.score,
"turn": r.metadata.turn,
})
})
.collect();
let payload = Value::Array(items).to_string();
Ok(meerkat_core::ops::ToolDispatchOutcome::from(
ToolResult::new(call.id.to_string(), payload, false),
))
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use meerkat_core::memory::{MemoryIndexRequest, MemoryIndexScope, MemoryMetadata, MemoryStore};
use meerkat_core::types::SessionId;
use serde_json::value::RawValue;
use std::time::SystemTime;
fn make_call(args_json: &str) -> (String, Box<RawValue>, String) {
let id = "test-call-1".to_string();
let raw = RawValue::from_string(args_json.to_string()).unwrap();
let name = TOOL_NAME.to_string();
(id, raw, name)
}
fn call_view<'a>(id: &'a str, raw: &'a RawValue, name: &'a str) -> ToolCallView<'a> {
ToolCallView {
id,
name,
args: raw,
}
}
fn meta(session_id: &SessionId) -> MemoryMetadata {
MemoryMetadata {
session_id: session_id.clone(),
turn: Some(1),
indexed_at: SystemTime::now(),
}
}
fn request(content: impl Into<String>, session_id: &SessionId) -> MemoryIndexRequest {
MemoryIndexRequest::new(
MemoryIndexScope::for_session(session_id.clone()),
content.into(),
meta(session_id),
)
.unwrap()
}
fn dispatcher(store: Arc<dyn MemoryStore>, session_id: &SessionId) -> MemorySearchDispatcher {
MemorySearchDispatcher::for_session(store, session_id.clone())
}
#[test]
fn test_tool_name() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let tools = dispatcher.tools();
assert_eq!(tools.len(), 1);
assert_eq!(tools[0].name, "memory_search");
}
#[test]
fn test_tool_schema_has_required_query() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let tools = dispatcher.tools();
let schema = &tools[0].input_schema;
assert_eq!(schema["type"], "object");
assert!(schema["properties"]["query"].is_object());
assert_eq!(schema["properties"]["query"]["type"], "string");
let required = schema["required"].as_array().unwrap();
let required_strs: Vec<&str> = required.iter().filter_map(|v| v.as_str()).collect();
assert!(required_strs.contains(&"query"));
}
#[test]
fn test_tool_schema_has_optional_limit() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let tools = dispatcher.tools();
let schema = &tools[0].input_schema;
assert!(schema["properties"]["limit"].is_object());
let required = schema["required"].as_array().unwrap();
let required_strs: Vec<&str> = required.iter().filter_map(|v| v.as_str()).collect();
assert!(!required_strs.contains(&"limit"));
}
#[tokio::test]
async fn test_search_returns_results() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let other_session_id = SessionId::new();
store
.index_scoped(request("The project codename is AURORA-7", &session_id))
.await
.unwrap();
store
.index_scoped(request("The budget was set at $42,000", &session_id))
.await
.unwrap();
store
.index_scoped(request(
"Meeting scheduled for next Tuesday",
&other_session_id,
))
.await
.unwrap();
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"query": "project codename"}"#);
let view = call_view(&id, &raw, &name);
let outcome = dispatcher.dispatch(view).await.unwrap();
assert!(!outcome.result.is_error);
let parsed: Vec<Value> = serde_json::from_str(&outcome.result.text_content()).unwrap();
assert!(!parsed.is_empty());
assert!(parsed[0]["content"].as_str().unwrap().contains("AURORA"));
assert!(parsed[0]["score"].as_f64().unwrap() > 0.0);
assert!(
parsed[0].get("session_id").is_none(),
"memory tool must not leak raw source session ids"
);
}
#[tokio::test]
async fn test_search_empty_store_returns_empty() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"query": "anything"}"#);
let view = call_view(&id, &raw, &name);
let outcome = dispatcher.dispatch(view).await.unwrap();
assert!(!outcome.result.is_error);
let parsed: Vec<Value> = serde_json::from_str(&outcome.result.text_content()).unwrap();
assert!(parsed.is_empty());
}
#[tokio::test]
async fn test_search_with_limit() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
for i in 0..10 {
store
.index_scoped(request(
format!("Memory entry {i} about testing"),
&session_id,
))
.await
.unwrap();
}
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"query": "testing", "limit": 3}"#);
let view = call_view(&id, &raw, &name);
let outcome = dispatcher.dispatch(view).await.unwrap();
let parsed: Vec<Value> = serde_json::from_str(&outcome.result.text_content()).unwrap();
assert_eq!(parsed.len(), 3);
}
#[tokio::test]
async fn test_search_default_limit() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
for i in 0..10 {
store
.index_scoped(request(
format!("Entry {i} about Rust programming"),
&session_id,
))
.await
.unwrap();
}
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"query": "Rust"}"#);
let view = call_view(&id, &raw, &name);
let outcome = dispatcher.dispatch(view).await.unwrap();
let parsed: Vec<Value> = serde_json::from_str(&outcome.result.text_content()).unwrap();
assert_eq!(parsed.len(), DEFAULT_LIMIT);
}
#[tokio::test]
async fn test_search_no_match_returns_empty() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
store
.index_scoped(request("The weather is sunny today", &session_id))
.await
.unwrap();
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"query": "quantum physics"}"#);
let view = call_view(&id, &raw, &name);
let outcome = dispatcher.dispatch(view).await.unwrap();
let parsed: Vec<Value> = serde_json::from_str(&outcome.result.text_content()).unwrap();
assert!(parsed.is_empty());
}
#[tokio::test]
async fn test_dispatch_wrong_tool_name() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let id = "test-1".to_string();
let raw = RawValue::from_string(r#"{"query": "test"}"#.to_string()).unwrap();
let name = "wrong_tool";
let view = ToolCallView {
id: &id,
name,
args: &raw,
};
let result = dispatcher.dispatch(view).await;
assert!(matches!(result, Err(ToolError::NotFound { .. })));
}
#[tokio::test]
async fn test_dispatch_invalid_args() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"not_query": "test"}"#);
let view = call_view(&id, &raw, &name);
let result = dispatcher.dispatch(view).await;
assert!(matches!(result, Err(ToolError::InvalidArguments { .. })));
}
#[tokio::test]
async fn test_limit_capped_at_20() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::SimpleMemoryStore::new());
let session_id = SessionId::new();
for i in 0..30 {
store
.index_scoped(request(
format!("Data point {i} about science"),
&session_id,
))
.await
.unwrap();
}
let dispatcher = dispatcher(store, &session_id);
let (id, raw, name) = make_call(r#"{"query": "science", "limit": 100}"#);
let view = call_view(&id, &raw, &name);
let outcome = dispatcher.dispatch(view).await.unwrap();
let parsed: Vec<Value> = serde_json::from_str(&outcome.result.text_content()).unwrap();
assert!(parsed.len() <= 20);
}
#[test]
fn test_usage_instructions_not_empty() {
let instructions = MemorySearchDispatcher::usage_instructions();
assert!(!instructions.is_empty());
assert!(instructions.contains("memory_search"));
}
#[test]
fn memory_tools_have_memory_provenance() {
let store: Arc<dyn MemoryStore> = Arc::new(crate::simple::SimpleMemoryStore::new());
let session_id = SessionId::new();
let dispatcher = dispatcher(store, &session_id);
let tools = dispatcher.tools();
assert_eq!(tools.len(), 1);
let prov = tools[0]
.provenance
.as_ref()
.expect("memory tool should have provenance");
assert_eq!(prov.kind, meerkat_core::types::ToolSourceKind::Memory);
assert_eq!(prov.source_id, "memory");
}
}