Skip to main content

aether_core/core/
agent_builder.rs

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
14/// Handle for communicating with a running Agent
15pub struct AgentHandle {
16    handle: JoinHandle<()>,
17}
18
19impl AgentHandle {
20    /// Abort the agent task immediately.
21    pub fn abort(&self) {
22        self.handle.abort();
23    }
24
25    /// Returns `true` if the agent task has finished.
26    pub fn is_finished(&self) -> bool {
27        self.handle.is_finished()
28    }
29
30    /// Wait for the agent task to complete.
31    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    /// Create a builder from a resolved `AgentSpec`.
82    ///
83    /// The LLM provider is derived from `spec.model` via `ModelProviderParser`.
84    /// `base_prompts` are prepended before the spec's own prompts.
85    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    /// Add a prompt to the system prompt.
116    ///
117    /// Multiple prompts are concatenated with double newlines.
118    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    /// Set the timeout for tool execution
130    ///
131    /// If a tool does not return a result within this duration, it will be marked as failed
132    /// and the agent will continue processing.
133    ///
134    /// Default: 60 minutes
135    pub fn tool_timeout(mut self, timeout: Duration) -> Self {
136        self.tool_timeout = timeout;
137        self
138    }
139
140    /// Configure context compaction settings.
141    ///
142    /// By default, agents automatically compact context when token usage exceeds
143    /// 85% of the context window, preventing overflow during long-running tasks.
144    ///
145    /// # Examples
146    /// ```ignore
147    /// // Custom threshold
148    /// agent(llm).compaction(CompactionConfig::with_threshold(0.9))
149    ///
150    /// // Disable compaction entirely
151    /// agent(llm).compaction(CompactionConfig::disabled())
152    ///
153    /// // Full customization
154    /// agent(llm).compaction(
155    ///     CompactionConfig::with_threshold(0.85)
156    ///         .keep_recent_tool_results(3)
157    ///         .min_messages(20)
158    /// )
159    /// ```
160    pub fn compaction(mut self, config: CompactionConfig) -> Self {
161        self.compaction_config = Some(config);
162        self
163    }
164
165    /// Disable context compaction entirely.
166    ///
167    /// Overflow errors from the model will be surfaced directly to callers.
168    pub fn disable_compaction(mut self) -> Self {
169        self.compaction_config = None;
170        self
171    }
172
173    /// Configure the maximum number of auto-continue attempts.
174    ///
175    /// When the LLM stops without making tool calls, the agent may inject a
176    /// continuation prompt and restart the LLM stream for resumable stop
177    /// reasons (for example, token length limits).
178    ///
179    /// This setting limits how many times the agent will attempt to continue
180    /// before giving up and ending the turn with [`TurnEvent::Ended`](crate::events::TurnEvent::Ended).
181    ///
182    /// Default: 3
183    ///
184    /// # Example
185    /// ```ignore
186    /// // Allow up to 5 auto-continue attempts
187    /// agent(llm).max_auto_continues(5)
188    ///
189    /// // Disable auto-continue entirely
190    /// agent(llm).max_auto_continues(0)
191    /// ```
192    pub fn max_auto_continues(mut self, max: u32) -> Self {
193        self.max_auto_continues = max;
194        self
195    }
196
197    /// Configure retry behavior for transient LLM provider failures.
198    pub fn retry(mut self, config: RetryConfig) -> Self {
199        self.retry_config = config;
200        self
201    }
202
203    /// Override the effective model context window in tokens.
204    pub fn context_window(mut self, context_window: Option<u32>) -> Self {
205        self.context_window = context_window;
206        self
207    }
208
209    /// Set the sampling controls (`temperature`, `top_p`, `max_tokens`) applied to
210    /// every model call this agent makes.
211    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    /// Pre-populate the context with conversation history (e.g. from a restored session).
222    ///
223    /// These messages are inserted after the system prompt.
224    pub fn messages(mut self, messages: Vec<ChatMessage>) -> Self {
225        self.initial_messages = messages;
226        self
227    }
228
229    /// Attach an observer of the agent's event stream.
230    pub fn observer(mut self, observer: Box<dyn AgentObserver>) -> Self {
231        self.observers.push(observer);
232        self
233    }
234
235    /// Record usage under `tracker`, which names this agent in usage events.
236    pub fn session_usage(mut self, tracker: SessionUsageTracker) -> Self {
237        self.session_usage = tracker;
238        self
239    }
240
241    /// Continue session totals from the last persisted usage event, e.g. when
242    /// resuming a session.
243    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}