Skip to main content

a_agent/agent/
agent.rs

1use std::sync::{Arc, Mutex};
2
3use anyhow::{Context, Result};
4use tokio_util::sync::CancellationToken;
5
6use crate::model::{ContentBlock, ModelMessage, ModelRequest, Role, StreamEvent};
7use crate::provider::{EventSink, Provider};
8use crate::session::SessionStore;
9use crate::tools::runner::ToolRunner;
10
11pub struct Agent {
12    provider: Arc<dyn Provider>,
13    tools: Arc<ToolRunner>,
14    store: Arc<Mutex<SessionStore>>,
15    session_id: String,
16    system_prompt: String,
17    max_cycles: usize,
18    context_window: Option<u64>,
19    max_output_tokens: u64,
20}
21
22#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct AgentResult {
24    pub final_text: Option<String>,
25    pub cycles: usize,
26}
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub struct ContextStatus {
30    pub used_tokens: u64,
31    pub provider_tokens: Option<u64>,
32    pub estimated_tokens: u64,
33    pub context_window: Option<u64>,
34    pub compact_at: Option<u64>,
35    pub max_output_tokens: u64,
36}
37
38struct ContextEstimate {
39    total: u64,
40    provider: Option<u64>,
41    estimated: u64,
42}
43
44impl Agent {
45    pub fn new(
46        provider: Arc<dyn Provider>,
47        tools: Arc<ToolRunner>,
48        store: Arc<Mutex<SessionStore>>,
49        session_id: String,
50        system_prompt: String,
51        max_cycles: usize,
52    ) -> Self {
53        Self {
54            provider,
55            tools,
56            store,
57            session_id,
58            system_prompt,
59            max_cycles,
60            context_window: None,
61            max_output_tokens: 0,
62        }
63    }
64
65    pub fn with_context_budget(
66        mut self,
67        context_window: Option<u64>,
68        max_output_tokens: u64,
69    ) -> Self {
70        self.context_window = context_window;
71        self.max_output_tokens = max_output_tokens;
72        self
73    }
74
75    pub async fn submit(
76        &self,
77        prompt: &str,
78        events: EventSink,
79        cancel: CancellationToken,
80    ) -> Result<AgentResult> {
81        self.compact_if_needed(Some(prompt), &events, &cancel)
82            .await?;
83        self.store
84            .lock()
85            .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
86            .append_item(
87                &self.session_id,
88                Role::User,
89                vec![ContentBlock::Text(prompt.into())],
90            )?;
91
92        for cycle in 1..=self.max_cycles {
93            if cancel.is_cancelled() {
94                anyhow::bail!("agent turn cancelled");
95            }
96            self.compact_if_needed(None, &events, &cancel).await?;
97            let messages = self
98                .store
99                .lock()
100                .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
101                .active_branch(&self.session_id)?
102                .into_iter()
103                .map(|item| ModelMessage {
104                    role: item.role,
105                    blocks: item.blocks,
106                })
107                .collect();
108            let request = ModelRequest {
109                system_prompt: self.system_prompt.clone(),
110                messages,
111                include_tools: true,
112            };
113            events.emit(StreamEvent::GenerationStart);
114            let turn = self
115                .provider
116                .stream_turn(request, events.clone(), cancel.clone())
117                .await?;
118            if !turn.blocks.is_empty() {
119                self.store
120                    .lock()
121                    .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
122                    .append_assistant_item(&self.session_id, turn.blocks.clone(), turn.usage)?;
123            }
124            if let Some(state) = &turn.provider_state {
125                self.store
126                    .lock()
127                    .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
128                    .set_provider_state(&self.session_id, "continuation", state)?;
129            }
130            if turn.tool_calls.is_empty() {
131                events.emit(StreamEvent::Done);
132                return Ok(AgentResult {
133                    final_text: turn.final_text(),
134                    cycles: cycle,
135                });
136            }
137            for call in &turn.tool_calls {
138                events.emit(StreamEvent::ToolExecutionStart {
139                    id: call.id.clone(),
140                });
141            }
142            let results = self
143                .tools
144                .execute_with(turn.tool_calls, events.clone(), cancel.clone())
145                .await;
146            for result in results {
147                events.emit(StreamEvent::ToolExecutionEnd {
148                    id: result.call_id.clone(),
149                    result: result.clone(),
150                });
151                self.store
152                    .lock()
153                    .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
154                    .append_item(
155                        &self.session_id,
156                        Role::Tool,
157                        vec![ContentBlock::ToolResult(result)],
158                    )?;
159            }
160        }
161        Err(anyhow::anyhow!(
162            "maximum agent cycles ({}) exceeded",
163            self.max_cycles
164        ))
165        .context("maximum agent cycles reached before a final response")
166    }
167
168    async fn compact_if_needed(
169        &self,
170        additional_prompt: Option<&str>,
171        events: &EventSink,
172        cancel: &CancellationToken,
173    ) -> Result<()> {
174        let Some(context_window) = self.context_window else {
175            return Ok(());
176        };
177        let threshold = context_window.saturating_sub(self.max_output_tokens);
178        let branch = self
179            .store
180            .lock()
181            .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
182            .active_branch(&self.session_id)?;
183        if let Some(summary_index) = branch.iter().rposition(is_compaction_summary)
184            && !branch[summary_index + 1..].iter().any(has_valid_usage)
185        {
186            return Ok(());
187        }
188        if branch.is_empty()
189            || estimate_context_tokens(&self.system_prompt, &branch, additional_prompt).total
190                <= threshold
191        {
192            return Ok(());
193        }
194        self.compact_branch(branch, events, cancel).await
195    }
196
197    async fn compact_branch(
198        &self,
199        branch: Vec<crate::model::ConversationItem>,
200        events: &EventSink,
201        cancel: &CancellationToken,
202    ) -> Result<()> {
203        let messages = branch
204            .iter()
205            .map(|item| ModelMessage {
206                role: item.role,
207                blocks: item.blocks.clone(),
208            })
209            .collect();
210        let request = ModelRequest {
211            system_prompt: "Summarize this coding-agent conversation for continuation. Preserve user goals, decisions, modified files, tool results, failures, unresolved work, and exact technical constraints. Return only the compact summary and do not call tools.".into(),
212            messages,
213            include_tools: false,
214        };
215        events.emit(StreamEvent::GenerationStart);
216        let result = self
217            .provider
218            .stream_turn(request, EventSink::default(), cancel.clone())
219            .await;
220        events.emit(StreamEvent::Done);
221        let turn = result.context("compact conversation")?;
222        if !turn.tool_calls.is_empty() {
223            anyhow::bail!("compaction model attempted to call tools");
224        }
225        let summary = turn
226            .final_text()
227            .context("compaction model returned no summary")?;
228        self.store
229            .lock()
230            .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
231            .replace_branch_with_summary(&self.session_id, &summary)?;
232        Ok(())
233    }
234
235    pub async fn compact(&self, events: EventSink, cancel: CancellationToken) -> Result<bool> {
236        let branch = self
237            .store
238            .lock()
239            .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
240            .active_branch(&self.session_id)?;
241        if branch.is_empty() {
242            return Ok(false);
243        }
244        self.compact_branch(branch, &events, &cancel).await?;
245        Ok(true)
246    }
247
248    pub fn context_status(&self) -> Result<ContextStatus> {
249        let branch = self
250            .store
251            .lock()
252            .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
253            .active_branch(&self.session_id)?;
254        let estimate = estimate_context_tokens(&self.system_prompt, &branch, None);
255        Ok(ContextStatus {
256            used_tokens: estimate.total,
257            provider_tokens: estimate.provider,
258            estimated_tokens: estimate.estimated,
259            context_window: self.context_window,
260            compact_at: self
261                .context_window
262                .map(|window| window.saturating_sub(self.max_output_tokens)),
263            max_output_tokens: self.max_output_tokens,
264        })
265    }
266
267    pub fn record_interruption(&self) -> Result<()> {
268        self.store
269            .lock()
270            .map_err(|_| anyhow::anyhow!("session store lock poisoned"))?
271            .append_turn_interrupted(&self.session_id)?;
272        Ok(())
273    }
274}
275
276fn is_compaction_summary(item: &crate::model::ConversationItem) -> bool {
277    item.blocks.iter().any(|block| {
278        matches!(block, ContentBlock::Text(text) if text.starts_with(crate::session::CONVERSATION_SUMMARY_PREFIX))
279    })
280}
281
282fn has_valid_usage(item: &crate::model::ConversationItem) -> bool {
283    item.role == Role::Assistant
284        && item
285            .usage
286            .and_then(|usage| usage.context_tokens())
287            .is_some_and(|tokens| tokens > 0)
288}
289
290fn estimate_context_tokens(
291    system_prompt: &str,
292    branch: &[crate::model::ConversationItem],
293    additional_prompt: Option<&str>,
294) -> ContextEstimate {
295    let usage_anchor = branch.iter().enumerate().rev().find_map(|(index, item)| {
296        (item.role == Role::Assistant)
297            .then_some(item.usage)
298            .flatten()
299            .and_then(|usage| usage.context_tokens())
300            .filter(|tokens| *tokens > 0)
301            .map(|tokens| (index, tokens))
302    });
303    let (start, provider, mut estimated) = usage_anchor.map_or_else(
304        || {
305            let tools =
306                serde_json::to_string(&crate::provider::tool_definitions()).unwrap_or_default();
307            (
308                0,
309                None,
310                estimate_text_tokens(system_prompt) + estimate_text_tokens(&tools),
311            )
312        },
313        |(index, tokens)| (index + 1, Some(tokens), 0),
314    );
315    for item in &branch[start..] {
316        estimated += estimate_blocks_tokens(&item.blocks);
317    }
318    if let Some(prompt) = additional_prompt {
319        estimated += estimate_text_tokens(prompt);
320    }
321    ContextEstimate {
322        total: provider.unwrap_or(0) + estimated,
323        provider,
324        estimated,
325    }
326}
327
328fn estimate_blocks_tokens(blocks: &[ContentBlock]) -> u64 {
329    let characters = blocks
330        .iter()
331        .map(|block| match block {
332            ContentBlock::Text(text) | ContentBlock::Reasoning(text) => text.chars().count(),
333            ContentBlock::ToolCall(call) => {
334                call.name.chars().count() + call.arguments.chars().count()
335            }
336            ContentBlock::ToolResult(result) => result.output.chars().count(),
337        })
338        .sum::<usize>();
339    characters.div_ceil(4) as u64
340}
341
342fn estimate_text_tokens(text: &str) -> u64 {
343    text.chars().count().div_ceil(4) as u64
344}