1use super::agent::{AgentConfig, AutoContinue, RetryConfig};
2use crate::agent_spec::AgentSpec;
3use crate::context::{CompactionConfig, SessionUsageTracker};
4use crate::core::{Agent, AgentDeps, Prompt, PromptCache, Result};
5use crate::events::{AgentEvent, AgentObserver, Command, TurnOutcome};
6use crate::mcp::McpHandle;
7use llm::parser::ModelProviderParser;
8use llm::{ChatMessage, Context, ModelSettings, SessionUsageEvent, StreamingModelProvider, ToolDefinition};
9use std::sync::Arc;
10use std::time::Duration;
11use tokio::sync::mpsc::{self, Receiver, Sender};
12use tokio::task::JoinHandle;
13
14pub struct AgentHandle {
16 handle: JoinHandle<()>,
17}
18
19impl AgentHandle {
20 pub fn abort(&self) {
22 self.handle.abort();
23 }
24
25 pub fn is_finished(&self) -> bool {
27 self.handle.is_finished()
28 }
29
30 pub async fn await_completion(self) {
32 let _ = self.handle.await;
33 }
34}
35
36pub async fn recv_agent_event(rx: &mut Receiver<AgentEvent>) -> AgentEvent {
37 rx.recv()
38 .await
39 .unwrap_or_else(|| AgentEvent::turn_ended(TurnOutcome::failed("Agent stopped before its turn ended")))
40}
41
42pub struct AgentBuilder {
43 llm: Arc<dyn StreamingModelProvider>,
44 prompts: Vec<Prompt>,
45 tool_definitions: Vec<ToolDefinition>,
46 initial_messages: Vec<ChatMessage>,
47 mcp: Option<McpHandle>,
48 channel_capacity: usize,
49 tool_timeout: Duration,
50 compaction_config: Option<CompactionConfig>,
51 max_auto_continues: u32,
52 retry_config: RetryConfig,
53 context_window: Option<u32>,
54 model_settings: ModelSettings,
55 observers: Vec<Box<dyn AgentObserver>>,
56 session_usage: SessionUsageTracker,
57 session_affinity_key: String,
58}
59
60impl AgentBuilder {
61 pub fn new(llm: Arc<dyn StreamingModelProvider>) -> Self {
62 Self {
63 llm,
64 prompts: Vec::new(),
65 tool_definitions: Vec::new(),
66 initial_messages: Vec::new(),
67 mcp: None,
68 channel_capacity: 1000,
69 tool_timeout: Duration::from_mins(60),
70 compaction_config: Some(CompactionConfig::default()),
71 max_auto_continues: 3,
72 retry_config: RetryConfig::default(),
73 context_window: None,
74 model_settings: ModelSettings::default(),
75 observers: Vec::new(),
76 session_usage: SessionUsageTracker::new("agent"),
77 session_affinity_key: uuid::Uuid::new_v4().to_string(),
78 }
79 }
80
81 pub async fn from_spec(spec: &AgentSpec, base_prompts: Vec<Prompt>, deps: &AgentDeps) -> Result<Self> {
86 let parser = ModelProviderParser::default().with_provider_connections(spec.provider_connections.clone());
87 let parser = match deps.oauth_credential_store.clone() {
88 Some(store) => parser.with_codex_provider(store),
89 None => parser,
90 };
91 let (provider, _) = parser.parse(&spec.model).await?;
92 let mut builder = Self::new(Arc::from(provider))
93 .context_window(spec.context_window)
94 .model_settings(spec.model_settings.clone())
95 .session_usage(SessionUsageTracker::new(&spec.name));
96
97 if let Some(key) = &deps.session_affinity_key {
98 builder = builder.session_affinity_key(key.clone());
99 }
100 if let Some(observer) = deps.observer(&spec.name) {
101 builder = builder.observer(observer);
102 }
103
104 for prompt in base_prompts {
105 builder = builder.system_prompt(prompt);
106 }
107
108 for prompt in &spec.prompts {
109 builder = builder.system_prompt(prompt.clone());
110 }
111
112 Ok(builder)
113 }
114
115 pub fn system_prompt(mut self, prompt: Prompt) -> Self {
119 self.prompts.push(prompt);
120 self
121 }
122
123 pub fn tools(mut self, mcp: McpHandle, tools: Vec<ToolDefinition>) -> Self {
124 self.tool_definitions = tools;
125 self.mcp = Some(mcp);
126 self
127 }
128
129 pub fn tool_timeout(mut self, timeout: Duration) -> Self {
136 self.tool_timeout = timeout;
137 self
138 }
139
140 pub fn compaction(mut self, config: CompactionConfig) -> Self {
161 self.compaction_config = Some(config);
162 self
163 }
164
165 pub fn disable_compaction(mut self) -> Self {
169 self.compaction_config = None;
170 self
171 }
172
173 pub fn max_auto_continues(mut self, max: u32) -> Self {
193 self.max_auto_continues = max;
194 self
195 }
196
197 pub fn retry(mut self, config: RetryConfig) -> Self {
199 self.retry_config = config;
200 self
201 }
202
203 pub fn context_window(mut self, context_window: Option<u32>) -> Self {
205 self.context_window = context_window;
206 self
207 }
208
209 pub fn model_settings(mut self, model_settings: ModelSettings) -> Self {
212 self.model_settings = model_settings;
213 self
214 }
215
216 pub fn session_affinity_key(mut self, key: impl Into<String>) -> Self {
217 self.session_affinity_key = key.into();
218 self
219 }
220
221 pub fn messages(mut self, messages: Vec<ChatMessage>) -> Self {
225 self.initial_messages = messages;
226 self
227 }
228
229 pub fn observer(mut self, observer: Box<dyn AgentObserver>) -> Self {
231 self.observers.push(observer);
232 self
233 }
234
235 pub fn session_usage(mut self, tracker: SessionUsageTracker) -> Self {
237 self.session_usage = tracker;
238 self
239 }
240
241 pub fn resume_usage(mut self, last: &SessionUsageEvent) -> Self {
244 self.session_usage.resume_from(last);
245 self
246 }
247
248 pub async fn spawn(self) -> Result<(Sender<Command>, Receiver<AgentEvent>, AgentHandle)> {
249 let mut prompt_cache = PromptCache::new(self.prompts);
250 let system_content = prompt_cache.render().await?;
251 let mut messages = Vec::new();
252
253 if !system_content.is_empty() {
254 messages.push(ChatMessage::system(system_content));
255 }
256
257 messages.extend(self.initial_messages);
258 let (command_tx, command_rx) = mpsc::channel::<Command>(self.channel_capacity);
259 let (message_tx, agent_event_rx) = mpsc::channel::<AgentEvent>(self.channel_capacity);
260 let mut context = Context::new(messages, self.tool_definitions);
261 context.set_model_settings(self.model_settings);
262 context.set_session_affinity_key(Some(self.session_affinity_key));
263
264 let config = AgentConfig {
265 llm: self.llm,
266 context,
267 mcp: self.mcp,
268 tool_timeout: self.tool_timeout,
269 compaction_config: self.compaction_config,
270 auto_continue: AutoContinue::new(self.max_auto_continues),
271 retry_config: self.retry_config,
272 context_window: self.context_window,
273 prompt_cache,
274 observers: self.observers,
275 session_usage: self.session_usage,
276 };
277
278 let agent = Agent::new(config, command_rx, message_tx);
279 let agent_handle = tokio::spawn(agent.run());
280
281 Ok((command_tx, agent_event_rx, AgentHandle { handle: agent_handle }))
282 }
283}
284
285#[cfg(test)]
286mod tests {
287 use super::*;
288 use crate::agent_spec::AgentSpecExposure;
289 use crate::events::TurnEvent;
290 use llm::ProviderConnectionOverrides;
291 use mcp_utils::client::ToolFilter;
292
293 #[tokio::test]
294 async fn test_agent_handle_is_finished() {
295 let handle = AgentHandle { handle: tokio::spawn(async {}) };
296 handle.await_completion().await;
297 }
298
299 #[tokio::test]
300 async fn test_agent_handle_abort() {
301 let handle = AgentHandle { handle: tokio::spawn(std::future::pending::<()>()) };
302 assert!(!handle.is_finished());
303 handle.abort();
304 while !handle.is_finished() {
305 tokio::task::yield_now().await;
306 }
307 }
308
309 #[tokio::test]
310 async fn recv_agent_event_returns_events_from_open_channel() {
311 let (tx, mut rx) = mpsc::channel(4);
312 tx.send(AgentEvent::Turn(TurnEvent::Started { content: vec![] })).await.unwrap();
313
314 assert_eq!(recv_agent_event(&mut rx).await, AgentEvent::Turn(TurnEvent::Started { content: vec![] }));
315 }
316
317 #[tokio::test]
318 async fn recv_agent_event_returns_failed_turn_end_once_channel_closes() {
319 let (tx, mut rx) = mpsc::channel::<AgentEvent>(4);
320 drop(tx);
321
322 let event = recv_agent_event(&mut rx).await;
323
324 assert!(matches!(event.turn_outcome(), Some(TurnOutcome::Failed { .. })), "got: {event:?}");
325 }
326
327 #[tokio::test]
328 async fn system_prompt_preserves_add_order() {
329 let builder = AgentBuilder::new(Arc::new(llm::testing::FakeLlmProvider::new(vec![])))
330 .system_prompt(Prompt::text("first"))
331 .system_prompt(Prompt::text("second"))
332 .system_prompt(Prompt::text("third"));
333
334 let rendered = Prompt::build_all(&builder.prompts).await.unwrap();
335
336 assert_eq!(rendered, "first\n\nsecond\n\nthird");
337 }
338
339 #[tokio::test]
340 async fn from_spec_applies_context_window_and_model_settings() {
341 let settings = ModelSettings { temperature: Some(0.0), max_tokens: Some(128), ..Default::default() };
342 let spec = AgentSpec {
343 name: "alloy".to_string(),
344 description: "alloy".to_string(),
345 model: "ollama:llama3.2,llamacpp:local".to_string(),
346 reasoning_effort: None,
347 model_settings: settings.clone(),
348 context_window: Some(200_000),
349 prompts: vec![],
350 provider_connections: ProviderConnectionOverrides::default(),
351 mcp_config_sources: Vec::new(),
352 exposure: AgentSpecExposure::both(),
353 tools: ToolFilter::default(),
354 };
355
356 let dependencies = AgentDeps::default().with_session_affinity_key("conversation-123");
357 let builder = AgentBuilder::from_spec(&spec, vec![], &dependencies).await.unwrap();
358
359 assert_eq!(builder.context_window, Some(200_000));
360 assert_eq!(builder.model_settings, settings);
361 assert_eq!(builder.session_affinity_key, "conversation-123");
362 }
363
364 #[tokio::test]
365 async fn from_spec_accepts_alloy_model_specs() {
366 let spec = AgentSpec {
367 name: "alloy".to_string(),
368 description: "alloy".to_string(),
369 model: "ollama:llama3.2,llamacpp:local".to_string(),
370 reasoning_effort: None,
371 model_settings: ModelSettings::default(),
372 context_window: None,
373 prompts: vec![],
374 provider_connections: ProviderConnectionOverrides::default(),
375 mcp_config_sources: Vec::new(),
376 exposure: AgentSpecExposure::both(),
377 tools: ToolFilter::default(),
378 };
379
380 let builder = AgentBuilder::from_spec(&spec, vec![], &AgentDeps::default()).await;
381 assert!(builder.is_ok());
382 }
383}