use crate::error::{SageError, SageResult};
use crate::llm::LlmClient;
use crate::session::{ProtocolViolation, SenderHandle, SessionId, SharedSessionRegistry};
use std::future::Future;
use tokio::sync::{mpsc, oneshot};
#[cfg(not(target_arch = "wasm32"))]
use tokio::task::JoinHandle;
#[cfg(not(target_arch = "wasm32"))]
pub struct AgentHandle<T> {
join: JoinHandle<SageResult<T>>,
message_tx: mpsc::Sender<Message>,
}
#[cfg(target_arch = "wasm32")]
pub struct AgentHandle<T> {
result_rx: oneshot::Receiver<SageResult<T>>,
message_tx: mpsc::Sender<Message>,
}
#[cfg(not(target_arch = "wasm32"))]
impl<T> AgentHandle<T> {
pub async fn result(self) -> SageResult<T> {
self.join.await?
}
}
#[cfg(target_arch = "wasm32")]
impl<T> AgentHandle<T> {
pub async fn result(self) -> SageResult<T> {
self.result_rx
.await
.map_err(|_| SageError::Agent("Agent task dropped".to_string()))?
}
}
impl<T> AgentHandle<T> {
pub async fn send<M>(&self, msg: M) -> SageResult<()>
where
M: serde::Serialize,
{
let message = Message::new(msg)?;
self.message_tx
.send(message)
.await
.map_err(|e| SageError::Agent(format!("Failed to send message: {e}")))
}
pub async fn send_message(&self, message: Message) -> SageResult<()> {
self.message_tx
.send(message)
.await
.map_err(|e| SageError::Agent(format!("Failed to send message: {e}")))
}
}
#[derive(Debug, Clone)]
pub struct Message {
pub payload: serde_json::Value,
pub session_id: Option<SessionId>,
pub sender: Option<SenderHandle>,
pub type_name: Option<String>,
}
impl Message {
pub fn new<T: serde::Serialize>(value: T) -> SageResult<Self> {
Ok(Self {
payload: serde_json::to_value(value)?,
session_id: None,
sender: None,
type_name: None,
})
}
pub fn with_session<T: serde::Serialize>(
value: T,
session_id: SessionId,
sender: SenderHandle,
type_name: impl Into<String>,
) -> SageResult<Self> {
Ok(Self {
payload: serde_json::to_value(value)?,
session_id: Some(session_id),
sender: Some(sender),
type_name: Some(type_name.into()),
})
}
#[must_use]
pub fn with_type_name(mut self, type_name: impl Into<String>) -> Self {
self.type_name = Some(type_name.into());
self
}
}
pub struct AgentContext<T> {
pub llm: LlmClient,
result_tx: Option<oneshot::Sender<T>>,
message_rx: mpsc::Receiver<Message>,
emitted: bool,
current_message: Option<Message>,
session_registry: SharedSessionRegistry,
agent_role: Option<String>,
}
impl<T> AgentContext<T> {
fn new(
llm: LlmClient,
result_tx: oneshot::Sender<T>,
message_rx: mpsc::Receiver<Message>,
session_registry: SharedSessionRegistry,
) -> Self {
Self {
llm,
result_tx: Some(result_tx),
message_rx,
emitted: false,
current_message: None,
session_registry,
agent_role: None,
}
}
pub fn set_role(&mut self, role: impl Into<String>) {
self.agent_role = Some(role.into());
}
#[must_use]
pub fn session_registry(&self) -> &SharedSessionRegistry {
&self.session_registry
}
pub fn emit(&mut self, value: T) -> SageResult<T>
where
T: Clone,
{
if self.emitted {
return Ok(value);
}
self.emitted = true;
if let Some(tx) = self.result_tx.take() {
let _ = tx.send(value.clone());
}
Ok(value)
}
pub async fn infer<R>(&self, prompt: &str) -> SageResult<R>
where
R: serde::de::DeserializeOwned,
{
self.llm.infer(prompt).await
}
pub async fn infer_string(&self, prompt: &str) -> SageResult<String> {
self.llm.infer_string(prompt).await
}
pub async fn receive<M>(&mut self) -> SageResult<M>
where
M: serde::de::DeserializeOwned,
{
let msg = self
.message_rx
.recv()
.await
.ok_or_else(|| SageError::Agent("Message channel closed".to_string()))?;
self.current_message = Some(msg.clone());
serde_json::from_value(msg.payload)
.map_err(|e| SageError::Agent(format!("Failed to deserialize message: {e}")))
}
#[cfg(not(target_arch = "wasm32"))]
pub async fn receive_timeout<M>(
&mut self,
timeout: std::time::Duration,
) -> SageResult<Option<M>>
where
M: serde::de::DeserializeOwned,
{
match tokio::time::timeout(timeout, self.message_rx.recv()).await {
Ok(Some(msg)) => {
self.current_message = Some(msg.clone());
let value = serde_json::from_value(msg.payload)
.map_err(|e| SageError::Agent(format!("Failed to deserialize message: {e}")))?;
Ok(Some(value))
}
Ok(None) => Err(SageError::Agent("Message channel closed".to_string())),
Err(_) => Ok(None), }
}
#[cfg(target_arch = "wasm32")]
pub async fn receive_timeout<M>(
&mut self,
timeout: std::time::Duration,
) -> SageResult<Option<M>>
where
M: serde::de::DeserializeOwned,
{
use futures::future::{select, Either};
use std::pin::pin;
let recv_fut = pin!(self.message_rx.recv());
let sleep_fut = pin!(sage_runtime_web::sleep(timeout));
match select(recv_fut, sleep_fut).await {
Either::Left((Some(msg), _)) => {
self.current_message = Some(msg.clone());
let value = serde_json::from_value(msg.payload)
.map_err(|e| SageError::Agent(format!("Failed to deserialize message: {e}")))?;
Ok(Some(value))
}
Either::Left((None, _)) => {
Err(SageError::Agent("Message channel closed".to_string()))
}
Either::Right((_, _)) => Ok(None), }
}
pub async fn receive_raw(&mut self) -> SageResult<Message> {
let msg = self
.message_rx
.recv()
.await
.ok_or_else(|| SageError::Agent("Message channel closed".to_string()))?;
self.current_message = Some(msg.clone());
Ok(msg)
}
pub fn set_current_message(&mut self, msg: Message) {
self.current_message = Some(msg);
}
pub fn clear_current_message(&mut self) {
self.current_message = None;
}
pub async fn reply<M: serde::Serialize>(&mut self, msg: M) -> SageResult<()> {
let current = self
.current_message
.as_ref()
.ok_or_else(|| SageError::from(ProtocolViolation::ReplyOutsideHandler))?;
let sender = current
.sender
.as_ref()
.ok_or_else(|| SageError::Agent("Message has no sender handle".to_string()))?;
sender.send(msg).await
}
pub async fn reply_with_protocol<M: serde::Serialize>(
&mut self,
msg: M,
msg_type: &str,
role: &str,
) -> SageResult<()> {
let current = self
.current_message
.as_ref()
.ok_or_else(|| SageError::from(ProtocolViolation::ReplyOutsideHandler))?;
if let Some(session_id) = current.session_id {
let mut registry = self.session_registry.write().await;
if let Some(session) = registry.get_mut(&session_id) {
if !session.state.can_send(msg_type, role) {
return Err(SageError::from(ProtocolViolation::UnexpectedMessage {
protocol: session.protocol.clone(),
expected: "valid reply".to_string(),
received: msg_type.to_string(),
state: session.state.state_name().to_string(),
}));
}
session.state.transition(msg_type)?;
}
}
let sender = current
.sender
.as_ref()
.ok_or_else(|| SageError::Agent("Message has no sender handle".to_string()))?;
sender.send(msg).await
}
pub async fn validate_protocol_receive(
&mut self,
msg_type: &str,
role: &str,
) -> SageResult<()> {
let current = match &self.current_message {
Some(msg) => msg,
None => return Ok(()), };
if let Some(session_id) = current.session_id {
let mut registry = self.session_registry.write().await;
if let Some(session) = registry.get_mut(&session_id) {
if !session.state.can_receive(msg_type, role) {
return Err(SageError::from(ProtocolViolation::UnexpectedMessage {
protocol: session.protocol.clone(),
expected: "valid message for current state".to_string(),
received: msg_type.to_string(),
state: session.state.state_name().to_string(),
}));
}
session.state.transition(msg_type)?;
if session.state.is_terminal() {
drop(registry);
self.session_registry.write().await.remove(&session_id);
}
}
}
Ok(())
}
pub async fn start_session(
&self,
protocol: String,
role: String,
state: Box<dyn crate::session::ProtocolStateMachine>,
partner: SenderHandle,
) -> SessionId {
let mut registry = self.session_registry.write().await;
let session_id = registry.next_id();
registry.start_session(session_id, protocol, role, state, partner);
session_id
}
#[must_use]
pub fn current_message(&self) -> Option<&Message> {
self.current_message.as_ref()
}
}
#[cfg(not(target_arch = "wasm32"))]
pub fn spawn<A, T, F>(agent: A) -> AgentHandle<T>
where
A: FnOnce(AgentContext<T>) -> F + Send + 'static,
F: Future<Output = SageResult<T>> + Send,
T: Send + 'static,
{
spawn_with_llm_config(agent, crate::llm::LlmConfig::from_env())
}
#[cfg(not(target_arch = "wasm32"))]
pub fn spawn_with_llm_config<A, T, F>(agent: A, llm_config: crate::llm::LlmConfig) -> AgentHandle<T>
where
A: FnOnce(AgentContext<T>) -> F + Send + 'static,
F: Future<Output = SageResult<T>> + Send,
T: Send + 'static,
{
let (result_tx, result_rx) = oneshot::channel();
let (message_tx, message_rx) = mpsc::channel(32);
let llm = LlmClient::new(llm_config);
let session_registry = crate::session::shared_registry();
let ctx = AgentContext::new(llm, result_tx, message_rx, session_registry);
let join = tokio::spawn(async move { agent(ctx).await });
drop(result_rx);
AgentHandle { join, message_tx }
}
#[cfg(target_arch = "wasm32")]
pub fn spawn<A, T, F>(agent: A) -> AgentHandle<T>
where
A: FnOnce(AgentContext<T>) -> F + 'static,
F: Future<Output = SageResult<T>> + 'static,
T: 'static,
{
spawn_with_llm_config(agent, crate::llm::LlmConfig::from_env())
}
#[cfg(target_arch = "wasm32")]
pub fn spawn_with_llm_config<A, T, F>(agent: A, llm_config: crate::llm::LlmConfig) -> AgentHandle<T>
where
A: FnOnce(AgentContext<T>) -> F + 'static,
F: Future<Output = SageResult<T>> + 'static,
T: 'static,
{
let (task_result_tx, task_result_rx) = oneshot::channel();
let (emit_tx, _emit_rx) = oneshot::channel();
let (message_tx, message_rx) = mpsc::channel(32);
let llm = LlmClient::new(llm_config);
let session_registry = crate::session::shared_registry();
let ctx = AgentContext::new(llm, emit_tx, message_rx, session_registry);
wasm_bindgen_futures::spawn_local(async move {
let result = agent(ctx).await;
let _ = task_result_tx.send(result);
});
AgentHandle {
result_rx: task_result_rx,
message_tx,
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde::{Deserialize, Serialize};
#[tokio::test]
async fn spawn_simple_agent() {
let handle = spawn(|mut ctx: AgentContext<i64>| async move { ctx.emit(42) });
let result = handle.result().await.expect("agent should succeed");
assert_eq!(result, 42);
}
#[tokio::test]
async fn spawn_agent_with_computation() {
let handle = spawn(|mut ctx: AgentContext<i64>| async move {
let sum = (1..=10).sum();
ctx.emit(sum)
});
let result = handle.result().await.expect("agent should succeed");
assert_eq!(result, 55);
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
struct TaskMessage {
id: u32,
content: String,
}
#[tokio::test]
async fn agent_receives_message() {
let handle = spawn(|mut ctx: AgentContext<String>| async move {
let msg: TaskMessage = ctx.receive().await?;
ctx.emit(format!("Got task {}: {}", msg.id, msg.content))
});
handle
.send(TaskMessage {
id: 42,
content: "Hello".to_string(),
})
.await
.expect("send should succeed");
let result = handle.result().await.expect("agent should succeed");
assert_eq!(result, "Got task 42: Hello");
}
#[tokio::test]
async fn agent_receives_multiple_messages() {
let handle = spawn(|mut ctx: AgentContext<i32>| async move {
let mut sum = 0;
for _ in 0..3 {
let n: i32 = ctx.receive().await?;
sum += n;
}
ctx.emit(sum)
});
for n in [10, 20, 30] {
handle.send(n).await.expect("send should succeed");
}
let result = handle.result().await.expect("agent should succeed");
assert_eq!(result, 60);
}
#[tokio::test]
async fn agent_receive_timeout() {
let handle = spawn(|mut ctx: AgentContext<String>| async move {
let result: Option<i32> = ctx
.receive_timeout(std::time::Duration::from_millis(10))
.await?;
match result {
Some(n) => ctx.emit(format!("Got {n}")),
None => ctx.emit("Timeout".to_string()),
}
});
let result = handle.result().await.expect("agent should succeed");
assert_eq!(result, "Timeout");
}
}