use futures::stream::StreamExt;
use molo::provider::OpenAiProvider;
use molo::tool::{SharedState, Tool, ToolError, ToolSchema};
use molo::{
ChatRequest, Message, MessageChannel, MpscChannel, Provider, StreamEvent, ToolCall,
ToolRegistry,
};
use schemars::JsonSchema;
use serde::Deserialize;
use std::io::Write;
use std::sync::Arc;
#[derive(Debug, Deserialize, JsonSchema)]
struct AskArgs {
#[schemars(description = "The question to ask the expert")]
question: String,
}
struct AskExpertTool {
channel: Arc<dyn MessageChannel>,
}
#[async_trait::async_trait]
impl Tool for AskExpertTool {
fn schema(&self) -> ToolSchema {
ToolSchema {
name: "ask_expert".into(),
description:
"Hands the question to the expert agent for an answer; returns the expert's reply."
.into(),
parameters: serde_json::to_value(schemars::schema_for!(AskArgs))
.expect("tool schema must serialize"),
}
}
async fn call(
&self,
arguments: serde_json::Value,
_state: &SharedState,
) -> Result<String, ToolError> {
let args: AskArgs = serde_json::from_value(arguments)?;
self.channel
.ask(&args.question)
.await
.map_err(|e| ToolError::Execution(e.to_string()))
}
}
struct Turn {
messages: Vec<Message>,
}
enum Mode {
Stream,
Chat,
}
const MAX_TOOL_ROUNDS: usize = 10;
#[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 mode = match std::env::args().nth(1).as_deref() {
Some("chat") => Mode::Chat,
_ => Mode::Stream,
};
let provider = OpenAiProvider::new(base_url.clone(), api_key.clone(), model.clone());
let (expert_ask, expert_side) = MpscChannel::pair();
let expert_provider = OpenAiProvider::new(base_url, api_key, model);
tokio::spawn(async move { expert_loop(&expert_provider, expert_side).await });
let mut registry = ToolRegistry::new();
registry.register(Calculator).register(AskExpertTool {
channel: Arc::new(expert_ask),
});
let tool_schemas = registry.schemas();
let mut messages = vec![Message::system(
"You are a helpful assistant. Use the calculator tool when you need to \
calculate; when you run into a specialized question (science, history, \
etc.), use the ask_expert tool to consult the expert.",
)];
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;
}
messages.push(Message::user(input));
let mut tool_rounds = 0;
loop {
tool_rounds += 1;
let turn = run_turn(&provider, &messages, &tool_schemas, &mode).await?;
let calls: Vec<ToolCall> = turn
.messages
.iter()
.flat_map(|m| match m {
Message::Assistant { tool_calls, .. } => tool_calls.to_vec(),
_ => Vec::new(),
})
.collect();
if calls.is_empty() {
messages.extend(turn.messages);
break;
}
if tool_rounds >= MAX_TOOL_ROUNDS {
println!("(reached max tool rounds {MAX_TOOL_ROUNDS}, stopping tool calls)");
messages.extend(turn.messages);
break;
}
messages.extend(turn.messages);
for call in calls {
println!(" → calling {}, arguments: {}", call.name, call.arguments);
let content = registry
.call(&call.name, &call.arguments, &SharedState::new())
.await
.unwrap_or_else(|e| e.to_string());
println!(" → {} returned: {content}", call.name);
messages.push(Message::ToolResult {
id: call.id,
content,
});
}
}
}
Ok(())
}
async fn expert_loop(
provider: &OpenAiProvider,
channel: MpscChannel,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let mut messages = vec![Message::system(
"You are a knowledgeable expert; answer the user's questions directly and do not use any tools.",
)];
loop {
let incoming = channel.recv().await?;
println!(" [expert] received question: {}", incoming.text());
messages.push(Message::user(incoming.text()));
let response = provider
.chat(ChatRequest {
messages: messages.clone(),
..Default::default()
})
.await?;
let content = match response.message {
Message::Assistant { content, .. } => content,
_ => unreachable!("the reply must be an Assistant message by contract"),
};
println!(" [expert] answer: {content}");
messages.push(Message::Assistant {
content: content.clone(),
reasoning: None,
tool_calls: vec![],
});
incoming.reply(content)?;
}
}
#[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())
}
}
async fn run_turn(
provider: &OpenAiProvider,
messages: &[Message],
tools: &[ToolSchema],
mode: &Mode,
) -> Result<Turn, Box<dyn std::error::Error>> {
let request = ChatRequest {
messages: messages.to_vec(),
tools: tools.to_vec(),
..Default::default()
};
match mode {
Mode::Stream => {
let mut stream = provider.stream_chat(request).await?;
let mut text = String::new();
let mut reasoning = String::new();
let mut calls = Vec::new();
let mut prefix_printed = false;
while let Some(event) = stream.next().await {
match event? {
StreamEvent::Delta(delta) => {
if !prefix_printed {
print!("assistant: ");
std::io::stdout().flush()?;
prefix_printed = true;
}
print!("{delta}");
std::io::stdout().flush()?;
text.push_str(&delta);
}
StreamEvent::Reasoning(chunk) => reasoning.push_str(&chunk),
StreamEvent::ToolCall {
id,
name,
arguments,
} => {
calls.push(ToolCall {
id,
name,
arguments,
});
}
StreamEvent::Done { .. } => {}
_ => {}
}
}
println!();
Ok(Turn {
messages: assemble_turn(text, reasoning, calls),
})
}
Mode::Chat => {
let response = provider.chat(request).await?;
let (content, reasoning, tool_calls) = match response.message {
Message::Assistant {
content,
reasoning,
tool_calls,
} => (content, reasoning, tool_calls),
_ => unreachable!("the reply must be an Assistant message by contract"),
};
if !content.is_empty() {
println!("assistant: {content}");
}
Ok(Turn {
messages: vec![Message::Assistant {
content,
reasoning,
tool_calls,
}],
})
}
}
}
fn assemble_turn(text: String, reasoning: String, calls: Vec<ToolCall>) -> Vec<Message> {
let mut messages = Vec::new();
if !text.is_empty() || !calls.is_empty() || !reasoning.is_empty() {
messages.push(Message::Assistant {
content: text,
reasoning: (!reasoning.is_empty()).then_some(reasoning),
tool_calls: calls,
});
}
messages
}