Skip to main content

phi_agent/bridge/
server.rs

1//! Protocol server — adapts [`AgentRuntime`] to the bridge protocol.
2//!
3//! Tool calls use a single-slot pattern: the serve loop pushes a receiver
4//! before each tool call, and ProxyTool pops it.  This handles sequential
5//! tool calls cleanly; parallel calls can be added later via a FIFO queue.
6
7use 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
18/// A tool-call result delivered back through the bridge slot: either the
19/// tool's content or an execution error.
20type ToolCallResult = AgentResult<Vec<Content>>;
21
22/// Bridge protocol server — wraps an [`AgentRuntime`] and exposes it over
23/// the NDJSON bridge protocol for external SDK consumption.
24///
25/// Manages sessions, tool registration via proxy tools, and event forwarding.
26/// Tool calls use a single-slot pattern: the serve loop pushes a receiver
27/// before each tool call, and ProxyTool pops it.
28#[derive(Clone)]
29pub struct ProtocolServer {
30    runtime: AgentRuntime,
31    /// Single-slot: the next tool call's response receiver.
32    /// serve loop pushes, ProxyTool pops.
33    slot: Arc<Mutex<Option<mpsc::UnboundedReceiver<ToolCallResult>>>>,
34    /// Map external_id → SessionId so that runs with the same
35    /// external_id reuse the same session.
36    sessions: Arc<Mutex<HashMap<String, SessionId>>>,
37}
38
39impl ProtocolServer {
40    /// Wrap an existing [`AgentRuntime`] in a protocol server.
41    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    /// Build a protocol server from an [`AgentBuilder`].
46    pub fn from_builder(builder: AgentBuilder) -> Result<Self, agent_base::AgentError> {
47        let runtime = builder.build()?;
48        Ok(Self::new(runtime))
49    }
50
51    /// Register a tool implemented on the SDK side.
52    ///
53    /// The tool's `call` will block until the SDK sends a `tool_result`
54    /// message through the bridge.
55    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    /// Set up the response channel for the NEXT tool call.
63    /// Returns the sender — keep it; send the result when the SDK replies.
64    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    /// Create a new session, optionally with an external ID for reuse.
71    pub async fn create_session(&self, external_id: Option<String>) -> (SessionId, Option<String>) {
72        let sid = self.runtime.create_session().await;
73        // NOTE: We intentionally do NOT set sid.external_id because
74        // agent_base's run_turn() hangs when external_id is Some.
75        // Session reuse is handled by get_or_create_session() which
76        // maintains its own external_id → SessionId map.
77        let ext = external_id.clone();
78        (sid, ext)
79    }
80
81    /// Get or create a session by external_id.
82    ///
83    /// If ``external_id`` is ``Some`` and a session with that id already
84    /// exists, it is reused (preserving conversation history).  Otherwise
85    /// a new session is created and registered.
86    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            // Create new and register
93            let (sid, _) = self.create_session(Some(ext.clone())).await;
94            sessions.insert(ext.clone(), sid.clone());
95            return sid;
96        }
97        // No external_id — always create new
98        self.create_session(None).await.0
99    }
100
101    /// Subscribe to runtime events broadcast by the agent.
102    pub fn subscribe_events(&self) -> tokio::sync::broadcast::Receiver<RuntimeEvent> {
103        self.runtime.subscribe_runtime_events()
104    }
105
106    /// Run a turn on the given session, forwarding events to the callback.
107    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    /// Cancel the currently running turn.
115    pub fn cancel(&self) {
116        self.runtime.cancel();
117    }
118
119    /// List all registered tools with their metadata, sorted by name.
120    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
127// ── ProxyTool ─────────────────────────────────────────────────────────
128
129struct 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    /// Yields one `Text` chunk, one `Stop` chunk, then ends — enough for the
180    /// react loop to complete a turn.
181    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    /// Register an "echo" proxy tool and return its handle from the registry.
242    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); // dropping the only sender closes the channel
304        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}