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
210 use std::sync::Arc;
211 use std::time::Duration;
212
213 use agent_base::llm_trait::response::FinishReason;
214 use agent_base::llm_trait::types::UsageInfo;
215 use agent_base::llm_trait::{
216 Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, LlmProvider, ProviderInfo,
217 };
218 use agent_base::{AgentBuilder, StreamChunk};
219
220 struct StubProvider;
221
222 #[async_trait::async_trait]
223 impl LlmProvider for StubProvider {
224 async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
225 Ok(ChatStream::new(Box::pin(futures_util::stream::iter(vec![
226 Ok(StreamChunk::Text("hello".to_string())),
227 Ok(StreamChunk::Stop {
228 finish_reason: Some("stop".to_string()),
229 }),
230 ]))))
231 }
232
233 async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
234 Ok(ChatResponse {
235 content: "hello".to_string(),
236 tool_calls: vec![],
237 usage: UsageInfo::default(),
238 finish_reason: FinishReason::Stop,
239 raw: None,
240 reasoning_content: None,
241 thinking_signature: None,
242 })
243 }
244
245 fn capabilities(&self) -> Capabilities {
246 Capabilities::default()
247 }
248
249 fn info(&self) -> ProviderInfo {
250 ProviderInfo {
251 name: "stub".to_string(),
252 model: "stub-model".to_string(),
253 version: None,
254 }
255 }
256 }
257
258 fn runtime() -> AgentRuntime {
259 AgentBuilder::new(Arc::new(StubProvider)).build().unwrap()
260 }
261
262 async fn wait_for_terminal(handle: &mut AgentHandle) -> Option<RuntimeEvent> {
263 let mut terminal = None;
264 for _ in 0..100 {
265 let ev = tokio::time::timeout(Duration::from_secs(5), handle.recv_event()).await;
266 match ev {
267 Ok(Some(e @ RuntimeEvent::RunFinished { .. }))
268 | Ok(Some(e @ RuntimeEvent::RunCancelled { .. })) => {
269 terminal = Some(e);
270 break;
271 }
272 Ok(Some(_)) => continue,
273 Ok(None) => break,
274 Err(_) => break,
275 }
276 }
277 terminal
278 }
279
280 #[tokio::test]
281 async fn test_send_input_and_recv_terminal() {
282 let mut handle = AgentHandle::new(runtime());
283 handle.send_input("hello").await.unwrap();
284 let terminal = wait_for_terminal(&mut handle).await;
285 assert!(
286 terminal.is_some(),
287 "expected a terminal event, got {terminal:?}"
288 );
289 }
290
291 #[tokio::test]
292 async fn test_send_input_with_session() {
293 let rt = runtime();
294 let session_id = rt.create_session().await;
295 let mut handle = AgentHandle::with_session(rt, session_id.clone());
296 handle
297 .send_input_with_session("hello", session_id)
298 .await
299 .unwrap();
300 let terminal = wait_for_terminal(&mut handle).await;
301 assert!(terminal.is_some());
302 }
303
304 #[tokio::test]
305 async fn test_runtime_accessor_and_cancel() {
306 let rt = runtime();
307 let handle = AgentHandle::new(rt.clone());
308 let session_id = handle.runtime().create_session().await;
310 assert!(rt.session(&session_id).await.is_some());
311 handle.cancel();
313 }
314
315 #[tokio::test]
316 async fn test_try_recv_event_initially_empty() {
317 let mut handle = AgentHandle::new(runtime());
318 assert!(handle.try_recv_event().is_none());
319 }
320}