use std::collections::{BTreeMap, HashMap};
use serde::{Deserialize, Serialize};
use crate::adapter::net::behavior::fold::capability_aggregation::TagMatcher;
use crate::adapter::net::behavior::ToolCapability;
pub fn description_metadata_key(tool_id: &str) -> String {
format!("tool::{tool_id}::description")
}
pub fn streaming_metadata_key(tool_id: &str) -> String {
format!("tool::{tool_id}::streaming")
}
pub fn tags_metadata_key(tool_id: &str) -> String {
format!("tool::{tool_id}::tags")
}
pub fn pricing_terms_metadata_key(tool_id: &str) -> String {
format!("tool::{tool_id}::pricing_terms")
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ToolDescriptor {
pub tool_id: String,
pub name: String,
pub version: String,
pub description: Option<String>,
pub input_schema: Option<String>,
pub output_schema: Option<String>,
pub requires: Vec<String>,
pub estimated_time_ms: u32,
pub stateless: bool,
pub streaming: bool,
pub tags: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub pricing_terms: Option<String>,
pub node_count: u32,
}
impl ToolDescriptor {
pub fn from_capability(cap: &ToolCapability, metadata: &BTreeMap<String, String>) -> Self {
let description = metadata
.get(&description_metadata_key(&cap.tool_id))
.cloned();
let streaming = metadata
.get(&streaming_metadata_key(&cap.tool_id))
.map(|s| s == "1")
.unwrap_or(false);
let tags = metadata
.get(&tags_metadata_key(&cap.tool_id))
.map(|raw| {
raw.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
})
.unwrap_or_default();
let pricing_terms = metadata
.get(&pricing_terms_metadata_key(&cap.tool_id))
.cloned();
Self {
tool_id: cap.tool_id.clone(),
name: cap.name.clone(),
version: cap.version.clone(),
description,
input_schema: cap.input_schema.clone(),
output_schema: cap.output_schema.clone(),
requires: cap.requires.clone(),
estimated_time_ms: cap.estimated_time_ms,
stateless: cap.stateless,
streaming,
tags,
pricing_terms,
node_count: 0,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ToolEvent {
Start {
tool_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
call_id: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
metadata: Option<serde_json::Value>,
},
Progress {
#[serde(default, skip_serializing_if = "Option::is_none")]
pct: Option<f32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
message: Option<String>,
},
Delta {
data: serde_json::Value,
},
Result {
data: serde_json::Value,
},
Error {
code: String,
message: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
details: Option<serde_json::Value>,
},
}
impl ToolEvent {
pub fn is_terminal(&self) -> bool {
matches!(self, Self::Result { .. } | Self::Error { .. })
}
}
pub struct ToolListWatch {
pub(crate) receiver: tokio::sync::mpsc::Receiver<ToolListChange>,
pub(crate) cancel: std::sync::Arc<tokio::sync::Notify>,
}
impl futures::Stream for ToolListWatch {
type Item = ToolListChange;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.receiver.poll_recv(cx)
}
}
impl ToolListWatch {
pub fn cancel(&self) {
self.cancel.notify_one();
}
pub fn cancel_handle(&self) -> std::sync::Arc<tokio::sync::Notify> {
self.cancel.clone()
}
pub async fn recv(&mut self) -> Option<ToolListChange> {
self.receiver.recv().await
}
pub fn try_recv(&mut self) -> Option<ToolListChange> {
self.receiver.try_recv().ok()
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum ToolListChange {
Added(ToolDescriptor),
Removed(ToolDescriptor),
NodeCountChanged {
descriptor: ToolDescriptor,
prev_node_count: u32,
},
}
pub const TOOL_METADATA_FETCH_SERVICE: &str = "tool.metadata.fetch";
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ToolMetadataRequest {
pub name: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "type", rename_all = "snake_case")]
#[expect(
clippy::large_enum_variant,
reason = "one transient response per tool.metadata.fetch RPC, decoded and \
immediately consumed — boxing the descriptor would add an \
allocation + wire-invisible indirection to every Found for a \
stack-size win nothing is sensitive to (ToolDescriptor crossed \
the 200-byte lint threshold when pricing_terms landed)"
)]
pub enum ToolMetadataResponse {
Found {
descriptor: ToolDescriptor,
},
NotFound {
name: String,
},
}
pub const TOOL_WATCH_SERVICE: &str = "tool.watch";
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WatchToolsRequest {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub matcher: Option<TagMatcher>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub interval_ms: Option<u64>,
}
#[derive(Debug, Clone, PartialEq)]
#[expect(
clippy::large_enum_variant,
reason = "one transient frame per streamed change, encoded and immediately \
consumed — boxing the change would add an allocation + \
wire-invisible indirection to every delta for a stack-size win \
nothing is sensitive to (same call as ToolMetadataResponse above)"
)]
pub enum ToolWatchFrame {
Change(ToolListChange),
Resync,
}
const TOOL_WATCH_FRAME_KINDS: [&str; 4] = ["added", "removed", "node_count_changed", "resync"];
impl Serialize for ToolWatchFrame {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
#[derive(Serialize)]
struct Wire<'a> {
#[serde(rename = "type")]
kind: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
descriptor: Option<&'a ToolDescriptor>,
#[serde(skip_serializing_if = "Option::is_none")]
prev_node_count: Option<u32>,
}
let wire = match self {
Self::Change(ToolListChange::Added(d)) => Wire {
kind: TOOL_WATCH_FRAME_KINDS[0],
descriptor: Some(d),
prev_node_count: None,
},
Self::Change(ToolListChange::Removed(d)) => Wire {
kind: TOOL_WATCH_FRAME_KINDS[1],
descriptor: Some(d),
prev_node_count: None,
},
Self::Change(ToolListChange::NodeCountChanged {
descriptor,
prev_node_count,
}) => Wire {
kind: TOOL_WATCH_FRAME_KINDS[2],
descriptor: Some(descriptor),
prev_node_count: Some(*prev_node_count),
},
Self::Resync => Wire {
kind: TOOL_WATCH_FRAME_KINDS[3],
descriptor: None,
prev_node_count: None,
},
};
wire.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for ToolWatchFrame {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
#[derive(Deserialize)]
struct Wire {
#[serde(rename = "type")]
kind: String,
#[serde(default)]
descriptor: Option<ToolDescriptor>,
#[serde(default)]
prev_node_count: Option<u32>,
}
use serde::de::Error;
let wire = Wire::deserialize(deserializer)?;
let descriptor =
|d: Option<ToolDescriptor>| d.ok_or_else(|| D::Error::missing_field("descriptor"));
match wire.kind.as_str() {
"added" => Ok(Self::Change(ToolListChange::Added(descriptor(
wire.descriptor,
)?))),
"removed" => Ok(Self::Change(ToolListChange::Removed(descriptor(
wire.descriptor,
)?))),
"node_count_changed" => Ok(Self::Change(ToolListChange::NodeCountChanged {
descriptor: descriptor(wire.descriptor)?,
prev_node_count: wire
.prev_node_count
.ok_or_else(|| D::Error::missing_field("prev_node_count"))?,
})),
"resync" => Ok(Self::Resync),
other => Err(D::Error::unknown_variant(other, &TOOL_WATCH_FRAME_KINDS)),
}
}
}
#[derive(Debug, Default)]
pub struct ToolMetadataRegistry {
inner: parking_lot::Mutex<RegistryState>,
change_signal: Option<std::sync::Arc<tokio::sync::watch::Sender<u64>>>,
}
#[derive(Debug, Default)]
struct RegistryState {
map: HashMap<String, ToolDescriptor>,
snapshot: Option<std::sync::Arc<[ToolDescriptor]>>,
}
impl ToolMetadataRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn with_change_signal(signal: std::sync::Arc<tokio::sync::watch::Sender<u64>>) -> Self {
Self {
inner: Default::default(),
change_signal: Some(signal),
}
}
fn signal_changed(&self) {
if let Some(tx) = &self.change_signal {
tx.send_modify(|g| *g = g.wrapping_add(1));
}
}
pub fn insert(&self, descriptor: ToolDescriptor) -> Option<ToolDescriptor> {
let (prev, changed) = {
let mut guard = self.inner.lock();
let changed = guard.map.get(&descriptor.tool_id) != Some(&descriptor);
if changed {
guard.snapshot = None;
}
let prev = guard.map.insert(descriptor.tool_id.clone(), descriptor);
(prev, changed)
};
if changed {
self.signal_changed();
}
prev
}
pub fn try_insert(&self, descriptor: ToolDescriptor) -> bool {
let inserted = {
let mut guard = self.inner.lock();
if guard.map.contains_key(&descriptor.tool_id) {
false
} else {
guard.map.insert(descriptor.tool_id.clone(), descriptor);
guard.snapshot = None;
true
}
};
if inserted {
self.signal_changed();
}
inserted
}
pub fn get(&self, name: &str) -> Option<ToolDescriptor> {
self.inner.lock().map.get(name).cloned()
}
pub fn remove(&self, name: &str) -> Option<ToolDescriptor> {
let prev = {
let mut guard = self.inner.lock();
let prev = guard.map.remove(name);
if prev.is_some() {
guard.snapshot = None;
}
prev
};
if prev.is_some() {
self.signal_changed();
}
prev
}
pub fn len(&self) -> usize {
self.inner.lock().map.len()
}
pub fn is_empty(&self) -> bool {
self.inner.lock().map.is_empty()
}
pub fn snapshot(&self) -> std::sync::Arc<[ToolDescriptor]> {
let mut guard = self.inner.lock();
if let Some(s) = &guard.snapshot {
return s.clone();
}
let snap: std::sync::Arc<[ToolDescriptor]> =
guard.map.values().cloned().collect::<Vec<_>>().into();
guard.snapshot = Some(snap.clone());
snap
}
}
#[cfg(test)]
mod tests {
use super::*;
fn cap(tool_id: &str) -> ToolCapability {
ToolCapability::new(tool_id, format!("Name for {tool_id}"))
.with_version("1.2.3")
.with_input_schema(r#"{"type":"object"}"#)
}
#[test]
fn metadata_keys_follow_existing_convention() {
assert_eq!(
description_metadata_key("web_search"),
"tool::web_search::description"
);
assert_eq!(
streaming_metadata_key("web_search"),
"tool::web_search::streaming"
);
assert_eq!(tags_metadata_key("web_search"), "tool::web_search::tags");
}
#[test]
fn descriptor_from_capability_picks_up_metadata_fields() {
let cap = cap("web_search");
let mut meta = BTreeMap::new();
meta.insert(
description_metadata_key("web_search"),
"Search the web.".to_string(),
);
meta.insert(streaming_metadata_key("web_search"), "1".to_string());
meta.insert(
tags_metadata_key("web_search"),
"web,research,external".to_string(),
);
let desc = ToolDescriptor::from_capability(&cap, &meta);
assert_eq!(desc.tool_id, "web_search");
assert_eq!(desc.version, "1.2.3");
assert_eq!(desc.description.as_deref(), Some("Search the web."));
assert!(desc.streaming);
assert_eq!(desc.tags, vec!["web", "research", "external"]);
assert_eq!(desc.input_schema.as_deref(), Some(r#"{"type":"object"}"#));
assert_eq!(
desc.node_count, 0,
"node_count is filled by the aggregator, not here"
);
}
fn sample_descriptor() -> ToolDescriptor {
ToolDescriptor::from_capability(&cap("web_search"), &BTreeMap::new())
}
#[test]
fn tool_watch_frame_json_matches_ffi_change_shape() {
let added = ToolWatchFrame::Change(ToolListChange::Added(sample_descriptor()));
let v: serde_json::Value = serde_json::to_value(&added).unwrap();
assert_eq!(v["type"], "added");
assert_eq!(v["descriptor"]["tool_id"], "web_search");
assert!(v.get("prev_node_count").is_none());
let ncc = ToolWatchFrame::Change(ToolListChange::NodeCountChanged {
descriptor: sample_descriptor(),
prev_node_count: 3,
});
let v: serde_json::Value = serde_json::to_value(&ncc).unwrap();
assert_eq!(v["type"], "node_count_changed");
assert_eq!(v["prev_node_count"], 3);
let v: serde_json::Value = serde_json::to_value(&ToolWatchFrame::Resync).unwrap();
assert_eq!(v, serde_json::json!({ "type": "resync" }));
}
#[test]
fn tool_watch_frame_round_trips_every_variant() {
let frames = vec![
ToolWatchFrame::Change(ToolListChange::Added(sample_descriptor())),
ToolWatchFrame::Change(ToolListChange::Removed(sample_descriptor())),
ToolWatchFrame::Change(ToolListChange::NodeCountChanged {
descriptor: sample_descriptor(),
prev_node_count: 7,
}),
ToolWatchFrame::Resync,
];
for frame in frames {
let json = serde_json::to_vec(&frame).unwrap();
let back: ToolWatchFrame = serde_json::from_slice(&json).unwrap();
assert_eq!(back, frame);
}
}
#[test]
fn tool_watch_frame_rejects_malformed_wire() {
let err = serde_json::from_str::<ToolWatchFrame>(r#"{"type":"added"}"#);
assert!(err.is_err(), "added without descriptor must fail");
let d = serde_json::to_value(sample_descriptor()).unwrap();
let raw = serde_json::json!({ "type": "node_count_changed", "descriptor": d });
assert!(serde_json::from_value::<ToolWatchFrame>(raw).is_err());
let err = serde_json::from_str::<ToolWatchFrame>(r#"{"type":"nonsense"}"#);
assert!(err.is_err(), "unknown frame kind must fail");
}
#[test]
fn watch_tools_request_omits_absent_options() {
let req = WatchToolsRequest {
matcher: None,
interval_ms: None,
};
assert_eq!(serde_json::to_string(&req).unwrap(), "{}");
let req = WatchToolsRequest {
matcher: Some(TagMatcher::Prefix {
value: "ai-tool:".into(),
}),
interval_ms: Some(250),
};
let json = serde_json::to_string(&req).unwrap();
let back: WatchToolsRequest = serde_json::from_str(&json).unwrap();
assert_eq!(back, req);
}
#[test]
fn change_signal_fires_on_real_mutations_only() {
let signal = std::sync::Arc::new(tokio::sync::watch::channel(0u64).0);
let reg = ToolMetadataRegistry::with_change_signal(signal.clone());
let generation = || *signal.borrow();
assert!(reg.remove("ghost").is_none());
assert_eq!(generation(), 0, "remove of an absent tool must not signal");
let desc = ToolDescriptor::from_capability(&cap("web_search"), &BTreeMap::new());
reg.insert(desc.clone());
assert_eq!(generation(), 1, "insert must signal");
reg.insert(desc.clone());
assert_eq!(generation(), 1, "idempotent re-insert must not signal");
let mut changed = desc;
changed.description = Some("now with a description".to_string());
reg.insert(changed);
assert_eq!(generation(), 2, "a changed replace must signal");
assert!(reg.remove("web_search").is_some());
assert_eq!(generation(), 3, "remove of a present tool must signal");
let _ = reg.get("web_search");
let _ = reg.snapshot();
let _ = reg.is_empty();
assert_eq!(generation(), 3, "reads must not signal");
}
#[test]
fn try_insert_rejects_duplicates_without_mutation_or_signal() {
let signal = std::sync::Arc::new(tokio::sync::watch::channel(0u64).0);
let reg = ToolMetadataRegistry::with_change_signal(signal.clone());
let generation = || *signal.borrow();
let original = ToolDescriptor::from_capability(&cap("web_search"), &BTreeMap::new());
assert!(reg.try_insert(original.clone()), "first insert commits");
assert_eq!(generation(), 1, "committed insert must signal");
let mut imposter = ToolDescriptor::from_capability(&cap("web_search"), &BTreeMap::new());
imposter.description = Some("imposter".to_string());
assert!(!reg.try_insert(imposter), "duplicate must be rejected");
assert_eq!(generation(), 1, "rejected duplicate must not signal");
assert_eq!(
reg.get("web_search").expect("entry present").description,
original.description,
"rejected duplicate must not replace the committed descriptor",
);
}
#[test]
fn descriptor_from_capability_handles_missing_metadata() {
let cap = cap("legacy");
let meta = BTreeMap::new();
let desc = ToolDescriptor::from_capability(&cap, &meta);
assert!(desc.description.is_none());
assert!(!desc.streaming);
assert!(desc.tags.is_empty());
}
#[test]
fn descriptor_tags_parsing_strips_whitespace_and_drops_empty() {
let cap = cap("messy");
let mut meta = BTreeMap::new();
meta.insert(tags_metadata_key("messy"), " a , b ,, c ,".to_string());
let desc = ToolDescriptor::from_capability(&cap, &meta);
assert_eq!(desc.tags, vec!["a", "b", "c"]);
}
#[test]
fn tool_event_serde_roundtrip_each_variant() {
let cases = vec![
ToolEvent::Start {
tool_id: "web_search".into(),
call_id: Some(42),
metadata: Some(serde_json::json!({"model": "claude-opus-4-7"})),
},
ToolEvent::Progress {
pct: Some(33.3),
message: Some("indexing".into()),
},
ToolEvent::Delta {
data: serde_json::json!({"token": "the"}),
},
ToolEvent::Result {
data: serde_json::json!({"results": ["a", "b"]}),
},
ToolEvent::Error {
code: "upstream_timeout".into(),
message: "took >30s".into(),
details: Some(serde_json::json!({"upstream": "anthropic"})),
},
];
for event in cases {
let encoded = serde_json::to_string(&event).expect("encode");
let decoded: ToolEvent = serde_json::from_str(&encoded).expect("decode");
assert_eq!(event, decoded, "round-trip must be byte-stable");
}
}
#[test]
fn tool_event_is_terminal_only_for_result_and_error() {
assert!(!ToolEvent::Start {
tool_id: "x".into(),
call_id: None,
metadata: None
}
.is_terminal());
assert!(!ToolEvent::Progress {
pct: None,
message: None
}
.is_terminal());
assert!(!ToolEvent::Delta {
data: serde_json::Value::Null
}
.is_terminal());
assert!(ToolEvent::Result {
data: serde_json::Value::Null
}
.is_terminal());
assert!(ToolEvent::Error {
code: "".into(),
message: "".into(),
details: None
}
.is_terminal());
}
#[test]
fn tool_event_optional_fields_omitted_when_none() {
let event = ToolEvent::Start {
tool_id: "x".into(),
call_id: None,
metadata: None,
};
let json = serde_json::to_string(&event).unwrap();
assert_eq!(json, r#"{"type":"start","tool_id":"x"}"#);
let event = ToolEvent::Progress {
pct: None,
message: None,
};
assert_eq!(
serde_json::to_string(&event).unwrap(),
r#"{"type":"progress"}"#
);
}
#[test]
fn from_capability_reads_pricing_terms_and_stays_free_without_the_key() {
let terms = "{\"object\":\"net.pricing.terms@1\"}";
let capability = cap("paid_tool");
let mut metadata = BTreeMap::new();
metadata.insert(pricing_terms_metadata_key("paid_tool"), terms.to_string());
let paid = ToolDescriptor::from_capability(&capability, &metadata);
assert_eq!(paid.pricing_terms.as_deref(), Some(terms));
let free = ToolDescriptor::from_capability(&capability, &BTreeMap::new());
assert_eq!(free.pricing_terms, None);
let json = serde_json::to_string(&free).unwrap();
assert!(!json.contains("pricing_terms"), "{json}");
}
fn descriptor(tool_id: &str) -> ToolDescriptor {
let cap = cap(tool_id);
ToolDescriptor::from_capability(&cap, &BTreeMap::new())
}
#[test]
fn tool_metadata_fetch_service_name_is_canonical() {
assert_eq!(TOOL_METADATA_FETCH_SERVICE, "tool.metadata.fetch");
}
#[test]
fn tool_metadata_response_serde_distinguishes_found_and_not_found() {
let found = ToolMetadataResponse::Found {
descriptor: descriptor("web_search"),
};
let not_found = ToolMetadataResponse::NotFound {
name: "missing".into(),
};
for resp in [&found, ¬_found] {
let encoded = serde_json::to_string(resp).unwrap();
let decoded: ToolMetadataResponse = serde_json::from_str(&encoded).unwrap();
assert_eq!(*resp, decoded, "round-trip must be byte-stable");
}
let found_json = serde_json::to_value(&found).unwrap();
assert_eq!(found_json["type"], "found");
let nf_json = serde_json::to_value(¬_found).unwrap();
assert_eq!(nf_json["type"], "not_found");
}
#[test]
fn tool_metadata_registry_insert_lookup_remove_roundtrip() {
let reg = ToolMetadataRegistry::new();
assert!(reg.is_empty());
assert_eq!(reg.len(), 0);
let desc = descriptor("web_search");
assert!(
reg.insert(desc.clone()).is_none(),
"first insert returns None"
);
assert_eq!(reg.len(), 1);
let got = reg.get("web_search").expect("get must find it");
assert_eq!(got, desc);
let prior = reg
.insert(desc.clone())
.expect("second insert returns prior");
assert_eq!(prior, desc);
let removed = reg.remove("web_search").expect("remove must find it");
assert_eq!(removed, desc);
assert!(reg.is_empty());
assert!(reg.get("web_search").is_none());
assert!(reg.remove("web_search").is_none());
}
#[test]
fn tool_metadata_registry_snapshot_returns_all_entries() {
let reg = ToolMetadataRegistry::new();
reg.insert(descriptor("a"));
reg.insert(descriptor("b"));
reg.insert(descriptor("c"));
let mut names: Vec<String> = reg.snapshot().iter().map(|d| d.tool_id.clone()).collect();
names.sort();
assert_eq!(names, vec!["a", "b", "c"]);
}
#[test]
fn tool_metadata_registry_snapshot_caches_until_mutation() {
let reg = ToolMetadataRegistry::new();
reg.insert(descriptor("a"));
let s1 = reg.snapshot();
let s2 = reg.snapshot();
assert!(
std::sync::Arc::ptr_eq(&s1, &s2),
"two consecutive snapshots without mutation must share the same Arc"
);
reg.insert(descriptor("b"));
let s3 = reg.snapshot();
assert!(
!std::sync::Arc::ptr_eq(&s1, &s3),
"insert must invalidate the cached snapshot"
);
let s4 = reg.snapshot();
assert!(
std::sync::Arc::ptr_eq(&s3, &s4),
"snapshot after insert must cache again"
);
reg.remove("a");
let s5 = reg.snapshot();
assert!(
!std::sync::Arc::ptr_eq(&s3, &s5),
"remove must invalidate the cached snapshot"
);
let s6 = reg.snapshot();
reg.remove("nonexistent");
let s7 = reg.snapshot();
assert!(
std::sync::Arc::ptr_eq(&s6, &s7),
"no-op remove must not invalidate the cached snapshot"
);
}
}