use std::collections::HashMap;
use std::sync::Arc;
use agent_base::{
AgentResult, AgentRuntime, Content, RunOutcome, RuntimeEvent, SessionId, Tool, ToolContext, ToolMetadata,
};
use agent_works::AgentBuilder;
use async_trait::async_trait;
use serde_json::Value;
use tokio::sync::{Mutex, mpsc};
type ToolCallResult = AgentResult<Vec<Content>>;
#[derive(Clone)]
pub struct ProtocolServer {
runtime: AgentRuntime,
slot: Arc<Mutex<Option<mpsc::UnboundedReceiver<ToolCallResult>>>>,
sessions: Arc<Mutex<HashMap<String, SessionId>>>,
}
impl ProtocolServer {
pub fn new(runtime: AgentRuntime) -> Self {
Self { runtime, slot: Arc::new(Mutex::new(None)), sessions: Arc::new(Mutex::new(HashMap::new())) }
}
pub fn from_builder(builder: AgentBuilder) -> Result<Self, agent_base::AgentError> {
let runtime = builder.build()?;
Ok(Self::new(runtime))
}
pub async fn register_tool(&self, name: String, description: String, parameters: Value) {
let proxy = ProxyTool { name, description, parameters, slot: self.slot.clone() };
let tools_arc = self.runtime.tools_mut();
let mut tools = tools_arc.write().await;
tools.register(proxy);
}
pub async fn prepare_tool_call(&self) -> mpsc::UnboundedSender<ToolCallResult> {
let (tx, rx) = mpsc::unbounded_channel();
*self.slot.lock().await = Some(rx);
tx
}
pub async fn create_session(&self, external_id: Option<String>) -> (SessionId, Option<String>) {
let sid = self.runtime.create_session().await;
let ext = external_id.clone();
(sid, ext)
}
pub async fn get_or_create_session(&self, external_id: Option<String>) -> SessionId {
if let Some(ref ext) = external_id {
let mut sessions = self.sessions.lock().await;
if let Some(sid) = sessions.get(ext) {
return sid.clone();
}
let (sid, _) = self.create_session(Some(ext.clone())).await;
sessions.insert(ext.clone(), sid.clone());
return sid;
}
self.create_session(None).await.0
}
pub fn subscribe_events(&self) -> tokio::sync::broadcast::Receiver<RuntimeEvent> {
self.runtime.subscribe_runtime_events()
}
pub async fn run_turn<F>(&self, sid: &SessionId, input: &str, f: F) -> AgentResult<RunOutcome>
where
F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
{
self.runtime.run_turn(sid.clone(), input, f).await
}
pub fn cancel(&self) {
self.runtime.cancel();
}
pub async fn list_tools(&self) -> Vec<ToolMetadata> {
let tools = self.runtime.tools_mut();
let registry = tools.read().await;
registry.metadatas()
}
}
struct ProxyTool {
name: String,
description: String,
parameters: Value,
slot: Arc<Mutex<Option<mpsc::UnboundedReceiver<ToolCallResult>>>>,
}
#[async_trait]
impl Tool for ProxyTool {
fn name(&self) -> &'static str {
Box::leak(self.name.clone().into_boxed_str())
}
fn description(&self) -> &'static str {
Box::leak(self.description.clone().into_boxed_str())
}
fn schema(&self) -> Value {
self.parameters.clone()
}
async fn call(&self, _args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
let mut rx = self
.slot
.lock()
.await
.take()
.ok_or_else(|| agent_base::AgentError::internal("no tool call slot prepared"))?;
match rx.recv().await {
Some(result) => result,
None => Ok(vec![Content::text("Tool call cancelled".to_string())]),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::builder::base_agent_builder;
use agent_base::ToolContext;
use async_trait::async_trait;
use futures_core::Stream;
use serde_json::json;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
struct StubClient;
struct StopStream {
state: u8,
}
impl Stream for StopStream {
type Item = Result<agent_base::StreamChunk, agent_base::llm_trait::LlmError>;
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match self.state {
0 => {
self.state = 1;
Poll::Ready(Some(Ok(agent_base::StreamChunk::Text("hello".to_string()))))
},
1 => {
self.state = 2;
Poll::Ready(Some(Ok(agent_base::StreamChunk::Stop { finish_reason: Some("stop".to_string()) })))
},
_ => Poll::Ready(None),
}
}
}
#[async_trait]
impl agent_base::llm_trait::LlmProvider for StubClient {
async fn stream(
&self,
_request: agent_base::llm_trait::ChatRequest,
) -> Result<agent_base::llm_trait::ChatStream, agent_base::llm_trait::LlmError> {
Ok(agent_base::llm_trait::ChatStream::new(Box::pin(StopStream { state: 0 })))
}
async fn chat(
&self,
_request: agent_base::llm_trait::ChatRequest,
) -> Result<agent_base::llm_trait::ChatResponse, agent_base::llm_trait::LlmError> {
Ok(agent_base::llm_trait::ChatResponse {
content: "hello".to_string(),
reasoning_content: None,
tool_calls: vec![],
usage: agent_base::UsageInfo::default(),
finish_reason: agent_base::llm_trait::FinishReason::Stop,
raw: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> agent_base::llm_trait::Capabilities {
agent_base::llm_trait::Capabilities::default()
}
fn info(&self) -> agent_base::llm_trait::ProviderInfo {
agent_base::llm_trait::ProviderInfo { name: "stub".to_string(), model: "stub".to_string(), version: None }
}
}
fn client() -> Arc<dyn agent_base::llm_trait::LlmProvider> {
Arc::new(StubClient)
}
fn runtime() -> agent_base::AgentRuntime {
base_agent_builder(client()).build().unwrap()
}
async fn register_echo(server: &ProtocolServer, rt: &agent_base::AgentRuntime) -> Arc<dyn agent_base::Tool> {
server.register_tool("echo".to_string(), "echo tool".to_string(), json!({ "type": "object" })).await;
let tools = rt.tools_mut();
let registry = tools.read().await;
registry.get("echo").expect("echo tool should be registered")
}
#[tokio::test(flavor = "multi_thread")]
async fn test_from_builder() {
let server = ProtocolServer::from_builder(base_agent_builder(client())).unwrap();
let _ = server;
}
#[tokio::test(flavor = "multi_thread")]
async fn test_register_and_list_tools() {
let rt = runtime();
let server = ProtocolServer::new(rt);
server.register_tool("echo".to_string(), "echo tool".to_string(), json!({ "type": "object" })).await;
let tools = server.list_tools().await;
let echo = tools.iter().find(|t| t.name == "echo").expect("echo tool should be listed");
assert_eq!(echo.description, "echo tool");
}
#[tokio::test(flavor = "multi_thread")]
async fn test_proxy_tool_call_without_slot_errors() {
let rt = runtime();
let server = ProtocolServer::new(rt.clone());
let tool = register_echo(&server, &rt).await;
let result = tool.call(&json!({}), &ToolContext::for_test()).await;
assert!(result.is_err());
}
#[tokio::test(flavor = "multi_thread")]
async fn test_proxy_tool_call_delivers_result() {
let rt = runtime();
let server = ProtocolServer::new(rt.clone());
let tool = register_echo(&server, &rt).await;
let tx = server.prepare_tool_call().await;
let args = json!({});
let ctx = ToolContext::for_test();
let call = tool.call(&args, &ctx);
tx.send(Ok(vec![Content::text("result".to_string())])).unwrap();
let result = call.await.unwrap();
assert_eq!(result.len(), 1);
match &result[0] {
Content::Text { text } => assert_eq!(text, "result"),
other => panic!("expected text content, got {other:?}"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_proxy_tool_call_cancelled_when_sender_dropped() {
let rt = runtime();
let server = ProtocolServer::new(rt.clone());
let tool = register_echo(&server, &rt).await;
let tx = server.prepare_tool_call().await;
drop(tx); let result = tool.call(&json!({}), &ToolContext::for_test()).await.unwrap();
assert_eq!(result.len(), 1);
match &result[0] {
Content::Text { text } => assert_eq!(text, "Tool call cancelled"),
other => panic!("expected text content, got {other:?}"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_create_session() {
let rt = runtime();
let server = ProtocolServer::new(rt);
let (_, ext) = server.create_session(None).await;
assert!(ext.is_none());
let (_, ext) = server.create_session(Some("ext".to_string())).await;
assert_eq!(ext.as_deref(), Some("ext"));
}
#[tokio::test(flavor = "multi_thread")]
async fn test_get_or_create_session_reuse() {
let rt = runtime();
let server = ProtocolServer::new(rt);
let a = server.get_or_create_session(Some("shared".to_string())).await;
let b = server.get_or_create_session(Some("shared".to_string())).await;
assert_eq!(a, b);
let c = server.get_or_create_session(None).await;
let d = server.get_or_create_session(None).await;
assert_ne!(c, d);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_subscribe_events_and_run_turn() {
let rt = runtime();
let server = ProtocolServer::new(rt);
let _rx = server.subscribe_events();
let sid = server.create_session(None).await.0;
let outcome = server.run_turn(&sid, "hi", |_| Ok(())).await;
assert!(outcome.is_ok());
server.cancel();
let _ = sid;
}
}