use std::collections::VecDeque;
use std::io::Write;
use futures::StreamExt;
use futures::stream::BoxStream;
use molo::agent::{Agent, AgentError};
use molo::memory::InMemoryMemory;
use molo::provider::OpenAiProvider;
use molo::tool::{SharedState, Tool, ToolError, ToolSchema};
use molo::{ChatRequest, Memory, Message, MessageChunk, Provider, StreamEvent, ToolCall};
use schemars::JsonSchema;
use serde::Deserialize;
#[derive(Debug, Deserialize, JsonSchema)]
struct CalcArgs {
#[schemars(description = "The math expression to evaluate, e.g. \"1 + 2 * 3\"")]
expression: String,
}
struct Calculator;
#[async_trait::async_trait]
impl Tool for Calculator {
fn schema(&self) -> ToolSchema {
ToolSchema {
name: "calculator".into(),
description: "Evaluates a math expression; supports basic arithmetic and parentheses, e.g. \"(1 + 2) * 3\".".into(),
parameters: serde_json::to_value(schemars::schema_for!(CalcArgs))
.expect("tool schema must serialize"),
}
}
async fn call(
&self,
arguments: serde_json::Value,
_state: &SharedState,
) -> Result<String, ToolError> {
let args: CalcArgs = serde_json::from_value(arguments)?;
let value =
evalexpr::eval(&args.expression).map_err(|e| ToolError::Execution(e.to_string()))?;
Ok(value.to_string())
}
}
struct CalculatorAgent {
provider: Box<dyn Provider>,
memory: Box<dyn Memory>,
tools: Vec<Box<dyn Tool>>,
max_tool_rounds: usize,
}
#[async_trait::async_trait]
impl Agent for CalculatorAgent {
async fn run(&mut self, input: &str) -> Result<String, AgentError> {
self.memory.record(Message::user(input)).await?;
let schemas: Vec<ToolSchema> = self.tools.iter().map(|t| t.schema()).collect();
for _ in 0..self.max_tool_rounds {
let response = self
.provider
.chat(ChatRequest {
messages: self.memory.context().await?,
tools: schemas.clone(),
..Default::default()
})
.await?;
let Message::Assistant {
content,
reasoning,
tool_calls,
} = response.message
else {
unreachable!("the reply must be an Assistant message by contract")
};
if !content.is_empty() || reasoning.is_some() || !tool_calls.is_empty() {
self.memory
.record(Message::Assistant {
content: content.clone(),
reasoning,
tool_calls: tool_calls.clone(),
})
.await?;
}
if tool_calls.is_empty() {
return Ok(content); }
for call in tool_calls {
let content = self.run_tool(&call.name, &call.arguments).await;
self.memory
.record(Message::ToolResult {
id: call.id,
content,
})
.await?;
}
}
Err(AgentError::TooManyToolRounds(self.max_tool_rounds))
}
async fn run_stream<'a>(
&'a mut self,
input: &'a str,
) -> Result<BoxStream<'a, Result<MessageChunk, AgentError>>, AgentError> {
self.memory.record(Message::user(input)).await?;
let schemas: Vec<ToolSchema> = self.tools.iter().map(|t| t.schema()).collect();
let max_rounds = self.max_tool_rounds;
let stream = futures::stream::unfold(
(
self,
0usize,
VecDeque::<Result<MessageChunk, AgentError>>::new(),
false,
),
move |(agent, rounds, mut pending, finished)| {
let schemas = schemas.clone();
async move {
if let Some(event) = pending.pop_front() {
return Some((event, (agent, rounds, pending, finished)));
}
if finished {
return None;
}
if rounds >= max_rounds {
return Some((
Err(AgentError::TooManyToolRounds(max_rounds)),
(agent, rounds, pending, true),
));
}
let messages = match agent.memory.context().await {
Ok(messages) => messages,
Err(e) => {
return Some((
Err(AgentError::Memory(e)),
(agent, rounds, pending, true),
));
}
};
let mut provider_stream = match agent
.provider
.stream_chat(ChatRequest {
messages,
tools: schemas,
..Default::default()
})
.await
{
Ok(stream) => stream,
Err(e) => {
return Some((
Err(AgentError::Provider(e)),
(agent, rounds, pending, true),
));
}
};
let mut text = String::new();
let mut reasoning = String::new();
let mut calls = Vec::new();
while let Some(event) = provider_stream.next().await {
match event {
Ok(StreamEvent::Delta(delta)) => text.push_str(&delta),
Ok(StreamEvent::Reasoning(chunk)) => reasoning.push_str(&chunk),
Ok(StreamEvent::ToolCall {
id,
name,
arguments,
}) => {
calls.push(ToolCall {
id,
name,
arguments,
});
}
Ok(StreamEvent::Done { .. }) => {}
Ok(_) => {}
Err(e) => {
return Some((
Err(AgentError::Provider(e)),
(agent, rounds + 1, pending, true),
));
}
}
}
if let Err(e) = agent
.memory
.record(Message::Assistant {
content: text.clone(),
reasoning: (!reasoning.is_empty()).then_some(reasoning),
tool_calls: calls.clone(),
})
.await
{
return Some((
Err(AgentError::Memory(e)),
(agent, rounds + 1, pending, true),
));
}
if !text.is_empty() {
pending.push_back(Ok(MessageChunk::Delta(text)));
}
for call in &calls {
pending.push_back(Ok(MessageChunk::ToolCall {
id: call.id.clone(),
name: call.name.clone(),
arguments: call.arguments.clone(),
}));
}
if calls.is_empty() {
pending.push_back(Ok(MessageChunk::Done(molo::RunSummary::default())));
return Some((
pending.pop_front().expect("just enqueued"),
(agent, rounds + 1, pending, true),
));
}
for call in calls {
let content = agent.run_tool(&call.name, &call.arguments).await;
pending.push_back(Ok(MessageChunk::ToolResult {
id: call.id,
name: call.name,
content,
}));
}
Some((
pending.pop_front().expect("just enqueued"),
(agent, rounds + 1, pending, false),
))
}
},
);
Ok(Box::pin(stream))
}
}
impl CalculatorAgent {
async fn run_tool(&self, name: &str, arguments: &str) -> String {
let Some(tool) = self.tools.iter().find(|t| t.schema().name == name) else {
return format!("tool not found: {name}");
};
let args = match serde_json::from_str(arguments) {
Ok(value) => value,
Err(e) => return format!("arguments are not valid JSON: {e}"),
};
match tool.call(args, &SharedState::new()).await {
Ok(text) => text,
Err(e) => format!("tool error: {e}"),
}
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
dotenvy::dotenv().ok();
let base_url =
std::env::var("MOLO_BASE_URL").unwrap_or_else(|_| "https://api.openai.com/v1".to_string());
let api_key = std::env::var("MOLO_API_KEY").unwrap_or_default();
let model = std::env::var("MOLO_MODEL").unwrap_or_else(|_| "gpt-4o-mini".to_string());
let provider = OpenAiProvider::new(base_url, api_key, model);
let tools: Vec<Box<dyn Tool>> = vec![Box::new(Calculator)];
let mut memory = InMemoryMemory::default();
memory
.record(Message::system(
"You are a helpful assistant. Use the calculator tool for calculations instead of doing math in your head.",
))
.await?;
let mut agent = CalculatorAgent {
provider: Box::new(provider),
memory: Box::new(memory),
tools,
max_tool_rounds: 10,
};
let mut input = String::new();
loop {
input.clear();
println!("user:");
let read = std::io::stdin().read_line(&mut input)?;
if read == 0 {
break; }
let input = input.trim();
if input.is_empty() {
continue;
}
if input == "exit" || input == "quit" {
break;
}
match agent.run_stream(input).await {
Ok(mut stream) => {
let mut prefix_printed = false;
while let Some(event) = stream.next().await {
match event {
Ok(MessageChunk::Delta(delta)) => {
if !prefix_printed {
print!("assistant: ");
std::io::stdout().flush()?;
prefix_printed = true;
}
print!("{delta}");
std::io::stdout().flush()?;
}
Ok(MessageChunk::ToolCall {
name, arguments, ..
}) => {
println!("\n → calling {name}, arguments: {arguments}");
}
Ok(MessageChunk::ToolResult { name, content, .. }) => {
println!(" → {name} returned: {content}");
}
Ok(MessageChunk::Done(_)) => break,
Ok(MessageChunk::Cancelled) => {
println!("\n[cancelled]");
break;
}
Ok(_) => {}
Err(e) => println!("\nerror: {e}"),
}
}
println!();
}
Err(e) => println!("error: {e}"),
}
}
Ok(())
}