use std::collections::HashMap;
use std::sync::Mutex;
use crate::protocol::{AgentCard, AgentSkill};
#[derive(Debug, thiserror::Error)]
pub enum RegistryError {
#[error("agent `{0}` is not registered")]
UnknownAgent(String),
#[error("agent `{0}` is already registered")]
AlreadyRegistered(String),
#[error("registry request failed: {0}")]
Http(String),
#[error("registry payload malformed: {0}")]
Parse(String),
}
#[derive(Debug, Default)]
pub struct AgentRegistry {
by_url: Mutex<HashMap<String, AgentCard>>,
}
impl AgentRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&self, card: AgentCard) -> Result<(), RegistryError> {
let mut by_url = self.by_url.lock().expect("registry lock poisoned");
if by_url.contains_key(&card.url) {
return Err(RegistryError::AlreadyRegistered(card.url));
}
by_url.insert(card.url.clone(), card);
Ok(())
}
pub fn upsert(&self, card: AgentCard) {
self.by_url
.lock()
.expect("registry lock poisoned")
.insert(card.url.clone(), card);
}
pub fn unregister(&self, url: &str) -> bool {
self.by_url
.lock()
.expect("registry lock poisoned")
.remove(url)
.is_some()
}
pub fn lookup(&self, url: &str) -> Option<AgentCard> {
self.by_url
.lock()
.expect("registry lock poisoned")
.get(url)
.cloned()
}
pub fn agents(&self) -> Vec<AgentCard> {
self.by_url
.lock()
.expect("registry lock poisoned")
.values()
.cloned()
.collect()
}
pub fn search_skill(&self, query: &str) -> Vec<AgentCard> {
let q = query.to_lowercase();
self.by_url
.lock()
.expect("registry lock poisoned")
.values()
.filter(|card| card.skills.iter().any(|s| skill_matches(s, &q)))
.cloned()
.collect()
}
pub fn filter_data_class(&self, class: &str) -> Vec<AgentCard> {
self.by_url
.lock()
.expect("registry lock poisoned")
.values()
.filter(|card| card.data_class.as_deref() == Some(class))
.cloned()
.collect()
}
pub fn len(&self) -> usize {
self.by_url.lock().expect("registry lock poisoned").len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
pub struct RegistryClient {
base_url: String,
http: reqwest::Client,
}
impl RegistryClient {
pub fn new(base_url: impl Into<String>) -> Self {
Self {
base_url: base_url.into(),
http: reqwest::Client::builder()
.no_proxy()
.build()
.expect("reqwest client builds"),
}
}
pub async fn fetch_catalog(&self) -> Result<Vec<AgentCard>, RegistryError> {
let url = format!("{}/registry.json", self.base_url);
let resp = self
.http
.get(&url)
.send()
.await
.map_err(|e| RegistryError::Http(e.to_string()))?;
if !resp.status().is_success() {
return Err(RegistryError::Http(format!(
"registry returned {}",
resp.status()
)));
}
resp.json::<Vec<AgentCard>>()
.await
.map_err(|e| RegistryError::Parse(e.to_string()))
}
pub async fn search_skill(&self, query: &str) -> Result<Vec<AgentCard>, RegistryError> {
let q = query.to_lowercase();
Ok(self
.fetch_catalog()
.await?
.into_iter()
.filter(|card| card.skills.iter().any(|s| skill_matches(s, &q)))
.collect())
}
}
fn skill_matches(skill: &AgentSkill, query: &str) -> bool {
skill.id.to_lowercase().contains(query)
|| skill.name.to_lowercase().contains(query)
|| skill.description.to_lowercase().contains(query)
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
fn card(url: &str, skill: &str, data_class: Option<&str>) -> AgentCard {
let mut card = AgentCard::new("agent", "a test agent", url).with_skill(AgentSkill::new(
skill,
skill,
format!("provides {skill}"),
));
if let Some(class) = data_class {
card = card.with_data_class(class);
}
card
}
#[test]
fn registry_register_lookup_unregister() {
let reg = AgentRegistry::new();
assert!(reg.is_empty());
reg.register(card("http://a", "summarize", None)).unwrap();
assert_eq!(reg.len(), 1);
assert!(reg.lookup("http://a").is_some());
assert!(reg.unregister("http://a"));
assert!(!reg.unregister("http://a"));
assert!(reg.is_empty());
}
#[test]
fn registry_rejects_duplicate_register_but_upsert_replaces() {
let reg = AgentRegistry::new();
reg.register(card("http://a", "summarize", None)).unwrap();
assert!(matches!(
reg.register(card("http://a", "translate", None)),
Err(RegistryError::AlreadyRegistered(_))
));
reg.upsert(card("http://a", "translate", None));
assert_eq!(reg.agents()[0].skills[0].id, "translate");
}
#[test]
fn registry_searches_skill_case_insensitively() {
let reg = AgentRegistry::new();
reg.upsert(card("http://sum", "summarize", None));
reg.upsert(card("http://translate", "translate", None));
let hits = reg.search_skill("Summ");
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].url, "http://sum");
assert!(reg.search_skill("nothing").is_empty());
}
#[test]
fn registry_filters_by_data_class() {
let reg = AgentRegistry::new();
reg.upsert(card("http://a", "summarize", Some("public")));
reg.upsert(card("http://b", "translate", Some("confidential")));
let public = reg.filter_data_class("public");
assert_eq!(public.len(), 1);
assert_eq!(public[0].url, "http://a");
assert!(reg.filter_data_class("internal").is_empty());
}
#[tokio::test]
async fn registry_client_fetches_remote_catalog() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let body = serde_json::json!([
{
"name": "summarizer",
"description": "summarizes",
"url": "http://sum",
"skills": [
{ "id": "summarize", "name": "summarize", "description": "provides summarize" }
],
"protocolVersion": "0.3.0",
"interfaces": {},
"securitySchemes": []
}
])
.to_string();
let resp = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nContent-Type: application/json\r\n\r\n{}",
body.len(),
body
);
while let Ok((mut stream, _)) = listener.accept().await {
let resp = resp.clone();
tokio::spawn(async move {
let mut chunk = [0u8; 4096];
let _ = tokio::time::timeout(
std::time::Duration::from_millis(200),
stream.read(&mut chunk),
)
.await;
let _ = stream.write_all(resp.as_bytes()).await;
});
}
});
let client = RegistryClient::new(format!("http://{addr}"));
let catalog = client.fetch_catalog().await.unwrap();
assert_eq!(catalog.len(), 1);
assert_eq!(catalog[0].url, "http://sum");
let hits = client.search_skill("summ").await.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].name, "summarizer");
}
}