1use crate::error::{AcpError, AcpResult};
4use hashbrown::HashMap;
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7use std::sync::Arc;
8use tokio::sync::RwLock;
9
10#[derive(Debug, Clone, Serialize, Deserialize)]
12pub struct AgentInfo {
13 id: String,
15
16 name: String,
18
19 pub(crate) base_url: String,
21
22 description: Option<String>,
24
25 capabilities: Vec<String>,
27
28 #[serde(default)]
30 metadata: HashMap<String, Value>,
31
32 #[serde(default = "default_online")]
34 online: bool,
35
36 last_seen: Option<String>,
38}
39
40fn default_online() -> bool {
41 true
42}
43
44#[derive(Clone)]
46pub struct AgentRegistry {
47 agents: Arc<RwLock<HashMap<String, AgentInfo>>>,
48}
49
50impl AgentRegistry {
51 pub(crate) fn new() -> Self {
53 Self { agents: Arc::new(RwLock::new(HashMap::new())) }
54 }
55
56 async fn register(&self, agent: AgentInfo) -> AcpResult<()> {
58 let mut agents = self.agents.write().await;
59 drop(agents.insert(agent.id.clone(), agent));
60 Ok(())
61 }
62
63 async fn unregister(&self, agent_id: &str) -> AcpResult<()> {
65 let mut agents = self.agents.write().await;
66 drop(agents.remove(agent_id));
67 Ok(())
68 }
69
70 pub(crate) async fn find(&self, agent_id: &str) -> AcpResult<AgentInfo> {
72 let agents = self.agents.read().await;
73 agents
74 .get(agent_id)
75 .cloned()
76 .ok_or_else(|| AcpError::AgentNotFound(agent_id.to_string()))
77 }
78
79 async fn find_by_capability(&self, capability: &str) -> AcpResult<Vec<AgentInfo>> {
81 let agents = self.agents.read().await;
82 let matching = agents
83 .values()
84 .filter(|a| a.online && a.capabilities.contains(&capability.to_string()))
85 .cloned()
86 .collect();
87 Ok(matching)
88 }
89
90 pub async fn list_all(&self) -> AcpResult<Vec<AgentInfo>> {
92 let agents = self.agents.read().await;
93 Ok(agents.values().cloned().collect())
94 }
95
96 pub async fn list_online(&self) -> AcpResult<Vec<AgentInfo>> {
98 let agents = self.agents.read().await;
99 Ok(agents.values().filter(|a| a.online).cloned().collect())
100 }
101
102 pub async fn update_status(&self, agent_id: &str, online: bool) -> AcpResult<()> {
104 let mut agents = self.agents.write().await;
105 if let Some(agent) = agents.get_mut(agent_id) {
106 agent.online = online;
107 agent.last_seen = Some(chrono::Utc::now().to_rfc3339());
108 Ok(())
109 } else {
110 Err(AcpError::AgentNotFound(agent_id.to_string()))
111 }
112 }
113
114 async fn count(&self) -> usize {
116 self.agents.read().await.len()
117 }
118
119 pub async fn clear(&self) {
121 self.agents.write().await.clear();
122 }
123}
124
125impl Default for AgentRegistry {
126 fn default() -> Self {
127 Self::new()
128 }
129}
130
131#[cfg(test)]
132mod tests {
133 use super::*;
134
135 #[tokio::test]
136 async fn test_agent_registry() {
137 let registry = AgentRegistry::new();
138
139 let agent = AgentInfo {
140 id: "test-agent".to_string(),
141 name: "Test Agent".to_string(),
142 base_url: "http://localhost:8080".to_string(),
143 description: Some("A test agent".to_string()),
144 capabilities: vec!["bash".to_string(), "python".to_string()],
145 metadata: HashMap::new(),
146 online: true,
147 last_seen: None,
148 };
149
150 registry.register(agent.clone()).await.unwrap();
151
152 let found = registry.find("test-agent").await.unwrap();
153 assert_eq!(found.id, "test-agent");
154
155 assert_eq!(registry.count().await, 1);
156
157 registry.unregister("test-agent").await.unwrap();
158 assert_eq!(registry.count().await, 0);
159 }
160
161 #[tokio::test]
162 async fn test_find_by_capability() {
163 let registry = AgentRegistry::new();
164
165 let agent1 = AgentInfo {
166 id: "agent-1".to_string(),
167 name: "Agent 1".to_string(),
168 base_url: "http://localhost:8080".to_string(),
169 description: None,
170 capabilities: vec!["bash".to_string()],
171 metadata: HashMap::new(),
172 online: true,
173 last_seen: None,
174 };
175
176 let agent2 = AgentInfo {
177 id: "agent-2".to_string(),
178 name: "Agent 2".to_string(),
179 base_url: "http://localhost:8081".to_string(),
180 description: None,
181 capabilities: vec!["bash".to_string(), "python".to_string()],
182 metadata: HashMap::new(),
183 online: true,
184 last_seen: None,
185 };
186
187 registry.register(agent1).await.unwrap();
188 registry.register(agent2).await.unwrap();
189
190 let bash_agents = registry.find_by_capability("bash").await.unwrap();
191 assert_eq!(bash_agents.len(), 2);
192
193 let python_agents = registry.find_by_capability("python").await.unwrap();
194 assert_eq!(python_agents.len(), 1);
195 }
196}