use std::sync::Arc;
use anyhow::{Result, anyhow};
use rullama_call_policy::BudgetGuard;
use rullama_core::{ChatOptions, Provider};
use rullama_tool_runtime::ToolExecutor;
use crate::chat_agent::ChatAgent;
pub struct AgentBuilder {
provider: Option<Arc<dyn Provider>>,
executor: Option<Arc<dyn ToolExecutor>>,
options: ChatOptions,
system_prompt: Option<String>,
max_iterations: Option<usize>,
tool_concurrency: Option<usize>,
summarization_keep_tail: Option<usize>,
budget: Option<BudgetGuard>,
}
impl AgentBuilder {
pub fn new() -> Self {
Self {
provider: None,
executor: None,
options: ChatOptions::default(),
system_prompt: None,
max_iterations: None,
tool_concurrency: None,
summarization_keep_tail: None,
budget: None,
}
}
pub fn provider(mut self, p: Arc<dyn Provider>) -> Self {
self.provider = Some(p);
self
}
pub fn tools(mut self, e: Arc<dyn ToolExecutor>) -> Self {
self.executor = Some(e);
self
}
pub fn options(mut self, o: ChatOptions) -> Self {
self.options = o;
self
}
pub fn system(mut self, prompt: impl Into<String>) -> Self {
self.system_prompt = Some(prompt.into());
self
}
pub fn max_iterations(mut self, rounds: usize) -> Self {
self.max_iterations = Some(rounds);
self
}
pub fn tool_concurrency(mut self, n: usize) -> Self {
self.tool_concurrency = Some(n);
self
}
pub fn summarization_keep_tail(mut self, keep: usize) -> Self {
self.summarization_keep_tail = Some(keep);
self
}
pub fn budget(mut self, guard: BudgetGuard) -> Self {
self.budget = Some(guard);
self
}
pub fn build_chat_agent(self) -> Result<ChatAgent> {
let provider = self.provider.ok_or_else(|| {
anyhow!("AgentBuilder: `provider` is required; call .provider(...) before building")
})?;
let executor = self.executor.ok_or_else(|| {
anyhow!("AgentBuilder: `tools` is required; call .tools(...) before building")
})?;
let mut agent = ChatAgent::new(provider, executor, self.options);
if let Some(sp) = self.system_prompt {
agent = agent.with_system_prompt(&sp);
}
if let Some(n) = self.max_iterations {
agent = agent.with_max_tool_rounds(n);
}
if let Some(n) = self.tool_concurrency {
agent = agent.with_tool_concurrency(n);
}
if let Some(n) = self.summarization_keep_tail {
agent = agent.with_summarization_keep_tail(n);
}
if let Some(guard) = self.budget {
agent = agent.with_budget(guard);
}
Ok(agent)
}
}
impl Default for AgentBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use rullama_tool_runtime::ToolRegistry;
fn fake_executor() -> Arc<dyn ToolExecutor> {
use rullama_core::ToolContext;
use rullama_tool_builtins::BuiltinToolExecutor;
Arc::new(BuiltinToolExecutor::new(
ToolRegistry::new(),
ToolContext::default(),
))
}
#[test]
fn missing_provider_errors() {
let result = AgentBuilder::new()
.tools(fake_executor())
.build_chat_agent();
let err = match result {
Ok(_) => panic!("expected build to fail without a provider"),
Err(e) => e.to_string(),
};
assert!(err.contains("`provider` is required"), "got: {err}");
}
#[test]
fn missing_tools_errors() {
let provider = Arc::new(rullama_test_fixtures::ScriptedProvider::always_text(
"test", "hi",
)) as Arc<dyn Provider>;
let result = AgentBuilder::new().provider(provider).build_chat_agent();
let err = match result {
Ok(_) => panic!("expected build to fail without tools"),
Err(e) => e.to_string(),
};
assert!(err.contains("`tools` is required"), "got: {err}");
}
#[test]
fn happy_path_builds() {
let provider = Arc::new(rullama_test_fixtures::ScriptedProvider::always_text(
"test", "hi",
)) as Arc<dyn Provider>;
let _agent = AgentBuilder::new()
.provider(provider)
.tools(fake_executor())
.system("you are helpful")
.max_iterations(20)
.tool_concurrency(2)
.build_chat_agent()
.expect("builder should succeed with provider + tools");
}
}