use async_trait::async_trait;
use everruns_core::atoms::{PostToolExecHook, PostToolExecHookPriority};
use everruns_core::capabilities::{Capability, CapabilityStatus, TOOL_SEARCH_TOOL_NAME};
use everruns_core::tool_types::{ToolCall, ToolDefinition, ToolResult};
use everruns_core::traits::ToolContext;
use everruns_core::typed_id::SessionId;
use serde_json::Value;
use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::{Arc, Mutex, MutexGuard};
pub(crate) const TOOL_REVEAL_CAPABILITY_ID: &str = "yolop_tool_reveal";
const MAX_TRACKED_SESSIONS: usize = 256;
#[derive(Default)]
pub(crate) struct RevealedTools {
inner: Mutex<RevealState>,
}
#[derive(Default)]
struct RevealState {
by_session: HashMap<SessionId, HashSet<String>>,
order: VecDeque<SessionId>,
}
impl RevealedTools {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn record(&self, session: SessionId, names: impl IntoIterator<Item = String>) {
let mut state = self.lock();
if !state.by_session.contains_key(&session) {
state.order.push_back(session);
while state.order.len() > MAX_TRACKED_SESSIONS {
if let Some(evicted) = state.order.pop_front() {
state.by_session.remove(&evicted);
}
}
}
state.by_session.entry(session).or_default().extend(names);
}
pub(crate) fn any_revealed(&self, session: SessionId, tools: &[&str]) -> bool {
let state = self.lock();
state
.by_session
.get(&session)
.is_some_and(|revealed| tools.iter().any(|name| revealed.contains(*name)))
}
fn lock(&self) -> MutexGuard<'_, RevealState> {
self.inner.lock().unwrap_or_else(|err| err.into_inner())
}
}
pub(crate) struct ToolRevealCapability {
reveals: Arc<RevealedTools>,
}
impl ToolRevealCapability {
pub(crate) fn new(reveals: Arc<RevealedTools>) -> Self {
Self { reveals }
}
}
#[async_trait]
impl Capability for ToolRevealCapability {
fn id(&self) -> &str {
TOOL_REVEAL_CAPABILITY_ID
}
fn name(&self) -> &str {
"Tool Reveal Tracking"
}
fn description(&self) -> &str {
"Tracks which deferred tool schemas `tool_search` has loaded, so capability \
prompt blocks can wait until their tools are callable."
}
fn status(&self) -> CapabilityStatus {
CapabilityStatus::Available
}
fn category(&self) -> Option<&str> {
Some("Guardrails")
}
fn post_tool_exec_hooks(&self) -> Vec<Arc<dyn PostToolExecHook>> {
vec![Arc::new(RevealTrackerHook {
reveals: self.reveals.clone(),
})]
}
}
struct RevealTrackerHook {
reveals: Arc<RevealedTools>,
}
#[async_trait]
impl PostToolExecHook for RevealTrackerHook {
fn priority(&self) -> PostToolExecHookPriority {
PostToolExecHookPriority::Normal
}
async fn after_exec(
&self,
tool_call: &ToolCall,
_tool_def: &ToolDefinition,
result: &mut ToolResult,
context: &ToolContext,
) {
if tool_call.name != TOOL_SEARCH_TOOL_NAME || result.error.is_some() {
return;
}
let loaded = loaded_tool_names(result.result.as_ref());
if !loaded.is_empty() {
self.reveals.record(context.session_id, loaded);
}
}
}
fn loaded_tool_names(result: Option<&Value>) -> Vec<String> {
result
.and_then(|value| value.get("loaded"))
.and_then(Value::as_array)
.map(|names| {
names
.iter()
.filter_map(Value::as_str)
.map(str::to_string)
.collect()
})
.unwrap_or_default()
}
#[cfg(test)]
mod tests {
use super::*;
use everruns_core::tool_types::{
BuiltinTool, DeferrablePolicy, ToolHints, ToolPolicy, ToolResult,
};
use serde_json::json;
fn tool_def(name: &str) -> ToolDefinition {
ToolDefinition::Builtin(BuiltinTool {
name: name.to_string(),
display_name: None,
description: String::new(),
parameters: json!({ "type": "object", "properties": {} }),
policy: ToolPolicy::Auto,
category: None,
deferrable: DeferrablePolicy::Automatic,
hints: ToolHints::default(),
full_parameters: None,
})
}
fn call(name: &str) -> ToolCall {
ToolCall {
id: format!("call-{name}"),
name: name.to_string(),
arguments: json!({}),
}
}
async fn run_hook(reveals: &Arc<RevealedTools>, session: SessionId, result: Value) {
let hook = RevealTrackerHook {
reveals: reveals.clone(),
};
let mut tool_result = ToolResult {
tool_call_id: "call".to_string(),
result: Some(result),
images: None,
error: None,
connection_required: None,
raw_output: None,
};
hook.after_exec(
&call(TOOL_SEARCH_TOOL_NAME),
&tool_def(TOOL_SEARCH_TOOL_NAME),
&mut tool_result,
&ToolContext::new(session),
)
.await;
}
#[tokio::test]
async fn tool_search_results_reveal_the_tools_they_loaded() {
let reveals = Arc::new(RevealedTools::new());
let session = SessionId::new();
assert!(
!reveals.any_revealed(session, &["get_config"]),
"nothing is revealed before the first tool_search"
);
run_hook(
&reveals,
session,
json!({ "query": "config", "loaded": ["get_config", "set_config"] }),
)
.await;
assert!(reveals.any_revealed(session, &["get_config"]));
assert!(
reveals.any_revealed(session, &["recall", "set_config"]),
"any listed tool being revealed is enough"
);
assert!(!reveals.any_revealed(session, &["remember"]));
}
#[tokio::test]
async fn reveals_do_not_leak_between_sessions() {
let reveals = Arc::new(RevealedTools::new());
let (first, second) = (SessionId::new(), SessionId::new());
run_hook(&reveals, first, json!({ "loaded": ["get_config"] })).await;
assert!(reveals.any_revealed(first, &["get_config"]));
assert!(
!reveals.any_revealed(second, &["get_config"]),
"a reveal in one session must not unhide prompt text in another"
);
}
#[tokio::test]
async fn empty_match_searches_reveal_nothing() {
let reveals = Arc::new(RevealedTools::new());
let session = SessionId::new();
run_hook(
&reveals,
session,
json!({ "tools": [], "available_tools": ["get_config"] }),
)
.await;
assert!(!reveals.any_revealed(session, &["get_config"]));
}
#[test]
fn tracked_sessions_are_bounded() {
let reveals = RevealedTools::new();
let first = SessionId::new();
reveals.record(first, ["get_config".to_string()]);
for _ in 0..MAX_TRACKED_SESSIONS {
reveals.record(SessionId::new(), ["get_config".to_string()]);
}
assert!(
!reveals.any_revealed(first, &["get_config"]),
"the oldest session is evicted once the bound is crossed"
);
assert_eq!(reveals.lock().by_session.len(), MAX_TRACKED_SESSIONS);
}
}