1use std::collections::HashMap;
8use std::sync::Arc;
9
10use agent_base::{
11 AgentResult, AgentRuntime, Content, RunOutcome, RuntimeEvent, SessionId, Tool, ToolContext, ToolMetadata,
12};
13use agent_works::AgentBuilder;
14use async_trait::async_trait;
15use serde_json::Value;
16use tokio::sync::{Mutex, mpsc};
17
18type ToolCallResult = AgentResult<Vec<Content>>;
21
22#[derive(Clone)]
29pub struct ProtocolServer {
30 runtime: AgentRuntime,
31 slot: Arc<Mutex<Option<mpsc::UnboundedReceiver<ToolCallResult>>>>,
34 sessions: Arc<Mutex<HashMap<String, SessionId>>>,
37}
38
39impl ProtocolServer {
40 pub fn new(runtime: AgentRuntime) -> Self {
42 Self { runtime, slot: Arc::new(Mutex::new(None)), sessions: Arc::new(Mutex::new(HashMap::new())) }
43 }
44
45 pub fn from_builder(builder: AgentBuilder) -> Result<Self, agent_base::AgentError> {
47 let runtime = builder.build()?;
48 Ok(Self::new(runtime))
49 }
50
51 pub async fn register_tool(&self, name: String, description: String, parameters: Value) {
56 let proxy = ProxyTool { name, description, parameters, slot: self.slot.clone() };
57 let tools_arc = self.runtime.tools_mut();
58 let mut tools = tools_arc.write().await;
59 tools.register(proxy);
60 }
61
62 pub async fn prepare_tool_call(&self) -> mpsc::UnboundedSender<ToolCallResult> {
65 let (tx, rx) = mpsc::unbounded_channel();
66 *self.slot.lock().await = Some(rx);
67 tx
68 }
69
70 pub async fn create_session(&self, external_id: Option<String>) -> (SessionId, Option<String>) {
72 let sid = self.runtime.create_session().await;
73 let ext = external_id.clone();
78 (sid, ext)
79 }
80
81 pub async fn get_or_create_session(&self, external_id: Option<String>) -> SessionId {
87 if let Some(ref ext) = external_id {
88 let mut sessions = self.sessions.lock().await;
89 if let Some(sid) = sessions.get(ext) {
90 return sid.clone();
91 }
92 let (sid, _) = self.create_session(Some(ext.clone())).await;
94 sessions.insert(ext.clone(), sid.clone());
95 return sid;
96 }
97 self.create_session(None).await.0
99 }
100
101 pub fn subscribe_events(&self) -> tokio::sync::broadcast::Receiver<RuntimeEvent> {
103 self.runtime.subscribe_runtime_events()
104 }
105
106 pub async fn run_turn<F>(&self, sid: &SessionId, input: &str, f: F) -> AgentResult<RunOutcome>
108 where
109 F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
110 {
111 self.runtime.run_turn(sid.clone(), input, f).await
112 }
113
114 pub fn cancel(&self) {
116 self.runtime.cancel();
117 }
118
119 pub async fn list_tools(&self) -> Vec<ToolMetadata> {
121 let tools = self.runtime.tools_mut();
122 let registry = tools.read().await;
123 registry.metadatas()
124 }
125}
126
127struct ProxyTool {
130 name: String,
131 description: String,
132 parameters: Value,
133 slot: Arc<Mutex<Option<mpsc::UnboundedReceiver<ToolCallResult>>>>,
134}
135
136#[async_trait]
137impl Tool for ProxyTool {
138 fn name(&self) -> &'static str {
139 Box::leak(self.name.clone().into_boxed_str())
140 }
141
142 fn description(&self) -> &'static str {
143 Box::leak(self.description.clone().into_boxed_str())
144 }
145
146 fn schema(&self) -> Value {
147 self.parameters.clone()
148 }
149
150 async fn call(&self, _args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
151 let mut rx = self
152 .slot
153 .lock()
154 .await
155 .take()
156 .ok_or_else(|| agent_base::AgentError::internal("no tool call slot prepared"))?;
157
158 match rx.recv().await {
159 Some(result) => result,
160 None => Ok(vec![Content::text("Tool call cancelled".to_string())]),
161 }
162 }
163}
164
165#[cfg(test)]
166mod tests {
167 use super::*;
168 use crate::agent::builder::base_agent_builder;
169 use agent_base::ToolContext;
170 use async_trait::async_trait;
171 use futures_core::Stream;
172 use serde_json::json;
173 use std::pin::Pin;
174 use std::sync::Arc;
175 use std::task::{Context, Poll};
176
177 struct StubClient;
178
179 struct StopStream {
182 state: u8,
183 }
184
185 impl Stream for StopStream {
186 type Item = Result<agent_base::StreamChunk, agent_base::llm_trait::LlmError>;
187
188 fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
189 match self.state {
190 0 => {
191 self.state = 1;
192 Poll::Ready(Some(Ok(agent_base::StreamChunk::Text("hello".to_string()))))
193 },
194 1 => {
195 self.state = 2;
196 Poll::Ready(Some(Ok(agent_base::StreamChunk::Stop { finish_reason: Some("stop".to_string()) })))
197 },
198 _ => Poll::Ready(None),
199 }
200 }
201 }
202
203 #[async_trait]
204 impl agent_base::llm_trait::LlmProvider for StubClient {
205 async fn stream(
206 &self,
207 _request: agent_base::llm_trait::ChatRequest,
208 ) -> Result<agent_base::llm_trait::ChatStream, agent_base::llm_trait::LlmError> {
209 Ok(agent_base::llm_trait::ChatStream::new(Box::pin(StopStream { state: 0 })))
210 }
211 async fn chat(
212 &self,
213 _request: agent_base::llm_trait::ChatRequest,
214 ) -> Result<agent_base::llm_trait::ChatResponse, agent_base::llm_trait::LlmError> {
215 Ok(agent_base::llm_trait::ChatResponse {
216 content: "hello".to_string(),
217 reasoning_content: None,
218 tool_calls: vec![],
219 usage: agent_base::UsageInfo::default(),
220 finish_reason: agent_base::llm_trait::FinishReason::Stop,
221 raw: None,
222 thinking_signature: None,
223 })
224 }
225 fn capabilities(&self) -> agent_base::llm_trait::Capabilities {
226 agent_base::llm_trait::Capabilities::default()
227 }
228 fn info(&self) -> agent_base::llm_trait::ProviderInfo {
229 agent_base::llm_trait::ProviderInfo { name: "stub".to_string(), model: "stub".to_string(), version: None }
230 }
231 }
232
233 fn client() -> Arc<dyn agent_base::llm_trait::LlmProvider> {
234 Arc::new(StubClient)
235 }
236
237 fn runtime() -> agent_base::AgentRuntime {
238 base_agent_builder(client()).build().unwrap()
239 }
240
241 async fn register_echo(server: &ProtocolServer, rt: &agent_base::AgentRuntime) -> Arc<dyn agent_base::Tool> {
243 server.register_tool("echo".to_string(), "echo tool".to_string(), json!({ "type": "object" })).await;
244 let tools = rt.tools_mut();
245 let registry = tools.read().await;
246 registry.get("echo").expect("echo tool should be registered")
247 }
248
249 #[tokio::test(flavor = "multi_thread")]
250 async fn test_from_builder() {
251 let server = ProtocolServer::from_builder(base_agent_builder(client())).unwrap();
252 let _ = server;
253 }
254
255 #[tokio::test(flavor = "multi_thread")]
256 async fn test_register_and_list_tools() {
257 let rt = runtime();
258 let server = ProtocolServer::new(rt);
259 server.register_tool("echo".to_string(), "echo tool".to_string(), json!({ "type": "object" })).await;
260
261 let tools = server.list_tools().await;
262 let echo = tools.iter().find(|t| t.name == "echo").expect("echo tool should be listed");
263 assert_eq!(echo.description, "echo tool");
264 }
265
266 #[tokio::test(flavor = "multi_thread")]
267 async fn test_proxy_tool_call_without_slot_errors() {
268 let rt = runtime();
269 let server = ProtocolServer::new(rt.clone());
270 let tool = register_echo(&server, &rt).await;
271
272 let result = tool.call(&json!({}), &ToolContext::for_test()).await;
273 assert!(result.is_err());
274 }
275
276 #[tokio::test(flavor = "multi_thread")]
277 async fn test_proxy_tool_call_delivers_result() {
278 let rt = runtime();
279 let server = ProtocolServer::new(rt.clone());
280 let tool = register_echo(&server, &rt).await;
281
282 let tx = server.prepare_tool_call().await;
283 let args = json!({});
284 let ctx = ToolContext::for_test();
285 let call = tool.call(&args, &ctx);
286 tx.send(Ok(vec![Content::text("result".to_string())])).unwrap();
287 let result = call.await.unwrap();
288
289 assert_eq!(result.len(), 1);
290 match &result[0] {
291 Content::Text { text } => assert_eq!(text, "result"),
292 other => panic!("expected text content, got {other:?}"),
293 }
294 }
295
296 #[tokio::test(flavor = "multi_thread")]
297 async fn test_proxy_tool_call_cancelled_when_sender_dropped() {
298 let rt = runtime();
299 let server = ProtocolServer::new(rt.clone());
300 let tool = register_echo(&server, &rt).await;
301
302 let tx = server.prepare_tool_call().await;
303 drop(tx); let result = tool.call(&json!({}), &ToolContext::for_test()).await.unwrap();
305
306 assert_eq!(result.len(), 1);
307 match &result[0] {
308 Content::Text { text } => assert_eq!(text, "Tool call cancelled"),
309 other => panic!("expected text content, got {other:?}"),
310 }
311 }
312
313 #[tokio::test(flavor = "multi_thread")]
314 async fn test_create_session() {
315 let rt = runtime();
316 let server = ProtocolServer::new(rt);
317
318 let (_, ext) = server.create_session(None).await;
319 assert!(ext.is_none());
320
321 let (_, ext) = server.create_session(Some("ext".to_string())).await;
322 assert_eq!(ext.as_deref(), Some("ext"));
323 }
324
325 #[tokio::test(flavor = "multi_thread")]
326 async fn test_get_or_create_session_reuse() {
327 let rt = runtime();
328 let server = ProtocolServer::new(rt);
329
330 let a = server.get_or_create_session(Some("shared".to_string())).await;
331 let b = server.get_or_create_session(Some("shared".to_string())).await;
332 assert_eq!(a, b);
333
334 let c = server.get_or_create_session(None).await;
335 let d = server.get_or_create_session(None).await;
336 assert_ne!(c, d);
337 }
338
339 #[tokio::test(flavor = "multi_thread")]
340 async fn test_subscribe_events_and_run_turn() {
341 let rt = runtime();
342 let server = ProtocolServer::new(rt);
343
344 let _rx = server.subscribe_events();
345 let sid = server.create_session(None).await.0;
346 let outcome = server.run_turn(&sid, "hi", |_| Ok(())).await;
347 assert!(outcome.is_ok());
348
349 server.cancel();
350 let _ = sid;
351 }
352}