use std::sync::Arc;
use dashmap::mapref::entry::Entry;
use dashmap::DashMap;
use super::tool_registry::ToolRegistry;
type CacheKey = (String, Option<usize>);
#[derive(Clone, Default)]
pub struct ToolRegistryCache {
entries: Arc<DashMap<CacheKey, Arc<ToolRegistry>>>,
}
impl ToolRegistryCache {
pub fn new() -> Self {
Self::default()
}
pub fn get_or_build(
&self,
agent_id: &str,
binding_index: Option<usize>,
base: &ToolRegistry,
allowed_tools: &[String],
) -> Arc<ToolRegistry> {
let key = (agent_id.to_string(), binding_index);
match self.entries.entry(key) {
Entry::Occupied(e) => Arc::clone(e.get()),
Entry::Vacant(slot) => {
let filtered = Arc::new(base.filtered_clone(allowed_tools));
slot.insert(Arc::clone(&filtered));
filtered
}
}
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use nexo_llm::ToolDef;
use serde_json::{json, Value};
use crate::agent::{AgentContext, ToolHandler};
struct NoopTool;
#[async_trait]
impl ToolHandler for NoopTool {
async fn call(&self, _ctx: &AgentContext, _args: Value) -> anyhow::Result<Value> {
Ok(json!({}))
}
}
fn tool_def(name: &str) -> ToolDef {
ToolDef {
name: name.into(),
description: String::new(),
parameters: json!({"type": "object"}),
}
}
fn base_registry() -> ToolRegistry {
let r = ToolRegistry::new();
r.register(tool_def("whatsapp_send_message"), NoopTool);
r.register(tool_def("memory_write"), NoopTool);
r.register(tool_def("memory_query"), NoopTool);
r.register(tool_def("browser_open"), NoopTool);
r
}
#[test]
fn filtered_registry_reflects_allowlist() {
let base = base_registry();
let cache = ToolRegistryCache::new();
let narrow = cache.get_or_build(
"ana",
Some(0),
&base,
&["whatsapp_send_message".to_string()],
);
assert!(narrow.contains("whatsapp_send_message"));
assert!(!narrow.contains("memory_write"));
assert!(!narrow.contains("browser_open"));
}
#[test]
fn wildcard_entry_keeps_everything() {
let base = base_registry();
let cache = ToolRegistryCache::new();
let full = cache.get_or_build("ana", Some(1), &base, &["*".to_string()]);
assert!(full.contains("whatsapp_send_message"));
assert!(full.contains("memory_write"));
assert!(full.contains("browser_open"));
}
#[test]
fn empty_allowlist_keeps_everything() {
let base = base_registry();
let cache = ToolRegistryCache::new();
let full = cache.get_or_build("legacy", None, &base, &[]);
assert_eq!(full.to_tool_defs().len(), 4);
}
#[test]
fn prefix_glob_matches() {
let base = base_registry();
let cache = ToolRegistryCache::new();
let mem_only = cache.get_or_build("ana", Some(2), &base, &["memory_*".to_string()]);
assert!(mem_only.contains("memory_write"));
assert!(mem_only.contains("memory_query"));
assert!(!mem_only.contains("whatsapp_send_message"));
}
#[test]
fn repeated_get_is_cache_hit() {
let base = base_registry();
let cache = ToolRegistryCache::new();
let a = cache.get_or_build("ana", Some(0), &base, &["*".to_string()]);
let b = cache.get_or_build("ana", Some(0), &base, &["*".to_string()]);
assert_eq!(cache.len(), 1);
assert!(Arc::ptr_eq(&a, &b));
}
#[test]
fn different_bindings_produce_independent_entries() {
let base = base_registry();
let cache = ToolRegistryCache::new();
let wa = cache.get_or_build("ana", Some(0), &base, &["whatsapp_send_message".into()]);
let tg = cache.get_or_build("ana", Some(1), &base, &["*".into()]);
assert_eq!(cache.len(), 2);
assert!(!Arc::ptr_eq(&wa, &tg));
assert_eq!(wa.to_tool_defs().len(), 1);
assert_eq!(tg.to_tool_defs().len(), 4);
}
#[test]
fn filtered_clone_leaves_base_untouched() {
let base = base_registry();
let _narrow = base.filtered_clone(&["whatsapp_send_message".to_string()]);
assert_eq!(base.to_tool_defs().len(), 4);
}
}