1use tokio::sync::mpsc;
2
3use agent_base::{AgentRuntime, RuntimeEvent, SessionId};
4
5pub struct AgentHandle {
36 cmd_tx: mpsc::Sender<AgentCommand>,
37 event_rx: mpsc::UnboundedReceiver<RuntimeEvent>,
38 runtime: AgentRuntime,
39 default_session_id: Option<SessionId>,
40}
41
42enum AgentCommand {
43 RunTurn {
44 session_id: SessionId,
45 input: String,
46 },
47}
48
49#[derive(Debug)]
50pub enum SendError {
51 ChannelClosed,
52}
53
54impl AgentHandle {
55 pub fn new(runtime: AgentRuntime) -> Self {
57 let (cmd_tx, cmd_rx) = mpsc::channel(32);
58 let (event_tx, event_rx) = mpsc::unbounded_channel();
59 let rt = runtime.clone();
60
61 tokio::spawn(async move {
63 let mut rx = cmd_rx;
64 while let Some(cmd) = rx.recv().await {
65 match cmd {
66 AgentCommand::RunTurn { session_id, input } => {
67 let tx = event_tx.clone();
68 let sid = session_id.clone();
69 let result = rt
70 .run_turn(session_id, &input, move |event| {
71 let _ = tx.send(event);
72 Ok(())
73 })
74 .await;
75
76 match &result {
77 Ok(_) => {}
78 Err(e) if e.is_cancelled() => {}
79 Err(e) => {
80 tracing::error!(error = %e, "run_turn failed");
81 let _ = event_tx.send(RuntimeEvent::RunFinished {
82 session_id: sid,
83 agent_id: None,
84 trace_id: None,
85 });
86 }
87 }
88 }
89 }
90 }
91 });
92
93 Self {
94 cmd_tx,
95 event_rx,
96 runtime,
97 default_session_id: None,
98 }
99 }
100
101 pub fn with_session(runtime: AgentRuntime, session_id: SessionId) -> Self {
104 let (cmd_tx, cmd_rx) = mpsc::channel(32);
105 let (event_tx, event_rx) = mpsc::unbounded_channel();
106 let rt = runtime.clone();
107
108 tokio::spawn(async move {
110 let mut rx = cmd_rx;
111 while let Some(cmd) = rx.recv().await {
112 match cmd {
113 AgentCommand::RunTurn { session_id, input } => {
114 let tx = event_tx.clone();
115 let sid = session_id.clone();
116 let result = rt
117 .run_turn(session_id, &input, move |event| {
118 let _ = tx.send(event);
119 Ok(())
120 })
121 .await;
122
123 match &result {
125 Ok(_) => {
126 }
128 Err(e) if e.is_cancelled() => {
129 }
131 Err(e) => {
132 tracing::error!(error = %e, "run_turn failed");
134 let _ = event_tx.send(RuntimeEvent::RunFinished {
135 session_id: sid,
136 agent_id: None,
137 trace_id: None,
138 });
139 }
140 }
141 }
142 }
143 }
144 });
145
146 Self {
147 cmd_tx,
148 event_rx,
149 runtime,
150 default_session_id: Some(session_id),
151 }
152 }
153
154 pub async fn send_input(&self, input: &str) -> Result<(), SendError> {
157 let session_id = match &self.default_session_id {
158 Some(id) => id.clone(),
159 None => self.runtime.create_session().await,
160 };
161 self.cmd_tx
162 .send(AgentCommand::RunTurn {
163 session_id,
164 input: input.to_string(),
165 })
166 .await
167 .map_err(|_| SendError::ChannelClosed)
168 }
169
170 pub async fn send_input_with_session(
172 &self,
173 input: &str,
174 session_id: SessionId,
175 ) -> Result<(), SendError> {
176 self.cmd_tx
177 .send(AgentCommand::RunTurn {
178 session_id,
179 input: input.to_string(),
180 })
181 .await
182 .map_err(|_| SendError::ChannelClosed)
183 }
184
185 pub async fn recv_event(&mut self) -> Option<RuntimeEvent> {
187 self.event_rx.recv().await
188 }
189
190 pub fn try_recv_event(&mut self) -> Option<RuntimeEvent> {
192 self.event_rx.try_recv().ok()
193 }
194
195 pub fn cancel(&self) {
197 self.runtime.cancel();
198 }
199
200 pub fn runtime(&self) -> &AgentRuntime {
202 &self.runtime
203 }
204}
205
206#[cfg(test)]
207mod tests {
208 use super::*;
209 use std::pin::Pin;
210 use std::sync::Arc;
211 use std::time::Duration;
212
213 use agent_base::{
214 AgentBuilder, AgentResult, ChatMessage, LlmCapabilities, ReasoningConfig, ResponseFormat,
215 StreamChunk, StreamClient,
216 };
217 use futures_core::Stream;
218 use serde_json::Value;
219
220 struct StubClient;
221
222 #[async_trait::async_trait]
223 impl StreamClient for StubClient {
224 async fn stream(
225 &self,
226 _messages: &[ChatMessage],
227 _tools: &[Value],
228 _reasoning: Option<&ReasoningConfig>,
229 _response_format: Option<&ResponseFormat>,
230 ) -> AgentResult<Pin<Box<dyn Stream<Item = AgentResult<StreamChunk>> + Send>>> {
231 Ok(Box::pin(futures_util::stream::iter(vec![
232 Ok(StreamChunk::Text("hello".to_string())),
233 Ok(StreamChunk::Stop {
234 finish_reason: Some("stop".to_string()),
235 }),
236 ])))
237 }
238
239 fn capabilities(&self) -> LlmCapabilities {
240 LlmCapabilities::default()
241 }
242 }
243
244 fn runtime() -> AgentRuntime {
245 AgentBuilder::new(Arc::new(StubClient)).build().unwrap()
246 }
247
248 async fn wait_for_terminal(handle: &mut AgentHandle) -> Option<RuntimeEvent> {
249 let mut terminal = None;
250 for _ in 0..100 {
251 let ev = tokio::time::timeout(Duration::from_secs(5), handle.recv_event()).await;
252 match ev {
253 Ok(Some(e @ RuntimeEvent::RunFinished { .. }))
254 | Ok(Some(e @ RuntimeEvent::RunCancelled { .. })) => {
255 terminal = Some(e);
256 break;
257 }
258 Ok(Some(_)) => continue,
259 Ok(None) => break,
260 Err(_) => break,
261 }
262 }
263 terminal
264 }
265
266 #[tokio::test]
267 async fn test_send_input_and_recv_terminal() {
268 let mut handle = AgentHandle::new(runtime());
269 handle.send_input("hello").await.unwrap();
270 let terminal = wait_for_terminal(&mut handle).await;
271 assert!(
272 terminal.is_some(),
273 "expected a terminal event, got {terminal:?}"
274 );
275 }
276
277 #[tokio::test]
278 async fn test_send_input_with_session() {
279 let rt = runtime();
280 let session_id = rt.create_session().await;
281 let mut handle = AgentHandle::with_session(rt, session_id.clone());
282 handle
283 .send_input_with_session("hello", session_id)
284 .await
285 .unwrap();
286 let terminal = wait_for_terminal(&mut handle).await;
287 assert!(terminal.is_some());
288 }
289
290 #[tokio::test]
291 async fn test_runtime_accessor_and_cancel() {
292 let rt = runtime();
293 let handle = AgentHandle::new(rt.clone());
294 let session_id = handle.runtime().create_session().await;
296 assert!(rt.session(&session_id).await.is_some());
297 handle.cancel();
299 }
300
301 #[tokio::test]
302 async fn test_try_recv_event_initially_empty() {
303 let mut handle = AgentHandle::new(runtime());
304 assert!(handle.try_recv_event().is_none());
305 }
306}