1use crate::{
31 completion::{CompletionModel, CompletionRequest, ModelChoice},
32 message::{AssistantContent, Message, ToolCall, ToolResult, UserContent},
33 tool::{ToolError, ToolSet},
34};
35use thiserror::Error;
36
37#[derive(Debug, Error)]
45pub enum AgentError<M: std::error::Error + 'static> {
46 #[error("Completion error: {0}")]
47 Completion(M),
48
49 #[error("Tool error: {0}")]
50 Tool(ToolError),
51
52 #[error("Exceeded maximum iterations ({0}) without a final text reply")]
54 MaxIterations(u32),
55}
56
57pub struct Agent<M> {
64 model: M,
65 preamble: Option<String>,
66 tools: ToolSet,
67 temperature: Option<f64>,
68 max_tokens: Option<u32>,
69 max_iterations: u32,
71 context: Vec<String>,
73}
74
75impl<M: CompletionModel> Agent<M> {
76 pub fn builder(model: M) -> AgentBuilder<M> {
78 AgentBuilder::new(model)
79 }
80
81 pub async fn prompt(&self, prompt: &str) -> Result<String, AgentError<M::Error>> {
86 self.chat(prompt, vec![]).await
87 }
88
89 pub async fn chat(
93 &self,
94 prompt: &str,
95 history: Vec<Message>,
96 ) -> Result<String, AgentError<M::Error>> {
97 let mut messages = self.build_messages(prompt, history);
98
99 for _ in 0..self.max_iterations {
100 let request = self.build_request(messages.clone());
101 let response = self.model.complete(request).await.map_err(AgentError::Completion)?;
102
103 match response.choice {
104 ModelChoice::Message(text) => return Ok(text),
105 ModelChoice::ToolCall(calls) => {
106 messages.push(Message::Assistant {
108 content: calls
109 .iter()
110 .map(|c| AssistantContent::ToolCall(c.clone()))
111 .collect(),
112 });
113
114 let mut results: Vec<UserContent> = Vec::with_capacity(calls.len());
116 for call in &calls {
117 let result = self.dispatch_tool(call).await;
118 results.push(UserContent::ToolResult(result));
119 }
120
121 messages.push(Message::User { content: results });
123 }
124 }
125 }
126
127 Err(AgentError::MaxIterations(self.max_iterations))
128 }
129
130 fn build_messages(&self, prompt: &str, mut history: Vec<Message>) -> Vec<Message> {
133 let mut messages: Vec<Message> = Vec::new();
134
135 if let Some(preamble) = &self.preamble {
136 messages.push(Message::system(preamble));
137 }
138
139 if !self.context.is_empty() {
142 let combined = self.context.join("\n\n");
143 messages.push(Message::user(combined));
144 }
145
146 messages.append(&mut history);
147 messages.push(Message::user(prompt));
148 messages
149 }
150
151 fn build_request(&self, messages: Vec<Message>) -> CompletionRequest {
152 let mut req = CompletionRequest::new(messages);
153 req.tools = self.tools.definitions();
154 req.temperature = self.temperature;
155 req.max_tokens = self.max_tokens;
156 req
157 }
158
159 async fn dispatch_tool(&self, call: &ToolCall) -> ToolResult {
160 let result = self.tools.call(&call.name, call.arguments.clone()).await;
161 let content = match result {
162 Ok(v) => v.to_string(),
163 Err(e) => format!("Error: {e}"),
165 };
166 ToolResult {
167 call_id: call.id.clone(),
168 name: call.name.clone(),
169 content,
170 }
171 }
172}
173
174pub struct AgentBuilder<M> {
178 model: M,
179 preamble: Option<String>,
180 tools: ToolSet,
181 temperature: Option<f64>,
182 max_tokens: Option<u32>,
183 max_iterations: u32,
184 context: Vec<String>,
185}
186
187impl<M: CompletionModel> AgentBuilder<M> {
188 fn new(model: M) -> Self {
189 Self {
190 model,
191 preamble: None,
192 tools: ToolSet::new(),
193 temperature: None,
194 max_tokens: None,
195 max_iterations: 10,
196 context: Vec::new(),
197 }
198 }
199
200 pub fn preamble(mut self, preamble: impl Into<String>) -> Self {
202 self.preamble = Some(preamble.into());
203 self
204 }
205
206 pub fn tool<T: crate::tool::Tool + 'static>(mut self, tool: T) -> Self {
208 self.tools.add(tool);
209 self
210 }
211
212 pub fn temperature(mut self, temperature: f64) -> Self {
214 self.temperature = Some(temperature);
215 self
216 }
217
218 pub fn max_tokens(mut self, max_tokens: u32) -> Self {
220 self.max_tokens = Some(max_tokens);
221 self
222 }
223
224 pub fn max_iterations(mut self, n: u32) -> Self {
226 self.max_iterations = n;
227 self
228 }
229
230 pub fn context(mut self, doc: impl Into<String>) -> Self {
232 self.context.push(doc.into());
233 self
234 }
235
236 pub fn build(self) -> Agent<M> {
238 Agent {
239 model: self.model,
240 preamble: self.preamble,
241 tools: self.tools,
242 temperature: self.temperature,
243 max_tokens: self.max_tokens,
244 max_iterations: self.max_iterations,
245 context: self.context,
246 }
247 }
248}