use tokio::sync::mpsc;
use agent_base::{AgentRuntime, RuntimeEvent, SessionId};
pub struct AgentHandle {
cmd_tx: mpsc::Sender<AgentCommand>,
event_rx: mpsc::UnboundedReceiver<RuntimeEvent>,
runtime: AgentRuntime,
default_session_id: Option<SessionId>,
}
enum AgentCommand {
RunTurn {
session_id: SessionId,
input: String,
},
}
#[derive(Debug)]
pub enum SendError {
ChannelClosed,
}
impl AgentHandle {
pub fn new(runtime: AgentRuntime) -> Self {
let (cmd_tx, cmd_rx) = mpsc::channel(32);
let (event_tx, event_rx) = mpsc::unbounded_channel();
let rt = runtime.clone();
tokio::spawn(async move {
let mut rx = cmd_rx;
while let Some(cmd) = rx.recv().await {
match cmd {
AgentCommand::RunTurn { session_id, input } => {
let tx = event_tx.clone();
let sid = session_id.clone();
let result = rt
.run_turn(session_id, &input, move |event| {
let _ = tx.send(event);
Ok(())
})
.await;
match &result {
Ok(_) => {}
Err(e) if e.is_cancelled() => {}
Err(e) => {
tracing::error!(error = %e, "run_turn failed");
let _ = event_tx.send(RuntimeEvent::RunFinished {
session_id: sid,
agent_id: None,
trace_id: None,
});
}
}
}
}
}
});
Self {
cmd_tx,
event_rx,
runtime,
default_session_id: None,
}
}
pub fn with_session(runtime: AgentRuntime, session_id: SessionId) -> Self {
let (cmd_tx, cmd_rx) = mpsc::channel(32);
let (event_tx, event_rx) = mpsc::unbounded_channel();
let rt = runtime.clone();
tokio::spawn(async move {
let mut rx = cmd_rx;
while let Some(cmd) = rx.recv().await {
match cmd {
AgentCommand::RunTurn { session_id, input } => {
let tx = event_tx.clone();
let sid = session_id.clone();
let result = rt
.run_turn(session_id, &input, move |event| {
let _ = tx.send(event);
Ok(())
})
.await;
match &result {
Ok(_) => {
}
Err(e) if e.is_cancelled() => {
}
Err(e) => {
tracing::error!(error = %e, "run_turn failed");
let _ = event_tx.send(RuntimeEvent::RunFinished {
session_id: sid,
agent_id: None,
trace_id: None,
});
}
}
}
}
}
});
Self {
cmd_tx,
event_rx,
runtime,
default_session_id: Some(session_id),
}
}
pub async fn send_input(&self, input: &str) -> Result<(), SendError> {
let session_id = match &self.default_session_id {
Some(id) => id.clone(),
None => self.runtime.create_session().await,
};
self.cmd_tx
.send(AgentCommand::RunTurn {
session_id,
input: input.to_string(),
})
.await
.map_err(|_| SendError::ChannelClosed)
}
pub async fn send_input_with_session(
&self,
input: &str,
session_id: SessionId,
) -> Result<(), SendError> {
self.cmd_tx
.send(AgentCommand::RunTurn {
session_id,
input: input.to_string(),
})
.await
.map_err(|_| SendError::ChannelClosed)
}
pub async fn recv_event(&mut self) -> Option<RuntimeEvent> {
self.event_rx.recv().await
}
pub fn try_recv_event(&mut self) -> Option<RuntimeEvent> {
self.event_rx.try_recv().ok()
}
pub fn cancel(&self) {
self.runtime.cancel();
}
pub fn runtime(&self) -> &AgentRuntime {
&self.runtime
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::time::Duration;
use agent_base::llm_trait::response::FinishReason;
use agent_base::llm_trait::types::UsageInfo;
use agent_base::llm_trait::{
Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, LlmProvider, ProviderInfo,
};
use agent_base::{AgentBuilder, StreamChunk};
struct StubProvider;
#[async_trait::async_trait]
impl LlmProvider for StubProvider {
async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
Ok(ChatStream::new(Box::pin(futures_util::stream::iter(vec![
Ok(StreamChunk::Text("hello".to_string())),
Ok(StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
}),
]))))
}
async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
Ok(ChatResponse {
content: "hello".to_string(),
tool_calls: vec![],
usage: UsageInfo::default(),
finish_reason: FinishReason::Stop,
raw: None,
reasoning_content: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> Capabilities {
Capabilities::default()
}
fn info(&self) -> ProviderInfo {
ProviderInfo {
name: "stub".to_string(),
model: "stub-model".to_string(),
version: None,
}
}
}
fn runtime() -> AgentRuntime {
AgentBuilder::new(Arc::new(StubProvider)).build().unwrap()
}
async fn wait_for_terminal(handle: &mut AgentHandle) -> Option<RuntimeEvent> {
let mut terminal = None;
for _ in 0..100 {
let ev = tokio::time::timeout(Duration::from_secs(5), handle.recv_event()).await;
match ev {
Ok(Some(e @ RuntimeEvent::RunFinished { .. }))
| Ok(Some(e @ RuntimeEvent::RunCancelled { .. })) => {
terminal = Some(e);
break;
}
Ok(Some(_)) => continue,
Ok(None) => break,
Err(_) => break,
}
}
terminal
}
#[tokio::test]
async fn test_send_input_and_recv_terminal() {
let mut handle = AgentHandle::new(runtime());
handle.send_input("hello").await.unwrap();
let terminal = wait_for_terminal(&mut handle).await;
assert!(
terminal.is_some(),
"expected a terminal event, got {terminal:?}"
);
}
#[tokio::test]
async fn test_send_input_with_session() {
let rt = runtime();
let session_id = rt.create_session().await;
let mut handle = AgentHandle::with_session(rt, session_id.clone());
handle
.send_input_with_session("hello", session_id)
.await
.unwrap();
let terminal = wait_for_terminal(&mut handle).await;
assert!(terminal.is_some());
}
#[tokio::test]
async fn test_runtime_accessor_and_cancel() {
let rt = runtime();
let handle = AgentHandle::new(rt.clone());
let session_id = handle.runtime().create_session().await;
assert!(rt.session(&session_id).await.is_some());
handle.cancel();
}
#[tokio::test]
async fn test_try_recv_event_initially_empty() {
let mut handle = AgentHandle::new(runtime());
assert!(handle.try_recv_event().is_none());
}
}